This commit is contained in:
@@ -18,29 +18,26 @@ class NftError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
# ---------- NftManager (textual-only) ----------
|
||||
# ---------- NftManager (textual vs JSON) ----------
|
||||
class NftManager:
|
||||
"""
|
||||
Thin wrapper around python-nftables exposing:
|
||||
- cmd execution (textual nft commands via Nftables.cmd())
|
||||
- json execution via Nftables.json_cmd() when available
|
||||
- convenience list_rules_text / list_rules_json / list_chain_text
|
||||
- cmd execution via Nftables.cmd() for textual output
|
||||
- json_cmd execution when available for JSON output (or fallback to cmd with -j)
|
||||
- helpers to list rules / chain text
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.nft = Nftables()
|
||||
# Try to prefer JSON for general listing; we'll toggle off for chain-list calls.
|
||||
# set_json_output is optional; don't rely on it for JSON path.
|
||||
try:
|
||||
# If set_json_output exists it's a convenience; we won't rely on it for JSON path.
|
||||
if hasattr(self.nft, "set_json_output"):
|
||||
self.nft.set_json_output(True)
|
||||
except Exception:
|
||||
logger.debug("set_json_output not available or ignored")
|
||||
|
||||
def cmd(self, text_cmd: str) -> Dict[str, Optional[str]]:
|
||||
"""
|
||||
Execute a textual nft command via Nftables.cmd().
|
||||
Returns dict { "rc": rc, "stdout": out_str, "stderr": err_str }.
|
||||
"""
|
||||
"""Execute textual nft command via Nftables.cmd()."""
|
||||
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)
|
||||
@@ -48,23 +45,20 @@ class NftManager:
|
||||
|
||||
def json_cmd(self, text_cmd: str) -> Tuple[int, str, str]:
|
||||
"""
|
||||
Execute a JSON-output nft command. Prefer Nftables.json_cmd if available
|
||||
(returns tuple (rc, stdout, stderr)). Otherwise, call `nft <cmd> -j` using cmd()
|
||||
and attempt to return the same tuple shape.
|
||||
Execute nft command expecting JSON output.
|
||||
Prefer Nftables.json_cmd when available (returns (rc, out, err)).
|
||||
Otherwise call cmd() with a '-j' suffix and return a similar tuple.
|
||||
"""
|
||||
# Prefer built-in json_cmd if present
|
||||
if hasattr(self.nft, "json_cmd"):
|
||||
try:
|
||||
res = self.nft.json_cmd(text_cmd)
|
||||
# Expecting (rc, out, err)
|
||||
if isinstance(res, (list, tuple)) and len(res) >= 3:
|
||||
return int(res[0]), res[1], res[2]
|
||||
except Exception as e:
|
||||
logger.debug("json_cmd failed, falling back to cmd with -j: %s", e)
|
||||
logger.debug("nft.json_cmd failed, falling back to cmd -j: %s", e)
|
||||
|
||||
# Fallback: call cmd with -j variant and parse output
|
||||
# Use 'list ruleset -j' or similar command suffixes as caller provides the whole command
|
||||
cmd_with_j = f"{text_cmd} -j" if "-j" not in text_cmd else text_cmd
|
||||
# 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 ""
|
||||
@@ -72,19 +66,14 @@ class NftManager:
|
||||
return rc, out, err
|
||||
|
||||
def list_rules(self) -> str:
|
||||
"""
|
||||
Return the textual ruleset as produced by 'nft list ruleset'.
|
||||
"""
|
||||
"""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 ""
|
||||
|
||||
def list_rules_json(self) -> Dict[str, Any]:
|
||||
"""
|
||||
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.
|
||||
"""
|
||||
"""Return parsed JSON from `nft -j list ruleset`."""
|
||||
rc, out, err = self.json_cmd("list ruleset")
|
||||
if rc != 0:
|
||||
raise NftError(f"nft list ruleset failed: {err}")
|
||||
@@ -99,41 +88,37 @@ class NftManager:
|
||||
def list_chain_text(self, family: str, table: str, chain: str) -> str:
|
||||
"""
|
||||
Return textual output of `nft list chain <family> <table> <chain>`.
|
||||
We prefer calling nft in textual mode (nft.cmd). If the wrapper can't return textual
|
||||
output and returns JSON, we attempt a best-effort reconstruction of the textual lines.
|
||||
Prefer textual cmd(); if it fails and JSON is returned, attempt a best-effort
|
||||
reconstruction of textual lines from JSON.
|
||||
"""
|
||||
cmd = f"list chain {family} {table} {chain}"
|
||||
# Attempt to call textual cmd (this uses self.nft.cmd)
|
||||
res = self.cmd(cmd)
|
||||
if res["rc"] != 0:
|
||||
# if the raw textual command failed, try JSON and attempt to reconstruct textual
|
||||
logger.debug("list_chain_text textual cmd rc!=0; trying JSON fallback: %s", res["stderr"])
|
||||
# try json_cmd
|
||||
if res["rc"] == 0:
|
||||
return res["stdout"] or ""
|
||||
|
||||
# textual call failed -> try JSON fallback and reconstruct
|
||||
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}")
|
||||
# attempt to build textual lines from JSON
|
||||
|
||||
try:
|
||||
parsed = json.loads(out)
|
||||
except Exception:
|
||||
# give up, return raw text (could be empty)
|
||||
# give up and return textual stdout (maybe empty)
|
||||
return res.get("stdout") or ""
|
||||
# build textual representation from JSON (best-effort)
|
||||
|
||||
lines: List[str] = []
|
||||
records = parsed.get("nftables") if isinstance(parsed, dict) else (parsed or [])
|
||||
if not isinstance(records, list):
|
||||
records = []
|
||||
for rec in records:
|
||||
if "rule" in rec:
|
||||
# try to use 'line' if present
|
||||
r = rec["rule"]
|
||||
# sometimes json includes 'expr' or 'handle' — we try to recompose a line
|
||||
line_parts: List[str] = []
|
||||
handle = r.get("handle")
|
||||
# attempt to render expr to short textual form
|
||||
expr = r.get("expr")
|
||||
parts: List[str] = []
|
||||
if isinstance(expr, list):
|
||||
# best-effort: reuse simple tokens
|
||||
for part in expr:
|
||||
if isinstance(part, dict):
|
||||
if "payload" in part:
|
||||
@@ -141,7 +126,7 @@ class NftManager:
|
||||
prot = p.get("protocol")
|
||||
field = p.get("field")
|
||||
if prot and field:
|
||||
line_parts.append(f"payload({prot}.{field})")
|
||||
parts.append(f"payload({prot}.{field})")
|
||||
continue
|
||||
if "match" in part:
|
||||
m = part["match"]
|
||||
@@ -152,55 +137,48 @@ class NftManager:
|
||||
prot = p.get("protocol")
|
||||
field = p.get("field")
|
||||
if prot and field:
|
||||
line_parts.append(f"{prot} {field} {right}")
|
||||
parts.append(f"{prot} {field} {right}")
|
||||
continue
|
||||
line_parts.append("match")
|
||||
parts.append("match")
|
||||
continue
|
||||
if "drop" in part:
|
||||
line_parts.append("drop")
|
||||
parts.append("drop")
|
||||
continue
|
||||
if "accept" in part:
|
||||
line_parts.append("accept")
|
||||
parts.append("accept")
|
||||
continue
|
||||
if "counter" in part:
|
||||
line_parts.append("counter")
|
||||
parts.append("counter")
|
||||
continue
|
||||
if "queue" in part:
|
||||
q = part["queue"]
|
||||
tok = "queue"
|
||||
if isinstance(q, dict):
|
||||
num = q.get("num") or q.get("number") or q.get("queue_number")
|
||||
token = "queue"
|
||||
if num is not None:
|
||||
token += f" num {num}"
|
||||
if q.get("bypass"):
|
||||
token += " bypass"
|
||||
line_parts.append(token)
|
||||
elif isinstance(q, (int, float)):
|
||||
line_parts.append(f"queue num {int(q)}")
|
||||
else:
|
||||
line_parts.append("queue")
|
||||
tok += f" num {num}"
|
||||
if q.get("bypass") or q.get("flags") == "bypass":
|
||||
tok += " bypass"
|
||||
parts.append(tok)
|
||||
continue
|
||||
# fallback to keys
|
||||
line_parts.append("+".join(sorted(part.keys())))
|
||||
if isinstance(q, (int, float)):
|
||||
parts.append(f"queue num {int(q)}")
|
||||
continue
|
||||
parts.append("queue")
|
||||
continue
|
||||
# fallback: join keys
|
||||
parts.append("+".join(sorted(part.keys())))
|
||||
else:
|
||||
line_parts.append(str(part))
|
||||
parts.append(str(part))
|
||||
else:
|
||||
# expr not list -> fallback to JSON dump of rule
|
||||
line_parts.append(json.dumps(r))
|
||||
# combine into single textual line; append handle if present
|
||||
text_line = " ".join(line_parts).strip()
|
||||
parts.append(json.dumps(r))
|
||||
text_line = " ".join(parts).strip()
|
||||
if handle is not None:
|
||||
text_line = f"{text_line} # handle {handle}"
|
||||
lines.append(text_line)
|
||||
return "\n".join(lines) if lines else (res.get("stdout") or "")
|
||||
# textual cmd succeeded
|
||||
return res.get("stdout") or ""
|
||||
|
||||
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 <family> <table> <chain> handle <handle>
|
||||
"""
|
||||
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}"
|
||||
@@ -215,81 +193,62 @@ router = APIRouter(prefix="/firewall", tags=["firewall"])
|
||||
mgr = NftManager()
|
||||
|
||||
|
||||
# ---------- Request/Response models (strongly typed) ----------
|
||||
# ---------- Request/Response models ----------
|
||||
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")
|
||||
cmd: str = Field(..., 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": ""}}
|
||||
rc: int
|
||||
stdout: Optional[str] = None
|
||||
stderr: Optional[str] = None
|
||||
|
||||
|
||||
# --- 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")
|
||||
handle: Optional[int] = None
|
||||
expr: Any
|
||||
text: str
|
||||
position: Optional[Any] = None
|
||||
comment: Optional[str] = None
|
||||
|
||||
|
||||
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)")
|
||||
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 = Field(..., description="Table family (inet/bridge/ipv4/...)")
|
||||
name: str = Field(..., description="Table name", example="filter")
|
||||
chains: List[ChainOut] = Field(..., description="Chains in this table")
|
||||
family: str
|
||||
name: str
|
||||
chains: List[ChainOut]
|
||||
|
||||
|
||||
class RulesetModel(BaseModel):
|
||||
tables: List[TableOut] = Field(..., description="Top-level tables list")
|
||||
tables: List[TableOut]
|
||||
|
||||
|
||||
# ---------- 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,
|
||||
}
|
||||
}
|
||||
family: str
|
||||
table: str
|
||||
chain: str
|
||||
expr: Any
|
||||
position: Optional[int] = None
|
||||
comment: Optional[str] = None
|
||||
|
||||
|
||||
# ruleset may be typed RulesetModel or raw textual string (fallback)
|
||||
RulesetValue = Optional[Union[RulesetModel, str]]
|
||||
RulesetValue = Optional[Union[Dict[str, Any], str]]
|
||||
|
||||
|
||||
class RulesetOut(BaseModel):
|
||||
ruleset: RulesetValue = Field(
|
||||
None,
|
||||
description="Parsed, strongly-typed ruleset (RulesetModel) or raw textual ruleset string if JSON is unavailable.",
|
||||
)
|
||||
ruleset: RulesetValue = None
|
||||
|
||||
|
||||
# ---------- Helpers to convert to desired shape ----------
|
||||
# ---------- Helpers ----------
|
||||
_handle_re = re.compile(r"\s+#\s*handle\s+\d+\s*$")
|
||||
|
||||
|
||||
@@ -319,9 +278,6 @@ def parse_priority(val: Any) -> Optional[int]:
|
||||
|
||||
|
||||
def rule_text_from_expr(expr: Any) -> str:
|
||||
"""
|
||||
Deterministic serializer to produce a compact UI-friendly string from expr list.
|
||||
"""
|
||||
if expr is None:
|
||||
return ""
|
||||
if isinstance(expr, list):
|
||||
@@ -348,8 +304,6 @@ def rule_text_from_expr(expr: Any) -> str:
|
||||
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:
|
||||
@@ -366,7 +320,7 @@ def rule_text_from_expr(expr: Any) -> str:
|
||||
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"):
|
||||
if q.get("bypass") or q.get("flags") == "bypass":
|
||||
token += " bypass"
|
||||
elif isinstance(q, (int, float)):
|
||||
token += f" num {int(q)}"
|
||||
@@ -374,8 +328,7 @@ def rule_text_from_expr(expr: Any) -> str:
|
||||
token += f" num {q}"
|
||||
tokens.append(token)
|
||||
else:
|
||||
keys = "+".join(sorted(part.keys()))
|
||||
tokens.append(keys)
|
||||
tokens.append("+".join(sorted(part.keys())))
|
||||
else:
|
||||
tokens.append(str(part))
|
||||
return " ".join(tokens)
|
||||
@@ -483,7 +436,6 @@ def build_predictable_ruleset(nft_json: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return result
|
||||
|
||||
|
||||
# ---------- New helper: populate_text_from_chain_text ----------
|
||||
def populate_text_from_chain_text(custom: Dict[str, Any]) -> None:
|
||||
tables = custom.get("tables") or []
|
||||
for t in tables:
|
||||
@@ -523,23 +475,13 @@ def populate_text_from_chain_text(custom: Dict[str, Any]) -> None:
|
||||
replaced = True
|
||||
break
|
||||
except Exception as e:
|
||||
logger.debug(
|
||||
"populate_text_from_chain_text: failed to get textual chain for %s %s %s: %s",
|
||||
fam,
|
||||
tname,
|
||||
cname,
|
||||
e,
|
||||
)
|
||||
logger.debug("populate_text_from_chain_text failed for %s %s %s: %s", fam, tname, cname, e)
|
||||
continue
|
||||
|
||||
|
||||
# ---------- Routes ----------
|
||||
|
||||
@router.get("/rules", response_model=RulesetOut, summary="List ruleset")
|
||||
def list_rules():
|
||||
"""
|
||||
Returns the ruleset in a stable, strongly-typed JSON shape derived from `nft -j list ruleset`.
|
||||
"""
|
||||
try:
|
||||
try:
|
||||
nft_json = mgr.list_rules_json()
|
||||
@@ -555,7 +497,7 @@ def list_rules():
|
||||
except Exception as e:
|
||||
logger.debug("list_rules: populate_text_from_chain_text failed: %s", e)
|
||||
|
||||
# IMPORTANT: return a plain dict (not a JSON string) so FastAPI/Pydantic serializes it naturally.
|
||||
# IMPORTANT: return native Python structure (dict) — not a JSON string.
|
||||
ruleset_model = RulesetModel.parse_obj(custom)
|
||||
return RulesetOut(ruleset=ruleset_model.dict())
|
||||
except NftError as e:
|
||||
@@ -570,7 +512,7 @@ def list_rules():
|
||||
"/rules",
|
||||
response_model=ExecResult,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
summary="Create rule (JSON, expr required; returns ExecResult with rc/stdout/stderr)",
|
||||
summary="Create rule (JSON, expr required; returns ExecResult)",
|
||||
)
|
||||
def create_rule_json(req: CreateRuleRequest):
|
||||
try:
|
||||
@@ -583,13 +525,9 @@ def create_rule_json(req: CreateRuleRequest):
|
||||
|
||||
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."
|
||||
)
|
||||
raise NftError("cannot render provided 'expr' to textual nft syntax; use /raw to run text command")
|
||||
|
||||
expr_text = rendered.strip()
|
||||
|
||||
if req.position is not None:
|
||||
try:
|
||||
pos = int(req.position)
|
||||
@@ -601,41 +539,33 @@ def create_rule_json(req: CreateRuleRequest):
|
||||
else:
|
||||
cmd = f"add rule {family} {table} {chain} {expr_text}"
|
||||
|
||||
logger.info("create_rule_json executing command: %s", cmd)
|
||||
|
||||
logger.info("create_rule_json executing: %s", cmd)
|
||||
res = mgr.cmd(cmd)
|
||||
raw_rc = res.get("rc")
|
||||
stdout = res.get("stdout") or ""
|
||||
stderr = res.get("stderr") or ""
|
||||
|
||||
try:
|
||||
rc = int(raw_rc)
|
||||
except Exception:
|
||||
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:
|
||||
return exec_res
|
||||
|
||||
if (rc < 0 or rc != 0) and stderr.strip() == "":
|
||||
logger.debug("create_rule_json: rc indicates failure but stderr empty; verifying rule presence")
|
||||
try:
|
||||
chain_text = mgr.list_chain_text(family, table, chain) or ""
|
||||
if expr_text and expr_text in chain_text:
|
||||
logger.info("create_rule_json: detected rule in chain after add; treating as success")
|
||||
logger.info("create_rule_json: detected rule after add; treating as success")
|
||||
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)
|
||||
logger.warning("create_rule_json: verification failed: %s", e_chain)
|
||||
|
||||
detail = f"nft command failed rc={rc}. stderr: {stderr!r}. cmd: {cmd}"
|
||||
logger.warning("create_rule_json failed: %s", detail)
|
||||
logger.warning(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))
|
||||
|
||||
Reference in New Issue
Block a user