diff --git a/backend/src/api/nft_manager.py b/backend/src/api/nft_manager.py index 787d539..a76736d 100644 --- a/backend/src/api/nft_manager.py +++ b/backend/src/api/nft_manager.py @@ -1,6 +1,7 @@ # 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 @@ -29,7 +30,7 @@ class NftManager: def __init__(self) -> None: self.nft = Nftables() - # set_json_output is optional; don't rely on it for JSON path. + # best-effort: don't strictly rely on set_json_output for JSON path try: if hasattr(self.nft, "set_json_output"): self.nft.set_json_output(True) @@ -41,7 +42,7 @@ class NftManager: 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} + return {"rc": int(rc), "stdout": out or "", "stderr": err or ""} def json_cmd(self, text_cmd: str) -> Tuple[int, str, str]: """ @@ -53,24 +54,21 @@ class NftManager: try: res = self.nft.json_cmd(text_cmd) if isinstance(res, (list, tuple)) and len(res) >= 3: - return int(res[0]), res[1], res[2] + return int(res[0]), res[1] or "", res[2] or "" except Exception as e: logger.debug("nft.json_cmd failed, falling back to cmd -j: %s", e) # fallback: append -j if not present and call textual cmd() cmd_with_j = text_cmd if "-j" in text_cmd else f"{text_cmd} -j" r = self.cmd(cmd_with_j) - rc = int(r.get("rc", -1) or -1) - out = r.get("stdout") or "" - err = r.get("stderr") or "" - return rc, out, err + return int(r.get("rc", -1) or -1), r.get("stdout", "") or "", r.get("stderr", "") or "" def list_rules(self) -> str: """Return textual ruleset from `nft list ruleset`.""" res = self.cmd("list ruleset") if res["rc"] != 0: raise NftError(f"nft list ruleset failed: {res['stderr']}") - return res["stdout"] or "" + return res["stdout"] def list_rules_json(self) -> Dict[str, Any]: """Return parsed JSON from `nft -j list ruleset`.""" @@ -100,7 +98,8 @@ class NftManager: logger.debug("list_chain_text: textual cmd failed (%s), attempting JSON fallback", res["stderr"]) rc, out, err = self.json_cmd(cmd) if rc != 0: - raise NftError(f"nft {cmd} failed: {err}") + # raise original textual error if JSON fallback doesn't work either + raise NftError(f"nft {cmd} failed: {res['stderr'] or err}") try: parsed = json.loads(out) @@ -166,7 +165,6 @@ class NftManager: continue parts.append("queue") continue - # fallback: join keys parts.append("+".join(sorted(part.keys()))) else: parts.append(str(part)) @@ -193,7 +191,7 @@ router = APIRouter(prefix="/firewall", tags=["firewall"]) mgr = NftManager() -# ---------- Request/Response models ---------- +# ---------- Request/Response models (only used for validation / docs) ---------- class RawCmdRequest(BaseModel): cmd: str = Field(..., example="add rule inet filter input ip saddr 10.0.0.0/8 drop") @@ -204,33 +202,6 @@ class ExecResult(BaseModel): stderr: Optional[str] = None -class RuleOut(BaseModel): - handle: Optional[int] = None - expr: Any - text: str - position: Optional[Any] = None - comment: Optional[str] = None - - -class ChainOut(BaseModel): - name: str - type: Optional[str] = None - hook: Optional[str] = None - priority: Optional[int] = None - policy: Optional[str] = None - rules: List[RuleOut] - - -class TableOut(BaseModel): - family: str - name: str - chains: List[ChainOut] - - -class RulesetModel(BaseModel): - tables: List[TableOut] - - class CreateRuleRequest(BaseModel): family: str table: str @@ -240,15 +211,7 @@ class CreateRuleRequest(BaseModel): comment: Optional[str] = None -# ruleset may be typed RulesetModel or raw textual string (fallback) -RulesetValue = Optional[Union[Dict[str, Any], str]] - - -class RulesetOut(BaseModel): - ruleset: RulesetValue = None - - -# ---------- Helpers ---------- +# ---------- Helpers to make predictable output ---------- _handle_re = re.compile(r"\s+#\s*handle\s+\d+\s*$") @@ -480,26 +443,32 @@ def populate_text_from_chain_text(custom: Dict[str, Any]) -> None: # ---------- Routes ---------- -@router.get("/rules", response_model=RulesetOut, summary="List ruleset") +@router.get("/rules", summary="List ruleset") def list_rules(): + """ + Returns JSON: { "ruleset": } + Ensures `ruleset` is a native dict when JSON is available and parsed. + """ try: try: nft_json = mgr.list_rules_json() except NftError as e: logger.debug("could not obtain nft JSON ruleset: %s", e) + # textual fallback - return plain textual ruleset string text = mgr.list_rules() - return RulesetOut(ruleset=text.strip() if text is not None else None) + return JSONResponse(content={"ruleset": text.strip() if text is not None else None}, status_code=200) + # build structured representation (native Python) custom = build_predictable_ruleset(nft_json) + # Try to replace rule['text'] with exact textual lines from nft list chain ... try: populate_text_from_chain_text(custom) except Exception as e: - logger.debug("list_rules: populate_text_from_chain_text failed: %s", e) + logger.debug("populate_text_from_chain_text failed: %s", e) - # IMPORTANT: return native Python structure (dict) — not a JSON string. - ruleset_model = RulesetModel.parse_obj(custom) - return RulesetOut(ruleset=ruleset_model.dict()) + # Return native structure (do NOT json.dumps) + return JSONResponse(content={"ruleset": custom}, status_code=200) except NftError as e: logger.exception("list_rules failed") raise HTTPException(status_code=500, detail=str(e)) @@ -541,19 +510,14 @@ def create_rule_json(req: CreateRuleRequest): logger.info("create_rule_json executing: %s", cmd) res = mgr.cmd(cmd) - raw_rc = res.get("rc") + rc = int(res.get("rc", -1) or -1) stdout = res.get("stdout") or "" stderr = res.get("stderr") or "" - try: - rc = int(raw_rc) - except Exception: - rc = -1 - - exec_res = ExecResult(rc=rc, stdout=stdout or None, stderr=stderr or None) if rc == 0: - return exec_res + return ExecResult(rc=rc, stdout=stdout or None, stderr=stderr or None) + # If rc indicates failure but stderr empty, check chain presence if (rc < 0 or rc != 0) and stderr.strip() == "": try: chain_text = mgr.list_chain_text(family, table, chain) or "" @@ -596,7 +560,7 @@ 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")) + return ExecResult(rc=rc, stdout=res.get("stdout") or None, stderr=res.get("stderr") or None) except Exception as e: logger.exception("exec_raw failed") raise HTTPException(status_code=500, detail=str(e))