diff --git a/backend/src/api/nftables_api.py b/backend/src/api/nftables_api.py index ea6b327..4930efc 100644 --- a/backend/src/api/nftables_api.py +++ b/backend/src/api/nftables_api.py @@ -1,52 +1,44 @@ +# fastapi_nft_replace.py """ -FastAPI app to manage nftables bridge rules dynamically. -- Removed idempotent-by-id behavior: POST /rules always adds rules (generates id if missing) -- Added DELETE /rules endpoint to delete one or many rules by id +nftables router for FastAPI to manage nftables bridge rules dynamically. + +Endpoints (mounted under /nft): +- GET /rules -> list active rules (read from nftables) +- POST /rules -> add one or many rules (append) +- DELETE /rules -> delete one or many rules by id +- PUT /rules -> replace entire ordered rule set via nft -f (returns new version) Notes: -- Runs nft(8) commands; the process must have sufficient privileges (run as root or via sudo). -- This implementation stores the full rule JSON inside the nft rule comment as base64 to allow round-trip parsing. -- The API keeps sniffer separate; this module only manages kernel rules. - -Endpoints: -- GET /rules -> list active rules (read from nftables) -- POST /rules -> add one or many rules -- DELETE /rules -> delete one or many rules by id - -Rule schema (example): -{ - "id": "optional-uuid-if-you-want", - "table": "mitm_tbl", - "chain": "forward", - "family": "bridge", - "match": { - "iif": "br0", - "oif": "eth1", - "meta_length": "100-200", # or single int as str/number - "ip_proto": "tcp", - "tcp_dport": 80 - }, - "action": {"type": "drop"} -} - -Supported matches in this example: iif, oif, meta_length, ip_proto, tcp_dport, udp_dport -Supported actions: drop, accept, queue (num), redirect (port) +- Must run with privileges to call `nft` (root or via sudo). +- Rules are stored in the nft rule comment as base64-encoded JSON for round-trip parsing. +- PUT /rules writes a temporary nft script and runs `nft -f `; this applies the new ordered rules. """ - -from fastapi import APIRouter, APIRouter, HTTPException +from fastapi import APIRouter, HTTPException from pydantic import BaseModel, Field -from typing import Optional, List, Dict, Any, Union +from typing import Optional, List, Dict, Any, Union, Tuple import subprocess import uuid import json import base64 import re +import tempfile +import os +import logging +# Router and logger --------------------------------------------------------- +router = APIRouter() +logger = logging.getLogger("nftables") +logger.debug("nftables router module loaded") + +# Defaults and in-memory version token -------------------------------------- DEFAULT_TABLE = "mitm_tbl" DEFAULT_CHAIN = "forward" DEFAULT_FAMILY = "bridge" +# in-memory version token updated on successful PUT +_current_version: Optional[str] = None + # ---------------------- Pydantic models ---------------------- class MatchModel(BaseModel): iif: Optional[str] @@ -56,11 +48,13 @@ class MatchModel(BaseModel): tcp_dport: Optional[int] udp_dport: Optional[int] + class ActionModel(BaseModel): type: str # drop | accept | queue | redirect queue_num: Optional[int] redirect_port: Optional[int] + class RuleModel(BaseModel): id: Optional[str] = Field(None, description="optional rule id; generated if missing") family: Optional[str] = Field(DEFAULT_FAMILY) @@ -69,35 +63,48 @@ class RuleModel(BaseModel): match: MatchModel action: ActionModel -# ---------------------- Utilities ---------------------- -def run_nft(args: List[str]) -> subprocess.CompletedProcess: +class ReplaceResult(BaseModel): + version: str + applied: bool + rules_count: int + + +# ---------------------- Utilities ---------------------- +def run_nft(args: List[str]) -> Tuple[str, str]: + """ + Run nft with given args. + Returns (stdout, stderr). + Raises RuntimeError on non-zero exit with stderr included. + """ + logger.debug("running nft: %s", " ".join(["nft"] + args)) try: - return subprocess.run(["nft"] + args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True) + proc = subprocess.run(["nft"] + args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True) + logger.debug("nft stdout: %s", proc.stdout.strip()) + return proc.stdout, proc.stderr except subprocess.CalledProcessError as e: - # raise with stderr for easier debugging + logger.error("nft failed: %s -- %s", " ".join(e.cmd), e.stderr.strip()) raise RuntimeError(f"nft failed: {' '.join(e.cmd)} -- {e.stderr.strip()}") -def ensure_table_chain(family: str, table: str, chain: str): +def ensure_table_chain(family: str, table: str, chain: str) -> None: # create table if missing (ignore error if exists) try: + logger.info("ensuring table %s.%s exists", family, table) run_nft(["add", "table", family, table]) except RuntimeError: - # already exists or failed; ignore existence error - pass + logger.debug("table %s.%s may already exist", family, table) # create chain if missing: forward chain with hook forward try: - # Note: use type filter hook forward priority 0 + logger.info("ensuring chain %s in table %s exists", chain, table) run_nft(["add", "chain", family, table, chain, "{", "type", "filter", "hook", "forward", "priority", "0", ";", "}"]) except RuntimeError: - # ignore if exists - pass + logger.debug("chain %s in table %s may already exist", chain, table) def encode_rule_comment(rule: Dict[str, Any]) -> str: - # store the rule JSON as base64 to avoid quoting/escaping issues inside nft comment + # store rule JSON as base64 to avoid quoting/escaping issues inside nft comment j = json.dumps(rule, separators=(",", ":")) b = base64.b64encode(j.encode()).decode() rid = rule.get("id") or "" @@ -115,28 +122,28 @@ def decode_comment_payload(comment: str) -> Optional[Dict[str, Any]]: j = base64.b64decode(b64.encode()).decode() return json.loads(j) except Exception: + logger.debug("failed to decode comment payload: %s", comment) return None -def build_nft_match_expr(match: MatchModel) -> List[str]: - expr: List[str] = [] +def build_nft_match_fragment(match: MatchModel) -> List[str]: + frag: List[str] = [] if match.iif: - expr += ["iif", match.iif] + frag += ["iif", f'"{match.iif}"'] if match.oif: - expr += ["oif", match.oif] + frag += ["oif", f'"{match.oif}"'] if match.meta_length is not None: - # accept either range string or int - expr += ["meta", "length", str(match.meta_length)] + frag += ["meta", "length", str(match.meta_length)] if match.ip_proto: - expr += ["ip", "protocol", match.ip_proto] + frag += ["ip", "protocol", match.ip_proto] if match.tcp_dport: - expr += ["tcp", "dport", str(match.tcp_dport)] + frag += ["tcp", "dport", str(match.tcp_dport)] if match.udp_dport: - expr += ["udp", "dport", str(match.udp_dport)] - return expr + frag += ["udp", "dport", str(match.udp_dport)] + return frag -def build_nft_action_expr(action: ActionModel) -> List[str]: +def build_nft_action_fragment(action: ActionModel) -> List[str]: if action.type == "drop": return ["drop"] if action.type == "accept": @@ -151,81 +158,77 @@ def build_nft_action_expr(action: ActionModel) -> List[str]: raise ValueError(f"unsupported action type: {action.type}") -def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]: - # ensure table/chain - ensure_table_chain(rule.family, rule.table, rule.chain) - - # ensure id - if not rule.id: - rule.id = str(uuid.uuid4()) - - # build nft command args - match_expr = build_nft_match_expr(rule.match) - action_expr = build_nft_action_expr(rule.action) - +def nft_rule_line_from_model(rule: RuleModel) -> str: + """ + Produce a single-line nft command: + add rule comment "" + """ + match_frag = build_nft_match_fragment(rule.match) + action_frag = build_nft_action_fragment(rule.action) comment = encode_rule_comment(rule.dict()) - args: List[str] = ["add", "rule", rule.family, rule.table, rule.chain] - args += match_expr - args += action_expr - args += ["comment", comment] - - try: - run_nft(args) - return {"id": rule.id, "status": "added"} - except RuntimeError as e: - raise HTTPException(status_code=500, detail=str(e)) + parts = ["add", "rule", rule.family, rule.table, rule.chain] + parts += match_frag + parts += action_frag + return " ".join(parts) + f' comment "{comment}"' def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]: try: - out = run_nft(["list", "table", family, table]).stdout + out, _ = run_nft(["list", "table", family, table]) except RuntimeError: + logger.debug("no table %s.%s found when listing rules", family, table) return [] results: List[Dict[str, Any]] = [] - # naive parse: nft prints individual rules as lines; look for comment "mitm_json:" for line in out.splitlines(): line = line.strip() if "comment" in line and "mitm_json:" in line: - # find comment payload part: comment "..." - # format often: "comment "mitm_id:... mitm_json:..."" try: - # extract between first pair of double quotes - first_quote = line.index('\"') - last_quote = line.rindex('\"') + first_quote = line.index('"') + last_quote = line.rindex('"') comment_str = line[first_quote + 1:last_quote] except ValueError: - # fallback: take substring after comment comment_str = line.split("comment", 1)[1].strip() parsed = decode_comment_payload(comment_str) if parsed is not None: results.append(parsed) + logger.debug("listed %d nft rules from %s.%s", len(results), family, table) return results -def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> List[str]: - """Delete rules whose comment contains mitm_id in ids. - Returns list of deleted ids. - """ +def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]: + ensure_table_chain(rule.family, rule.table, rule.chain) + if not rule.id: + rule.id = str(uuid.uuid4()) + + cmd_text = nft_rule_line_from_model(rule) try: - out = run_nft(["list", "chain", family, table, chain, "-a"]).stdout + # run as discrete args to avoid shell quoting issues + run_nft(cmd_text.split()) + logger.info("added nft rule id=%s family=%s table=%s chain=%s", rule.id, rule.family, rule.table, rule.chain) + return {"id": rule.id, "status": "added"} + except RuntimeError as e: + logger.error("failed to add rule id=%s: %s", rule.id, e) + raise HTTPException(status_code=500, detail=str(e)) + + +def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> List[str]: + try: + out, _ = run_nft(["list", "chain", family, table, chain, "-a"]) except RuntimeError: + logger.debug("no chain %s in table %s.%s when attempting delete", chain, family, table) return [] deleted: List[str] = [] - - # nft -a prints rules; each rule line may contain a comment and a trailing handle number: "... comment \\"mitm_id:...\\" ... handle 5" for line in out.splitlines(): if "comment" in line and "mitm_id:" in line: - # extract comment between quotes try: - q1 = line.index('\"') - q2 = line.index('\"', q1 + 1) + q1 = line.index('"') + q2 = line.index('"', q1 + 1) comment_str = line[q1 + 1:q2] except ValueError: - # fallback: substring comment_str = line.split("comment", 1)[1] parsed = decode_comment_payload(comment_str) @@ -233,10 +236,8 @@ def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> continue rid = parsed.get("id") if rid in ids: - # find handle number m = re.search(r"handle\s+(\d+)", line) if not m: - # try to find handle on the next token(s) - naive fallback parts = line.split() if "handle" in parts: hi = parts.index("handle") @@ -251,24 +252,23 @@ def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> try: run_nft(["delete", "rule", family, table, chain, "handle", handle]) + logger.info("deleted nft rule id=%s handle=%s", rid, handle) deleted.append(rid) except RuntimeError: - # ignore deletion errors for now + logger.error("failed to delete nft rule id=%s handle=%s", rid, handle) continue return deleted + # ---------------------- API endpoints ---------------------- - -router = APIRouter() - @router.get("/rules") def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE): rules = list_rules_from_nft(family, table) - return {"count": len(rules), "rules": rules} + return {"count": len(rules), "rules": rules, "version": _current_version} + @router.post("/rules") def post_rules(payload: Union[RuleModel, List[RuleModel]]): - # accept either single or list rules = payload if isinstance(payload, list) else [payload] results = [] for r in rules: @@ -276,12 +276,88 @@ def post_rules(payload: Union[RuleModel, List[RuleModel]]): results.append(res) return {"results": results} + @router.delete("/rules") -def delete_rules(payload: Union[str, List[str]], family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): - """Delete one or many rules by id. - Payload can be a single id string or a list of ids. - """ +def delete_rules(payload: Union[str, List[str]], family: Optional[str] = DEFAULT_FAMILY, + table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): ids = [payload] if isinstance(payload, str) else payload deleted = delete_rules_by_ids(ids, family, table, chain) results = [{"id": i, "deleted": i in deleted} for i in ids] return {"results": results} + + +@router.put("/rules", response_model=ReplaceResult) +def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY, + table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): + """ + Replace entire ordered rule set by generating an nft script and applying via `nft -f`. + + Behavior: + 1. Validate rules and ensure family/table/chain match (if provided). + 2. Ensure the table/chain exist. + 3. Generate nft script that flushes the chain and adds rules in given order. + 4. Write script to a secure temp file and run `nft -f `. + 5. On success update in-memory version token and return it. + + Note: nft runs script sequentially; if nft errors mid-script, partial state may exist. + For stricter atomicity, implement the temp-table swap approach. + """ + # validate per-rule family/table/chain if present + for r in rules: + if r.family and r.family != family: + raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}") + if r.table and r.table != table: + raise HTTPException(status_code=400, detail=f"rule table mismatch: {r.table} != {table}") + if r.chain and r.chain != chain: + raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}") + + # ensure table/chain exist + ensure_table_chain(family, table, chain) + + # ensure rule IDs + for r in rules: + if not r.id: + r.id = str(uuid.uuid4()) + + # build nft script lines + lines: List[str] = [] + # flush chain (clear existing ordered rules) + lines.append(f"flush chain {family} {table} {chain}") + + # add ordered rules + for r in rules: + try: + lines.append(nft_rule_line_from_model(r)) + except Exception as e: + raise HTTPException(status_code=400, detail=f"invalid rule: {e}") + + script = "\n".join(lines) + "\n" + + tmpfile_path: Optional[str] = None + try: + with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf: + tmpfile_path = tf.name + tf.write(script) + tf.flush() + os.fsync(tf.fileno()) + logger.info("wrote nft script to %s; applying...", tmpfile_path) + + # apply the file + try: + run_nft(["-f", tmpfile_path]) + except RuntimeError as e: + logger.error("failed applying nft script: %s", e) + raise HTTPException(status_code=500, detail=f"failed applying nft script: {e}") + + # success: bump version token + global _current_version + _current_version = str(uuid.uuid4()) + logger.info("applied nft ruleset successfully; version=%s", _current_version) + return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules)) + finally: + if tmpfile_path and os.path.exists(tmpfile_path): + try: + os.remove(tmpfile_path) + except Exception: + logger.debug("failed to remove temp nft script %s", tmpfile_path) +