diff --git a/backend/src/api/nft_manager.py b/backend/src/api/nft_manager.py index 707cb9a..e0cd803 100644 --- a/backend/src/api/nft_manager.py +++ b/backend/src/api/nft_manager.py @@ -1,226 +1,324 @@ # app.py from typing import Any, Dict, List, Optional, Tuple, Union from fastapi import FastAPI, APIRouter, HTTPException, status -from fastapi.responses import JSONResponse +from pydantic import BaseModel, Field import logging import json import re -# python-nftables wrapper +# libnftables (we call textual commands through its .cmd() method) from nftables import Nftables # type: ignore +# ---------- logging ---------- logging.basicConfig(level=logging.INFO) -logger = logging.getLogger("nft_api") +logger = logging.getLogger("nft_api_raw_only") -# ---------- Errors ---------- +# ---------- Exceptions ---------- class NftError(RuntimeError): pass -# ---------- 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 ---------- +# ---------- NftManager (textual-only) ---------- 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 + Thin wrapper around python-nftables exposing: + - cmd execution (textual nft commands via Nftables.cmd()) + - convenience list_rules_text / list_rules_json / list_chain_text + We prefer JSON globally, but for per-chain textual listing we temporarily disable JSON + so the output matches `nft list chain ...` textual rule lines. """ def __init__(self) -> None: self.nft = Nftables() - # best-effort: prefer JSON output where appropriate, but we'll explicitly choose cmd/json_cmd + # Try to prefer JSON for general listing; we'll toggle off for chain-list calls. try: - if hasattr(self.nft, "set_json_output"): - # do not rely on global toggles for each call - self.nft.set_json_output(True) + self.nft.set_json_output(True) except Exception: - logger.debug("set_json_output not available") + logger.debug("set_json_output not available or ignored") - def cmd(self, text_cmd: str) -> Tuple[int, str, str]: + def cmd(self, text_cmd: str) -> Dict[str, Optional[str]]: """ - Execute a textual nft command via Nftables.cmd() - Returns (rc, stdout, stderr) where stdout/stderr are strings (possibly empty). + Execute a textual nft command via Nftables.cmd(). + Returns dict { "rc": rc, "stdout": out_str, "stderr": err_str }. """ 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]: - """ - 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, falling back to -j: %s", e) - - # 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) - - def list_rules_text(self) -> str: - rc, out, err = self.cmd("list ruleset") if rc != 0: - raise NftError(f"nft list ruleset failed: {err}") - return out + logger.warning("nft cmd rc=%s stderr=%s cmd=%s", rc, err, text_cmd) + return {"rc": rc, "stdout": out, "stderr": err} + + def list_rules(self) -> str: + """ + Return the textual ruleset as produced by 'nft list ruleset'. + Uses the textual command path. + """ + res = self.cmd("list ruleset") + if res["rc"] != 0: + raise NftError(f"nft list ruleset failed: {res['stderr']}") + return res["stdout"] 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. + Try to obtain nft -j list ruleset (JSON). Returns parsed JSON dict on success. + Raises NftError on failure or when output cannot be parsed as JSON. """ - rc, out, err = self.cmd("list ruleset -j") - if rc != 0: - raise NftError(f"nft list ruleset failed (json): {err}") - try: - 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 + cmd_variants = ["list ruleset -j", "list ruleset"] + last_err = None + for c in cmd_variants: + res = self.cmd(c) + if res["rc"] != 0: + last_err = res["stderr"] + continue + out = res["stdout"] + if not out: + last_err = "empty output" + continue + try: + parsed = json.loads(out) + return parsed + except json.JSONDecodeError as e: + last_err = f"json decode error: {e}" + continue + raise NftError(f"unable to get JSON ruleset: {last_err}") def list_chain_text(self, family: str, table: str, chain: str) -> str: """ - Prefer textual listing via cmd(). If textual listing fails or returns JSON, - attempt to use json_cmd() and reconstruct readable lines. + Return textual output of `nft list chain `. + This tries to temporarily disable JSON output so the wrapper returns the textual + representation used by `nft list ruleset`. If disabling JSON is not possible, + we attempt to parse returned JSON (as a last resort), but the preferred path is + to get textual output. """ cmd = f"list chain {family} {table} {chain}" - rc, out, err = self.cmd(cmd) - if rc == 0 and out and not (out.strip().startswith("{") or out.strip().startswith("[")): - # Looks like proper textual output - return out - - # Either cmd returned error, or returned JSON-like text; try JSON path - rcj, outj, errj = self.json_cmd(cmd) - if rcj != 0: - # prefer the textual error if present - raise NftError(f"nft {cmd} failed: {err or errj}") - + # Attempt to temporarily disable JSON output on the wrapper (best-effort). + json_toggled = False + res = {"rc": -1, "stdout": "", "stderr": "unknown"} try: - parsed = _safe_load_json_unwrap(outj) - except json.JSONDecodeError: - # if parsing fails, return whatever textual output we have (maybe empty) - return out + if hasattr(self.nft, "set_json_output"): + try: + # Turn off JSON output to force textual output for this call. + self.nft.set_json_output(False) + json_toggled = True + except Exception: + logger.debug("could not toggle set_json_output(False); will try command anyway") + res = self.cmd(cmd) + finally: + # Restore JSON output preference if we toggled it. + if json_toggled and hasattr(self.nft, "set_json_output"): + try: + self.nft.set_json_output(True) + except Exception: + logger.debug("failed to restore set_json_output(True)") - # 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 + if res["rc"] != 0: + raise NftError(f"nft {cmd} failed: {res['stderr']}") - lines: List[str] = [] - for rec in recs: - if "rule" in rec: - r = rec["rule"] - expr = r.get("expr") - handle = r.get("handle") - parts: List[str] = [] - if isinstance(expr, list): - for el in expr: - if isinstance(el, dict): - if "match" in el: - m = el["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"] + out = res["stdout"] or "" + # If the output looks like JSON (starts with '{' or '['), try a safe fallback: + s = out.strip() + if s.startswith("{") or s.startswith("["): + # Best-effort: parse JSON and attempt to extract rule textual forms if present. + try: + parsed = json.loads(s) + # parsed may be the whole ruleset (nftables list) or a list; find any "rule" objects + rule_lines: List[str] = [] + # parsed might be dict with "nftables" or a list of records + records = parsed.get("nftables") if isinstance(parsed, dict) else parsed + if not isinstance(records, list): + records = [] + for rec in records: + if "rule" in rec: + r = rec["rule"] + expr = r.get("expr") + if isinstance(expr, list): + tokens: List[str] = [] + for part in expr: + if isinstance(part, dict) and "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): + p = left["payload"] + prot = p.get("protocol") + field = p.get("field") + if prot and field: + tokens.append(f"{prot} {field} {right}") + continue + tokens.append("match") + elif isinstance(part, dict) and "payload" in part: + p = part["payload"] prot = p.get("protocol") field = p.get("field") - if prot and field: - parts.append(f"{prot} {field} {right}") - continue - parts.append("match") - continue - if "payload" in el: - p = el["payload"] - prot = p.get("protocol"); field = p.get("field") - if prot and field: - parts.append(f"payload({prot}.{field})") - continue - parts.append("payload") - continue - if "drop" in el: - parts.append("drop"); continue - if "accept" in el: - parts.append("accept"); continue - if "counter" in el: - parts.append("counter"); continue - if "queue" in el: - q = el["queue"] - tok = "queue" - if isinstance(q, dict): - 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": - tok += " bypass" - elif isinstance(q, (int, float)): - tok += f" num {int(q)}" - parts.append(tok); continue - parts.append("+".join(sorted(el.keys()))) + tokens.append(f"payload({prot}.{field})") + elif isinstance(part, dict) and "drop" in part: + tokens.append("drop") + elif isinstance(part, dict) and "accept" in part: + tokens.append("accept") + elif isinstance(part, dict) and "counter" in part: + tokens.append("counter") + elif isinstance(part, dict) and "queue" in part: + # handle fallback queue textualization + q = part["queue"] + if isinstance(q, dict): + num = q.get("num") or q.get("number") or q.get("queue_number") or q.get("from") or q.get("range") + tok = "queue" + if num is not None: + tok += f" num {num}" + if q.get("bypass"): + tok += " bypass" + tokens.append(tok) + else: + # numeric or string value + if isinstance(q, (int, float)): + tokens.append(f"queue num {int(q)}") + else: + tokens.append(f"queue num {q}") + else: + # fallback for unknown dict token + if isinstance(part, dict): + tokens.append("+".join(part.keys())) + else: + tokens.append(str(part)) + rule_lines.append(" ".join(tokens)) else: - parts.append(str(el)) - else: - parts.append(json.dumps(r)) - txt = " ".join([p for p in parts if p]).strip() - if handle is not None: - txt = f"{txt} # handle {handle}" - lines.append(txt) - return "\n".join(lines) if lines else out + rule_lines.append(json.dumps(r)) + if rule_lines: + return "\n".join(rule_lines) + except Exception: + logger.debug("fallback JSON parsing of chain output failed; returning raw output") + + return out + + def delete_rule_by_handle_text(self, family: str, table: str, chain: str, handle: int) -> None: + """ + Delete a rule by handle using textual nft command: + delete rule
handle + """ + if not isinstance(handle, int) or handle <= 0: + raise ValueError("handle must be a positive integer") + # construct textual command + cmd = f"delete rule {family} {table} {chain} handle {handle}" + res = self.cmd(cmd) + if res["rc"] != 0: + raise NftError(f"delete rule failed: {res['stderr']}") + + +# ---------- FastAPI + Router ---------- +app = FastAPI(title="Unrestricted nftables API (json create)") +router = APIRouter(prefix="/firewall", tags=["firewall"]) +mgr = NftManager() + + +# ---------- Request/Response models (strongly typed) ---------- +class RawCmdRequest(BaseModel): + cmd: str = Field(..., description="Textual nft command to execute", example="add rule inet filter input ip saddr 10.0.0.0/8 drop") + + +class ExecResult(BaseModel): + rc: int = Field(..., description="Return code from nft execution", example=0) + stdout: Optional[str] = Field(None, description="Standard output from nft", example="") + stderr: Optional[str] = Field(None, description="Standard error from nft", example="") + + class Config: + schema_extra = {"example": {"rc": 0, "stdout": "ok", "stderr": ""}} + + +# --- Strong models returned to frontend --- +class RuleOut(BaseModel): + handle: Optional[int] = Field(None, description="The rule handle (unique per rule), if available", example=3) + expr: Any = Field(..., description="Machine-readable nft expression (original nft JSON expr).") + text: str = Field(..., description="Deterministic short display string derived from expr", example="ip protocol icmp drop") + position: Optional[Any] = Field(None, description="Optional position metadata from nft if present") + comment: Optional[str] = Field(None, description="Optional comment attached to the rule") + + +class ChainOut(BaseModel): + name: str = Field(..., description="Chain name", example="forward") + type: Optional[str] = Field(None, description="Chain type (e.g. filter, nat, route)") + hook: Optional[str] = Field(None, description="Hook (input/forward/output/ingress/egress) if present") + priority: Optional[int] = Field(None, description="Hook priority if present") + policy: Optional[str] = Field(None, description="Chain policy (accept/drop) if present") + rules: List[RuleOut] = Field(..., description="Rules in this chain (ordered)") + + +class TableOut(BaseModel): + family: str = Field(..., description="Table family (inet/bridge/ipv4/...)") + name: str = Field(..., description="Table name", example="filter") + chains: List[ChainOut] = Field(..., description="Chains in this table") + + +class RulesetModel(BaseModel): + tables: List[TableOut] = Field(..., description="Top-level tables list") + + +# ---------- New: CreateRuleRequest (JSON, expr required) ---------- +class CreateRuleRequest(BaseModel): + family: str = Field(..., description="Table family (e.g. inet, bridge, ip, ip6)", example="bridge") + table: str = Field(..., description="Table name (e.g. filter)", example="filter") + chain: str = Field(..., description="Chain name (e.g. forward)", example="forward") + expr: Any = Field(..., description="nft JSON expression (machine-readable). This field is required for JSON rule creation.") + position: Optional[int] = Field(None, description="Optional insertion position (zero-based). If provided endpoint will insert at that position.") + comment: Optional[str] = Field(None, description="Optional comment") + + class Config: + schema_extra = { + "example": { + "family": "bridge", + "table": "filter", + "chain": "forward", + "expr": [{"match": {"left": {"payload": {"protocol": "ip", "field": "protocol"}}, "op": "==", "right": "icmp"}}, {"drop": None}], + "position": 0, + } + } + + +# ruleset may be typed RulesetModel or raw textual string (fallback) +RulesetValue = Optional[Union[RulesetModel, str]] + + +class RulesetOut(BaseModel): + ruleset: RulesetValue = Field( + None, + description="Parsed, strongly-typed ruleset (RulesetModel) or raw textual ruleset string if JSON is unavailable.", + ) + + +# ---------- Helpers to convert to desired shape ---------- +_handle_re = re.compile(r"\s+#\s*handle\s+\d+\s*$") -# ---------- Utilities to convert nft JSON -> predictable model ---------- def parse_priority(val: Any) -> Optional[int]: + """ + Robustly parse a priority value returned in various nft JSON shapes. + Accepts: + - int -> returns unchanged + - numeric string -> parsed int + - dict -> tries common nested keys ('priority', 'prio') + Returns None if not parseable. + """ if val is None: return None + # if it's already an int if isinstance(val, int): return val + # numeric string if isinstance(val, str): s = val.strip() + # try integer parse try: return int(s) except Exception: try: + # sometimes it's "0.0" or similar return int(float(s)) except Exception: return None + # nested dicts sometimes appear if isinstance(val, dict): + # look for common keys for key in ("priority", "prio"): if key in val: return parse_priority(val.get(key)) + # try nested dict values for v in val.values(): p = parse_priority(v) if p is not None: @@ -230,7 +328,9 @@ def parse_priority(val: Any) -> Optional[int]: def rule_text_from_expr(expr: Any) -> str: """ - Conservative compact textualifier for expr. + Deterministic serializer to produce a compact UI-friendly string from expr list. + Covers common constructs; falls back to JSON dump for unknown constructs. + (Used for display in GET /rules). """ if expr is None: return "" @@ -238,155 +338,333 @@ def rule_text_from_expr(expr: Any) -> str: tokens: List[str] = [] for part in expr: if isinstance(part, dict): + # common tokens 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") + m = part["match"] + left = m.get("left") + right = m.get("right") + if isinstance(left, dict) and "payload" in left and isinstance(right, str): + p = left["payload"] + prot = p.get("protocol") + field = p.get("field") if prot and 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") + p = part["payload"] + prot = p.get("protocol") + field = p.get("field") if prot and field: tokens.append(f"payload({prot}.{field})") continue tokens.append("payload") + elif "cmp" in part or "binary" in part: + tokens.append("cmp") elif "drop" in part: tokens.append("drop") elif "accept" in part: tokens.append("accept") elif "counter" in part: tokens.append("counter") + elif "tcp" in part or "udp" in part: + proto = "tcp" if "tcp" in part else "udp" + tokens.append(proto) elif "queue" in part: - q = part["queue"]; tok = "queue" + q = part["queue"] + token = "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": - tok += " bypass" + token += f" num {num}" + if q.get("bypass"): + token += " bypass" elif isinstance(q, (int, float)): - tok += f" num {int(q)}" - tokens.append(tok) + token += f" num {int(q)}" + elif isinstance(q, str): + token += f" num {q}" + tokens.append(token) else: - tokens.append("+".join(sorted(part.keys()))) + keys = "+".join(sorted(part.keys())) + tokens.append(keys) else: tokens.append(str(part)) return " ".join(tokens) - try: - return str(expr) - except Exception: - return json.dumps(expr) + return str(expr) def build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]: """ - Convert nft -j structure into deterministic dict with tables/chains/rules. + Convert nft -j list ruleset parsed JSON into a deterministic, predictable JSON: + { + "tables": [ + { "family": ..., "name": ..., "chains": [ { "name": ..., "type": ..., "hook": ..., "priority": ..., "policy": ..., "rules": [ { handle, expr, text } ] } ] } + ] + } """ result: Dict[str, Any] = {"tables": []} items = nft_json.get("nftables", []) if isinstance(nft_json, dict) else (nft_json or []) - tables: Dict[Tuple[str, str], Dict[str, Any]] = {} + # Build intermediate map: (family, table) -> {family, name, chains: {chain_name: {"name", "type", "hook", "priority", "policy", "rules":[]}}} + tables: Dict[Tuple[str, str], Dict[str, Any]] = {} for rec in items: + # table records if "table" in rec: - t = rec["table"]; fam = t.get("family"); name = t.get("name") + t = rec["table"] + fam = t.get("family") + name = t.get("name") if fam and name: tables.setdefault((fam, name), {"family": fam, "name": name, "chains": {}}) + # chain records: capture chain metadata elif "chain" in rec: - c = rec["chain"] - fam = c.get("family") or (c.get("table") or {}).get("family") - tname = c.get("table") or (c.get("table") or {}).get("name") - cname = c.get("name") - if fam and tname and cname: - tables.setdefault((fam, tname), {"family": fam, "name": tname, "chains": {}}) - chains = tables[(fam, tname)]["chains"] - existing = chains.get(cname) - ch_type = c.get("type"); ch_hook = c.get("hook") - ch_prio = parse_priority(c.get("priority") if "priority" in c else c.get("prio") if "prio" in c else c.get("priority")) - ch_policy = c.get("policy") + ch = rec["chain"] + # chain may include family/table or nested table reference + fam = ch.get("family") or (ch.get("table") or {}).get("family") + table_name = ch.get("table") or (ch.get("table") or {}).get("name") + cname = ch.get("name") + if fam and table_name and cname: + tables.setdefault((fam, table_name), {"family": fam, "name": table_name, "chains": {}}) + chains_map = tables[(fam, table_name)]["chains"] + + # existing chain (maybe created earlier by rule processing) + existing = chains_map.get(cname) + # extract metadata robustly + ch_type = ch.get("type") + ch_hook = ch.get("hook") + # try multiple keys for priority/prio shapes + ch_priority = parse_priority(ch.get("priority") if "priority" in ch else ch.get("prio") if "prio" in ch else ch.get("priority", None)) + # also attempt to parse nested shapes if present (some nft JSON variations) + if ch_priority is None: + ch_priority = parse_priority(ch.get("hook") if isinstance(ch.get("hook"), dict) else None) + + ch_policy = ch.get("policy") + if existing is None: - chains[cname] = {"name": cname, "type": ch_type, "hook": ch_hook, "priority": ch_prio, "policy": ch_policy, "rules": []} + chains_map[cname] = { + "name": cname, + "type": ch_type, + "hook": ch_hook, + "priority": ch_priority, + "policy": ch_policy, + "rules": [], + } else: - if existing.get("type") is None and ch_type is not None: - existing["type"] = ch_type - if existing.get("hook") is None and ch_hook is not None: - existing["hook"] = ch_hook - if existing.get("priority") is None and ch_prio is not None: - existing["priority"] = ch_prio - if existing.get("policy") is None and ch_policy is not None: - existing["policy"] = ch_policy + # merge into placeholder (do not overwrite existing rules) + if isinstance(existing, dict): + if existing.get("type") is None and ch_type is not None: + existing["type"] = ch_type + if existing.get("hook") is None and ch_hook is not None: + existing["hook"] = ch_hook + if existing.get("priority") is None and ch_priority is not None: + existing["priority"] = ch_priority + if existing.get("policy") is None and ch_policy is not None: + existing["policy"] = ch_policy + # rule records elif "rule" in rec: r = rec["rule"] - fam = r.get("family"); tname = r.get("table"); cname = r.get("chain") - if not (fam and tname and cname): - continue - 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: Dict[str, Any] = { - "handle": r.get("handle"), - "expr": r.get("expr"), - "text": rule_text_from_expr(r.get("expr")), - } - if "position" in r: - rule_obj["position"] = r["position"] - if "comment" in r: - rule_obj["comment"] = r["comment"] - chains[cname]["rules"].append(rule_obj) + fam = r.get("family") + table_name = r.get("table") + chain_name = r.get("chain") + handle = r.get("handle") + expr = r.get("expr") + if fam and table_name and chain_name: + tables.setdefault((fam, table_name), {"family": fam, "name": table_name, "chains": {}}) + chains_map = tables[(fam, table_name)]["chains"] + # ensure chain placeholder exists, with possible metadata defaults + chains_map.setdefault(chain_name, {"name": chain_name, "type": None, "hook": None, "priority": None, "policy": None, "rules": []}) - # salvage nested chain metadata if present - if isinstance(r.get("chain"), dict): - csub = r.get("chain") - if chains[cname].get("priority") is None: - p = parse_priority(csub.get("priority") if "priority" in csub else csub.get("prio")) - if p is not None: - chains[cname]["priority"] = p - if chains[cname].get("type") is None and csub.get("type") is not None: - chains[cname]["type"] = csub.get("type") - if chains[cname].get("hook") is None and csub.get("hook") is not None: - chains[cname]["hook"] = csub.get("hook") - if chains[cname].get("policy") is None and csub.get("policy") is not None: - chains[cname]["policy"] = csub.get("policy") + rule_obj: Dict[str, Any] = { + "handle": handle, + "expr": expr, + "text": rule_text_from_expr(expr), + } + # include other useful metadata if present + if "position" in r: + rule_obj["position"] = r["position"] + if "comment" in r: + rule_obj["comment"] = r["comment"] + chains_map[chain_name]["rules"].append(rule_obj) - # materialize lists deterministically + # Attempt to salvage chain metadata from rule record if present + # some nft JSON may include 'chain' subfields inside rule record + # e.g. r.get('chain') might be an object - handle that defensively + if isinstance(r.get("chain"), dict): + csub = r.get("chain") + # try to parse nested priority + if chains_map[chain_name].get("priority") is None: + parsed_prio = parse_priority(csub.get("priority") if "priority" in csub else csub.get("prio")) + if parsed_prio is not None: + chains_map[chain_name]["priority"] = parsed_prio + # type/hook/policy from nested if present + if chains_map[chain_name].get("type") is None and csub.get("type") is not None: + chains_map[chain_name]["type"] = csub.get("type") + if chains_map[chain_name].get("hook") is None and csub.get("hook") is not None: + chains_map[chain_name]["hook"] = csub.get("hook") + if chains_map[chain_name].get("policy") is None and csub.get("policy") is not None: + chains_map[chain_name]["policy"] = csub.get("policy") + + # Convert map to sorted lists for deterministic order, and include chain metadata for (fam, tname) in sorted(tables.keys(), key=lambda k: (k[0], k[1])): tdata = tables[(fam, tname)] chains_list: List[Dict[str, Any]] = [] for cname in sorted(tdata["chains"].keys()): - ch = tdata["chains"][cname] - chains_list.append({ - "name": ch.get("name"), - "type": ch.get("type"), - "hook": ch.get("hook"), - "priority": ch.get("priority"), - "policy": ch.get("policy"), - "rules": ch.get("rules", []), - }) + chdata = tdata["chains"][cname] + chains_list.append( + { + "name": chdata.get("name"), + "type": chdata.get("type"), + "hook": chdata.get("hook"), + "priority": chdata.get("priority"), + "policy": chdata.get("policy"), + "rules": chdata.get("rules", []), + } + ) result["tables"].append({"family": fam, "name": tname, "chains": chains_list}) + return result -def enrich_text_from_chain(custom: Dict[str, Any], mgr: NftManager) -> None: +# ---------- Helpers to render expr -> textual nft (best-effort) ---------- +def expr_to_text(expr: Any) -> Optional[str]: """ - Replace rule['text'] with exact textual lines from `nft list chain` when possible. + Best-effort renderer that converts a typical nft JSON expr (list) into a textual + fragment suitable to append to 'add rule
...'. + Returns None when it cannot deterministically render the provided expr. + Supported cases (common): + - [{'match': {'left': {'payload': {'protocol':'ip','field':'protocol'}}, 'op':'==', 'right':'icmp'}}, {'drop': None}] + -> 'ip protocol icmp drop' + - payload / tcp / udp / counter / accept + - queue tokens and optional bypass support + This intentionally does not attempt to support every nft JSON construct. """ - for table in custom.get("tables", []): - fam = table.get("family"); tname = table.get("name") + if expr is None: + return "" + if isinstance(expr, str): + return expr + if not isinstance(expr, list): + # unsupported top-level type + return None + + parts: List[str] = [] + for element in expr: + if isinstance(element, dict): + # handle drop/accept/counter directly + if "drop" in element: + parts.append("drop") + continue + if "accept" in element: + parts.append("accept") + continue + if "counter" in element: + parts.append("counter") + continue + + # queue support: allow {"queue": 1} or {"queue": {"num":1, "bypass": True}} etc. + if "queue" in element: + q = element["queue"] + token = "queue" + if isinstance(q, dict): + 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: + token += f" num {num}" + if q.get("bypass"): + token += " bypass" + elif isinstance(q, (int, float)): + token += f" num {int(q)}" + elif isinstance(q, str): + token += f" num {q}" + parts.append(token) + continue + + # match left/right payload equals -> ip protocol icmp, or ip saddr/daddr + if "match" in element: + m = element["match"] + left = m.get("left") + right = m.get("right") + # payload matches + if isinstance(left, dict) and "payload" in left and isinstance(right, (str, int)): + p = left["payload"] + prot = p.get("protocol") + field = p.get("field") + # common: protocol field match (protocol == icmp) + if prot and field and isinstance(right, str): + # ip vs ip6 decision is left to the frontend; here we render 'ip protocol icmp' (works for many setups) + if field == "protocol": + parts.append(f"{prot} {field} {right}") + continue + # payload might be l4 ports etc; produce generic payload(...) token + parts.append(f"payload({prot}.{field}) {right}") + continue + # fallback for match: try to stringify right + parts.append("match") + continue + + # payload shorthand + if "payload" in element: + p = element["payload"] + prot = p.get("protocol") + field = p.get("field") + if prot and field: + parts.append(f"payload({prot}.{field})") + continue + parts.append("payload") + continue + + # tcp/udp as nested dicts sometimes appear + if "tcp" in element or "udp" in element: + proto = "tcp" if "tcp" in element else "udp" + val = element.get(proto) + # attempt to detect dport/sport keys + if isinstance(val, dict): + if "dport" in val: + parts.append(f"{proto} dport {val['dport']}") + continue + if "sport" in val: + parts.append(f"{proto} sport {val['sport']}") + continue + parts.append(proto) + continue + + # cmp/binary operators etc — not supported deterministically + # return None to indicate we can't safely render this expr + return None + else: + # non-dict token (string/number) + parts.append(str(element)) + + # join tokens + return " ".join(parts).strip() + + +# ---------- New helper: populate_text_from_chain_text ---------- +def populate_text_from_chain_text(custom: Dict[str, Any]) -> None: + """ + Replace rule['text'] in the 'custom' predictable ruleset with the exact textual + rule lines as produced by `nft list chain
` when possible. + + This modifies `custom` in-place. If textual listing for a chain fails, we fall + back to the existing rule['text'] that was produced from JSON. + """ + tables = custom.get("tables") or [] + for t in tables: + fam = t.get("family") + tname = t.get("name") if not fam or not tname: continue - for ch in table.get("chains", []): + for ch in t.get("chains", []): cname = ch.get("name") if not cname: continue try: - txt = mgr.list_chain_text(fam, tname, cname) or "" - lines = [ln.rstrip() for ln in txt.splitlines() if ln.strip()] + chain_text = mgr.list_chain_text(fam, tname, cname) or "" + lines = [ln.rstrip() for ln in chain_text.splitlines() if ln.strip() != ""] + # build handle -> line map handle_map: Dict[str, str] = {} for ln in lines: - m = _handle_re.search(ln) + m = re.search(r"\bhandle\s+(\d+)\b", ln) if m: handle_map[m.group(1)] = ln.strip() + for rule in ch.get("rules", []): replaced = False h = rule.get("handle") @@ -395,107 +673,210 @@ def enrich_text_from_chain(custom: Dict[str, Any], mgr: NftManager) -> None: if key in handle_map: rule["text"] = handle_map[key] replaced = True + if not replaced: - probe = rule.get("text") or rule_text_from_expr(rule.get("expr")) + # fallback: try to find a line that contains the JSON-derived compact text fragment + expr = rule.get("expr") + probe = rule.get("text") or rule_text_from_expr(expr) if probe: + # try longest-first strategy (not strictly necessary here) — simple substring match for ln in lines: if probe in ln: rule["text"] = ln.strip() replaced = True break + # if still not replaced, keep existing rule["text"] except Exception as e: - logger.debug("enrich_text_from_chain failed for %s %s %s: %s", fam, tname, cname, e) + logger.debug( + "populate_text_from_chain_text: failed to get textual chain for %s %s %s: %s", + fam, + tname, + cname, + e, + ) continue -# ---------- FastAPI app ---------- -app = FastAPI(title="nft API (robust unwrap)") -router = APIRouter(prefix="/firewall", tags=["firewall"]) -mgr = NftManager() +# ---------- Routes ---------- - -@router.get("/rules", response_model=None) -def get_rules(): +@router.get("/rules", response_model=RulesetOut, summary="List ruleset") +def list_rules(): """ - Return: JSONResponse({"ruleset": }) - If nft -j is available -> returns native dict under "ruleset". - Else -> returns textual ruleset string. + Returns the ruleset in a stable, strongly-typed JSON shape derived from `nft -j list ruleset`. + + Structure: + { "ruleset": { "tables": [ { "family": ..., "name": ..., "chains": [ { "name": ..., "type": ..., "hook": ..., "priority": ..., "policy": ..., "rules": [ { "handle", "expr", "text" } ] } ] } ] } } + + Fallback: + - If nft JSON is unavailable, falls back to returning the raw textual ruleset string. """ try: try: nft_json = mgr.list_rules_json() except NftError as e: - logger.debug("json listing unavailable: %s", e) - text = mgr.list_rules_text() - return JSONResponse(content={"ruleset": text.strip() if text is not None else None}, status_code=200) + logger.debug("could not obtain nft JSON ruleset: %s", e) + text = mgr.list_rules() + return RulesetOut(ruleset=text.strip() if text is not None else None) custom = build_predictable_ruleset(nft_json) - 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, no double-encoding) - return JSONResponse(content={"ruleset": custom}, status_code=200) + # Enrich rule['text'] by attempting to fetch the exact textual nft rule lines + # as printed by `nft list chain
`. This is best-effort and + # will not fail the overall listing if textual retrieval fails for some chains. + try: + populate_text_from_chain_text(custom) + except Exception as e: + logger.debug("list_rules: populate_text_from_chain_text failed: %s", e) + + ruleset_model = RulesetModel.parse_obj(custom) + return RulesetOut(ruleset=ruleset_model) except NftError as e: - logger.exception("get_rules: nft error") + logger.exception("list_rules failed") raise HTTPException(status_code=500, detail=str(e)) except Exception as e: - logger.exception("get_rules internal") + logger.exception("list_rules internal error") raise HTTPException(status_code=500, detail=str(e)) -@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) - 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", response_model=None, status_code=201) -def create_rule(req: Dict[str, Any]): +@router.post( + "/rules", + response_model=ExecResult, + status_code=status.HTTP_201_CREATED, + summary="Create rule (JSON, expr required; returns ExecResult with rc/stdout/stderr)", +) +def create_rule_json(req: CreateRuleRequest): """ - 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. + Create a rule from JSON (expr required). + - If req.position is provided, uses: insert rule
position + - Otherwise, uses: add rule
(append) + - If rendering fails: 400 instructing the client to use POST /firewall/raw + - Returns ExecResult on success (201) or on error (400) with stdout/stderr in body. + - If nft wrapper returns an invalid rc but the command produced no stderr, we double-check the chain + to see if the new rule is present; if present we treat as success. """ - 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") - rendered = rule_text_from_expr(expr) - if rendered is None: - raise HTTPException(status_code=400, detail="cannot render expr to textual rule; use /raw") - expr_text = rendered.strip() - if pos is not None: + try: + family = req.family + table = req.table + chain = req.chain + + if req.expr is None: + raise NftError("field 'expr' is required for JSON rule creation") + + rendered = expr_to_text(req.expr) + if rendered is None: + raise NftError( + "cannot render provided 'expr' to textual nft syntax. " + "Please use POST /firewall/raw to execute the textual nft command." + ) + + expr_text = rendered.strip() + + # If a position is explicitly requested, use the 'insert rule ... position ...' form. + # 'add rule ... position ...' is not supported by some nft versions / syntaxes. + if req.position is not None: + try: + pos = int(req.position) + # clamp pos to >= 0 + if pos < 0: + pos = 0 + except Exception: + pos = 0 + cmd = f"insert rule {family} {table} {chain} position {pos} {expr_text}" + else: + cmd = f"add rule {family} {table} {chain} {expr_text}" + + logger.info("create_rule_json executing command: %s", cmd) + + res = mgr.cmd(cmd) + # res expected {"rc": rc, "stdout": out, "stderr": err} + raw_rc = res.get("rc") + stdout = res.get("stdout") or "" + stderr = res.get("stderr") or "" + + # Coerce rc to int safely; if not int-like, set -1 to indicate unknown. try: - pos_i = int(pos) - if pos_i < 0: - pos_i = 0 + rc = int(raw_rc) except Exception: - pos_i = 0 - cmd = f"insert rule {family} {table} {chain} position {pos_i} {expr_text}" - else: - cmd = f"add rule {family} {table} {chain} {expr_text}" - rc, out, err = mgr.cmd(cmd) - # 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}") + rc = -1 + + logger.info("nft cmd rc=%s stdout=%r stderr=%r cmd=%s", rc, stdout, stderr, cmd) + + exec_res = ExecResult(rc=rc, stdout=stdout or None, stderr=stderr or None) + + # If rc == 0 — success + if rc == 0: + return exec_res + + # Handle the annoying case: wrapper returned invalid rc (<0) or non-zero, + # but stderr is empty. The command may have succeeded nevertheless. + if (rc < 0 or rc != 0) and stderr.strip() == "": + logger.debug("create_rule_json: rc indicates failure but stderr empty; verifying rule presence") + + # Attempt to verify the rule exists by listing the chain and searching for a textual match. + # We use list_chain_text because it returns textual rule lines we can search for the preview text. + try: + chain_text = mgr.list_chain_text(family, table, chain) or "" + # Simple presence check: the textual fragment we attempted to add should be present + # as a substring in the chain listing (e.g. "ip protocol icmp drop"). + if expr_text and expr_text in chain_text: + logger.info("create_rule_json: detected rule in chain after add; treating as success") + # return success ExecResult with rc=0 to indicate success to client + return ExecResult(rc=0, stdout=stdout or None, stderr=stderr or None) + else: + logger.debug("create_rule_json: rule not found in chain text; chain_text=%r", chain_text) + except Exception as e_chain: + logger.warning("create_rule_json: failed to list chain for verification: %s", e_chain) + + # If we reach here -> treat as error: return 400 with exec_res in body. + # FastAPI cannot both raise HTTPException and include ExecResult as body easily, so raise HTTPException + # with detail that includes stderr and the executed cmd. + detail = f"nft command failed rc={rc}. stderr: {stderr!r}. cmd: {cmd}" + logger.warning("create_rule_json failed: %s", detail) + raise HTTPException(status_code=400, detail=detail) + + except NftError as e: + logger.warning("create_rule_json NftError: %s", e) + raise HTTPException(status_code=400, detail=str(e)) + except HTTPException: + # re-raise HTTPException so we don't wrap it again + raise + except Exception as e: + logger.exception("create_rule_json internal error") + raise HTTPException(status_code=500, detail=str(e)) -@router.delete("/rules/{handle}", status_code=status.HTTP_204_NO_CONTENT) + +@router.delete("/rules/{handle}", status_code=status.HTTP_204_NO_CONTENT, summary="Delete rule by handle") def delete_rule(handle: int, family: str = "inet", table: str = "filter", chain: str = "input"): - if not isinstance(handle, int) or handle <= 0: - raise HTTPException(status_code=400, detail="invalid handle") - cmd = f"delete rule {family} {table} {chain} handle {handle}" - rc, out, err = mgr.cmd(cmd) - if rc != 0: - raise HTTPException(status_code=500, detail=f"delete failed rc={rc} stderr={err}") - return JSONResponse(status_code=204, content={}) + """ + Delete a rule by handle using textual nft command. + Command executed: + delete rule
handle + """ + try: + mgr.delete_rule_by_handle_text(family=family, table=table, chain=chain, handle=handle) + except ValueError as e: + logger.warning("delete_rule client error: %s", e) + raise HTTPException(status_code=400, detail=str(e)) + except NftError as e: + logger.exception("delete_rule failed") + raise HTTPException(status_code=500, detail=str(e)) + except Exception as e: + logger.exception("delete_rule internal error") + raise HTTPException(status_code=500, detail=str(e)) +@router.post("/raw", response_model=ExecResult, summary="Execute raw textual nft command") +def exec_raw(req: RawCmdRequest): + """ + Execute an arbitrary textual nft command and return structured {rc, stdout, stderr}. + """ + try: + res = mgr.cmd(req.cmd) + rc = int(res.get("rc", -1) or -1) + return ExecResult(rc=rc, stdout=res.get("stdout"), stderr=res.get("stderr")) + except Exception as e: + logger.exception("exec_raw failed") + raise HTTPException(status_code=500, detail=str(e)) + app.include_router(router) \ No newline at end of file