# app.py from typing import Any, Dict, List, Optional, Tuple, Union from fastapi import FastAPI, APIRouter, HTTPException, status from pydantic import BaseModel, Field, ValidationError import logging import json import re # 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_raw_only") # ---------- Exceptions ---------- class NftError(RuntimeError): pass # ---------- NftManager (textual-only) ---------- class NftManager: def __init__(self) -> None: self.nft = Nftables() try: # prefer JSON output globally where available self.nft.set_json_output(True) self.nft.set_handle_output(True) except Exception: logger.debug("set_json_output/set_handle_output not available") def cmd(self, text_cmd: str) -> Dict[str, Optional[str]]: rc, out, err = self.nft.cmd(text_cmd) if rc != 0: logger.warning("nft cmd rc=%s stderr=%s cmd=%s", rc, err, text_cmd) return {"rc": rc, "stdout": out, "stderr": err} def list_rules_json(self) -> Dict[str, Any]: res = self.cmd("list ruleset") if res["rc"] != 0: raise NftError(f"nft list ruleset failed: {res['stderr']}") out = res["stdout"] if not out: raise NftError("nft list ruleset returned empty output") try: return json.loads(out) except json.JSONDecodeError as e: raise NftError(f"json decode error: {e}") def list_rules_text(self) -> str: # best-effort: temporarily disable JSON output so we get textual form json_toggled = False try: if hasattr(self.nft, "set_json_output"): try: self.nft.set_json_output(False) json_toggled = True except Exception: logger.debug("could not toggle set_json_output(False)") res = self.cmd("list ruleset") finally: if json_toggled and hasattr(self.nft, "set_json_output"): try: self.nft.set_json_output(True) except Exception: logger.debug("could not restore set_json_output(True)") if res["rc"] != 0: raise NftError(f"nft list ruleset failed: {res['stderr']}") return res["stdout"] or "" def delete_rule_by_handle_text(self, family: str, table: str, chain: str, handle: int) -> None: if not isinstance(handle, int) or handle <= 0: raise ValueError("handle must be a positive integer") 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() # ---------- Models ---------- class RawCmdRequest(BaseModel): cmd: str = Field(..., description="Textual nft command to execute") class ExecResult(BaseModel): rc: int = Field(..., description="Return code from nft execution") stdout: Optional[str] = Field(None) stderr: Optional[str] = Field(None) class RuleOut(BaseModel): handle: Optional[int] = Field(None) expr: Any = Field(..., description="Machine-readable nft expression (original nft JSON expr).") text: str = Field(..., description="Deterministic short display string derived from expr") position: Optional[Any] = Field(None) comment: Optional[str] = Field(None) class ChainOut(BaseModel): name: str = Field(...) type: Optional[str] = Field(None) hook: Optional[str] = Field(None) priority: Optional[int] = Field(None) policy: Optional[str] = Field(None) rules: List[RuleOut] = Field(...) class TableOut(BaseModel): family: str = Field(...) name: str = Field(...) chains: List[ChainOut] = Field(...) class RulesetModel(BaseModel): tables: List[TableOut] = Field(...) class CreateRuleRequest(BaseModel): family: str table: str chain: str expr: Any position: Optional[int] comment: Optional[str] RulesetValue = Optional[Union[RulesetModel, str]] class RulesetOut(BaseModel): ruleset: RulesetValue # ---------- Helpers ---------- def parse_priority(val: Any) -> Optional[int]: if val is None: return None if isinstance(val, int): return val if isinstance(val, str): s = val.strip() try: return int(s) except Exception: try: return int(float(s)) except Exception: return None 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 build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]: """ Convert nft -j list ruleset parsed JSON into a deterministic, predictable JSON: """ 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]] = {} for rec in items: if "table" in rec: t = rec["table"] fam = t.get("family") name = t.get("name") if fam and name: tables.setdefault((fam, name), {"family": fam, "name": name, "chains": {}}) elif "chain" in rec: ch = rec["chain"] 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 = chains_map.get(cname) ch_type = ch.get("type") ch_hook = ch.get("hook") ch_priority = parse_priority(ch.get("priority") if "priority" in ch else ch.get("prio") if "prio" in ch else ch.get("priority", None)) 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_map[cname] = { "name": cname, "type": ch_type, "hook": ch_hook, "priority": ch_priority, "policy": ch_policy, "rules": [], } else: 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 elif "rule" in rec: r = rec["rule"] 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"] chains_map.setdefault(chain_name, {"name": chain_name, "type": None, "hook": None, "priority": None, "policy": None, "rules": []}) # do NOT change expr shape here; keep it exactly as NFT JSON provided rule_obj: Dict[str, Any] = {"handle": handle, "expr": expr, "text": ""} 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) # salvage chain-level metadata from rule record if present if isinstance(r.get("chain"), dict): csub = r.get("chain") 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 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 to lists (deterministic order) 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()): 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 # ---------- Text parsing helpers (enrichment only) ---------- def parse_ruleset_text(nft_text: str) -> Dict[Tuple[str, str, str], List[Dict[str, Optional[Union[str, int]]]]]: result: Dict[Tuple[str, str, str], List[Dict[str, Optional[Union[str, int]]]]] = {} if not nft_text: return result table_re = re.compile(r"^\s*table\s+(\S+)\s+(\S+)\s*\{") chain_re = re.compile(r"^\s*chain\s+(\S+)\s*\{") handle_re = re.compile(r"#\s*handle\s*(\d+)\b") # only skip semicolon-terminated chain metadata lines (type/hook/priority/policy) chain_meta_re = re.compile(r"^\s*(type\b|hook\b|priority\b|policy\b)\b.*;") current_family = None current_table = None current_chain = None for raw_ln in nft_text.splitlines(): ln = raw_ln.rstrip("\n") s = ln.strip() m_table = table_re.match(ln) if m_table: current_family = m_table.group(1) current_table = m_table.group(2) current_chain = None continue m_chain = chain_re.match(ln) if m_chain and current_family and current_table: current_chain = m_chain.group(1) key = (current_family, current_table, current_chain) result.setdefault(key, []) continue if current_family and current_table and current_chain: if s == "" or s == "{" or s == "}": continue if chain_meta_re.match(s): # skip chain metadata lines only continue m_handle = handle_re.search(s) handle_val: Optional[int] = None if m_handle: try: handle_val = int(m_handle.group(1)) except Exception: handle_val = None key = (current_family, current_table, current_chain) result.setdefault(key, []).append({"line": ln.strip(), "handle": handle_val}) return result def populate_text_from_ruleset_text(custom: Dict[str, Any], nft_text: str) -> None: """ Enrich JSON-derived 'custom' structure in-place by setting only rule['text'] when a reliable textual mapping is found. Do not change expr or other types. """ if not nft_text: return parsed = parse_ruleset_text(nft_text) for table in custom.get("tables", []): fam = table.get("family") tname = table.get("name") if not fam or not tname: continue for chain in table.get("chains", []): cname = chain.get("name") if not cname: continue key = (fam, tname, cname) textual_entries = parsed.get(key, []) if not textual_entries: continue handle_map: Dict[int, str] = {} ordered_lines: List[str] = [] for ent in textual_entries: ln = ent.get("line") or "" h = ent.get("handle") ordered_lines.append(ln) if isinstance(h, int): handle_map[h] = ln rules = chain.get("rules", []) for idx, rule in enumerate(rules): # ONLY update 'text' when we can map a textual line h = rule.get("handle") mapped: Optional[str] = None if isinstance(h, int) and h in handle_map: mapped = handle_map[h] else: pos = rule.get("position") if isinstance(pos, int) and 0 <= pos < len(ordered_lines): mapped = ordered_lines[pos] elif idx < len(ordered_lines): mapped = ordered_lines[idx] # final substring probe (safe) if mapped is None: probe = rule.get("text") if probe: for ln in ordered_lines: if probe in ln: mapped = ln break if mapped is not None: # ensure we only write a str into 'text' try: rule["text"] = str(mapped) except Exception: rule["text"] = mapped # should be str already # ---------- Normalization helper (lightweight and safe) ---------- def normalize_custom_for_model(custom: Dict[str, Any]) -> None: """ Make minimal, safe guarantees required by Pydantic: - rule['expr'] must exist (if None -> set to empty list) - rule['text'] must be a str (if missing -> derived string) - rule['handle'] coerced to int or None Do NOT change any other shapes. """ for t in custom.get("tables", []): for ch in t.get("chains", []): rules = ch.get("rules", []) or [] for r in rules: # expr: if missing or None => set to [] (preserves Any) if "expr" not in r or r.get("expr") is None: r["expr"] = [] # text: ensure string if "text" not in r or r.get("text") is None: r["text"] = "" else: if not isinstance(r["text"], str): try: r["text"] = str(r["text"]) except Exception: r["text"] = "" # handle: coerce to int or None h = r.get("handle") if isinstance(h, str): try: r["handle"] = int(h) except Exception: r["handle"] = None elif isinstance(h, float): try: r["handle"] = int(h) except Exception: r["handle"] = None elif not isinstance(h, int): r["handle"] = None # ---------- Routes ---------- @router.get("/rules", response_model=RulesetOut, summary="List ruleset") def list_rules(): """ Returns JSON-derived ruleset (RulesetModel) and enriches each rule['text'] with the textual nft rule line when possible. This function will not replace JSON-derived 'expr' or other data types — enrichment is additive only. """ try: try: nft_json = mgr.list_rules_json() except NftError as e: logger.debug("could not obtain nft JSON ruleset: %s", e) raise HTTPException(status_code=500, detail=f"Failed to obtain nft JSON ruleset: {e}") # best-effort textual snapshot for enrichment nft_text = "" try: nft_text = mgr.list_rules_text() except Exception: logger.debug("could not obtain textual nft ruleset snapshot") # Build canonical JSON-derived shape (source of truth) custom = build_predictable_ruleset(nft_json) # Enrich only the 'text' field in-place using the textual snapshot try: if nft_text: populate_text_from_ruleset_text(custom, nft_text) except Exception as e: logger.debug("populate_text_from_ruleset_text failed: %s", e) # Normalize minimally for model validation normalize_custom_for_model(custom) # Debug counts num_tables = len(custom.get("tables", [])) num_rules = sum(len(ch.get("rules", [])) for t in custom.get("tables", []) for ch in t.get("chains", [])) logger.info("list_rules: returning tables=%d rules=%d", num_tables, num_rules) # RETURN a shape matching response_model=RulesetOut try: ruleset_model = RulesetModel.parse_obj(custom) except ValidationError as ve: # log full validation error for debugging and return 500 with message logger.exception("RulesetModel validation failed: %s", ve) raise HTTPException(status_code=500, detail=f"Internal: ruleset validation failed: {ve}") return {"ruleset": ruleset_model} except NftError as e: logger.exception("list_rules failed") raise HTTPException(status_code=500, detail=str(e)) except HTTPException: raise except Exception as e: logger.exception("list_rules internal error") raise HTTPException(status_code=500, detail=str(e)) @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"): 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): 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)