This commit is contained in:
@@ -22,14 +22,13 @@ import uuid
|
||||
from enum import Enum
|
||||
import asyncio
|
||||
import subprocess
|
||||
import re
|
||||
|
||||
# Try to import NFTables binding (various pyroute2 layouts)
|
||||
try:
|
||||
# common location
|
||||
from pyroute2.nftables.main import NFTables as NFTablesBinding # type: ignore
|
||||
except Exception:
|
||||
try:
|
||||
# alternate location
|
||||
from pyroute2.nftables import NFTables as NFTablesBinding # type: ignore
|
||||
except Exception:
|
||||
NFTablesBinding = None # will fall back to subprocess wrapper
|
||||
@@ -102,12 +101,9 @@ class ReplaceResult(BaseModel):
|
||||
|
||||
# ---------------------- resilient NFT wrapper selection ----------------------
|
||||
class NFTSubprocessWrapper:
|
||||
"""
|
||||
Fallback wrapper calling the `nft` CLI via subprocess.
|
||||
"""
|
||||
"""Fallback wrapper calling the `nft` CLI via subprocess."""
|
||||
|
||||
def __init__(self):
|
||||
# log availability
|
||||
try:
|
||||
subprocess.run(["nft", "--version"], check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
|
||||
logger.debug("nft CLI available for subprocess wrapper")
|
||||
@@ -115,7 +111,6 @@ class NFTSubprocessWrapper:
|
||||
logger.warning("nft CLI not available or couldn't be invoked: %s", e)
|
||||
|
||||
def run(self, cmd: str) -> Dict[str, Any]:
|
||||
# naive split is acceptable for the limited command patterns used here
|
||||
args = ["nft"] + cmd.split()
|
||||
try:
|
||||
proc = subprocess.run(args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True)
|
||||
@@ -177,7 +172,6 @@ class NFTBindingWrapper:
|
||||
last_exc = e
|
||||
|
||||
if not self._constructed:
|
||||
# try no-arg
|
||||
try:
|
||||
self._binding = binding_cls()
|
||||
self._constructed = True
|
||||
@@ -191,10 +185,8 @@ class NFTBindingWrapper:
|
||||
# If the binding exposes an async setup coroutine, ensure we can run it synchronously.
|
||||
setup_coro = getattr(self._binding, "setup_endpoint", None)
|
||||
if setup_coro and asyncio.iscoroutinefunction(setup_coro):
|
||||
# If event loop running, we cannot await here — signal unsuitability.
|
||||
if asyncio.get_event_loop().is_running():
|
||||
raise RuntimeError("pyroute2 NFTables requires async setup but event loop is already running")
|
||||
# otherwise run it to complete setup
|
||||
try:
|
||||
asyncio.get_event_loop().run_until_complete(setup_coro())
|
||||
except Exception as e:
|
||||
@@ -231,7 +223,6 @@ def make_nft_wrapper():
|
||||
"""
|
||||
if NFTablesBinding is not None:
|
||||
try:
|
||||
# If an event loop is running (uvicorn), prefer subprocess to avoid async binding setup.
|
||||
if asyncio.get_event_loop().is_running():
|
||||
logger.info("asyncio loop is running; skipping binding and using subprocess wrapper")
|
||||
raise RuntimeError("event loop running")
|
||||
@@ -240,7 +231,6 @@ def make_nft_wrapper():
|
||||
return w
|
||||
except Exception as e:
|
||||
logger.warning("pyroute2 binding unavailable/synchronous construction failed: %s; falling back to nft CLI", e)
|
||||
# Fallback
|
||||
logger.info("using nft CLI subprocess wrapper")
|
||||
return NFTSubprocessWrapper()
|
||||
|
||||
@@ -309,6 +299,10 @@ def parse_comment_from_exprs(exprs: List[Dict[str, Any]]) -> Optional[str]:
|
||||
|
||||
|
||||
def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert pyroute2/nft `rule` JSON entry into the simple dict format returned by the API.
|
||||
Best-effort: handles common expression shapes; if not present returns defaults.
|
||||
"""
|
||||
exprs = []
|
||||
if "rule" in entry:
|
||||
r = entry["rule"]
|
||||
@@ -423,6 +417,128 @@ def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
# ---------------------- helpers for normalized output ----------------------
|
||||
def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Convert various wrapper outputs into a list of nft 'entries'.
|
||||
Handles shapes:
|
||||
- {'nftables': [ ... ]}
|
||||
- {'out': '<json text>'}
|
||||
- list([...])
|
||||
- dict (single entry)
|
||||
- {'out': '<plain text>'} -> returned as [{'text': <...>}]
|
||||
"""
|
||||
if not out:
|
||||
return []
|
||||
if isinstance(out, dict):
|
||||
if "nftables" in out and isinstance(out["nftables"], list):
|
||||
return out["nftables"]
|
||||
if "out" in out and isinstance(out["out"], str):
|
||||
text = out["out"]
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
if isinstance(parsed, dict) and "nftables" in parsed:
|
||||
return parsed["nftables"]
|
||||
if isinstance(parsed, list):
|
||||
return parsed
|
||||
return [parsed]
|
||||
except Exception:
|
||||
# textual fallback
|
||||
return [{"text": text}]
|
||||
# single dict -> wrap
|
||||
return [out]
|
||||
if isinstance(out, list):
|
||||
return out
|
||||
if isinstance(out, str):
|
||||
return [{"text": out}]
|
||||
return []
|
||||
|
||||
|
||||
def _parse_textual_chain_output(text: str, family: str, table: str, chain: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Parse textual chain dump for simple rules as a last-resort fallback.
|
||||
Looks for lines inside the chain and extracts simple patterns.
|
||||
"""
|
||||
results: List[Dict[str, Any]] = []
|
||||
lines = text.splitlines()
|
||||
for raw in lines:
|
||||
line = raw.strip()
|
||||
if not line:
|
||||
continue
|
||||
if line.startswith("type ") or line.startswith("policy ") or line.startswith("chain ") or line.startswith("table "):
|
||||
continue
|
||||
if line in ("{", "}"):
|
||||
continue
|
||||
|
||||
comment = None
|
||||
cm = re.search(r'comment\s+"([^"]+)"', line)
|
||||
if cm:
|
||||
comment = cm.group(1)
|
||||
line_no_comment = re.sub(r'comment\s+"[^"]+"', '', line)
|
||||
else:
|
||||
line_no_comment = line
|
||||
|
||||
match: Dict[str, Any] = {}
|
||||
action: Dict[str, Any] = {"type": "accept"}
|
||||
|
||||
m = re.search(r'\bip\s+protocol\s+([A-Za-z0-9_+-]+)\b', line_no_comment, flags=re.IGNORECASE)
|
||||
if m:
|
||||
proto = m.group(1).lower()
|
||||
if proto == "icmp":
|
||||
match["ip_proto"] = 1
|
||||
elif proto == "tcp":
|
||||
match["ip_proto"] = 6
|
||||
elif proto == "udp":
|
||||
match["ip_proto"] = 17
|
||||
else:
|
||||
try:
|
||||
match["ip_proto"] = int(proto)
|
||||
except Exception:
|
||||
match["ip_proto"] = proto
|
||||
|
||||
m2 = re.search(r'\btcp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE)
|
||||
if m2:
|
||||
match["tcp_dport"] = int(m2.group(1))
|
||||
|
||||
m3 = re.search(r'\budp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE)
|
||||
if m3:
|
||||
match["udp_dport"] = int(m3.group(1))
|
||||
|
||||
m4 = re.search(r'\biif\s+"?([^"\s]+)"?\b', line_no_comment, flags=re.IGNORECASE)
|
||||
if m4:
|
||||
match["iif"] = m4.group(1)
|
||||
m5 = re.search(r'\boif\s+"?([^"\s]+)"?\b', line_no_comment, flags=re.IGNORECASE)
|
||||
if m5:
|
||||
match["oif"] = m5.group(1)
|
||||
|
||||
if re.search(r'\bdrop\b', line_no_comment, flags=re.IGNORECASE):
|
||||
action = {"type": "drop"}
|
||||
elif re.search(r'\baccept\b', line_no_comment, flags=re.IGNORECASE):
|
||||
action = {"type": "accept"}
|
||||
elif re.search(r'\bqueue\b', line_no_comment, flags=re.IGNORECASE):
|
||||
q = re.search(r'queue\s+num\s+(\d+)', line_no_comment, flags=re.IGNORECASE)
|
||||
if q:
|
||||
action = {"type": "queue", "queue_num": int(q.group(1))}
|
||||
else:
|
||||
action = {"type": "queue"}
|
||||
elif re.search(r'\bredirect\b', line_no_comment, flags=re.IGNORECASE) or re.search(r'\bto\s+:\d+\b', line_no_comment, flags=re.IGNORECASE):
|
||||
mredir = re.search(r':(\d+)', line_no_comment)
|
||||
if mredir:
|
||||
action = {"type": "redirect", "redirect_port": int(mredir.group(1))}
|
||||
else:
|
||||
action = {"type": "redirect"}
|
||||
|
||||
results.append({
|
||||
"family": family,
|
||||
"table": table,
|
||||
"chain": chain,
|
||||
"match": match,
|
||||
"action": action,
|
||||
"comment": comment,
|
||||
})
|
||||
return results
|
||||
|
||||
|
||||
# ---------------------- high-level operations ----------------------
|
||||
def ensure_table_chain(family: Union[str, Family], table: str, chain: str):
|
||||
fam = family.value if isinstance(family, Family) else family
|
||||
@@ -436,44 +552,64 @@ def ensure_table_chain(family: Union[str, Family], table: str, chain: str):
|
||||
logger.debug("add_chain may have failed/exists: %s", e)
|
||||
|
||||
|
||||
def list_rules_from_nft(family: Union[str, Family], table: str) -> List[Dict[str, Any]]:
|
||||
def list_rules_from_nft(family: Union[str, Family], table: str, chain: str = DEFAULT_CHAIN) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Return reconstructed rules from kernel state.
|
||||
|
||||
Strategy:
|
||||
1. Try `list chain <family> <table> <chain> -a` (best for rules + handles).
|
||||
2. If that fails or returns textual output, fall back to `list ruleset` and extract all 'rule' entries.
|
||||
3. If structured JSON is unavailable, parse textual output as last resort.
|
||||
"""
|
||||
fam = family.value if isinstance(family, Family) else family
|
||||
|
||||
# primary: list chain -a
|
||||
try:
|
||||
out = NFTC.list_table(fam, table)
|
||||
raw = NFTC.list_chain(fam, table, chain)
|
||||
logger.debug("raw list_chain output: %s", str(raw)[:2000])
|
||||
except Exception as e:
|
||||
logger.debug("list_table failed: %s", e)
|
||||
logger.debug("list_chain failed (%s); falling back to list ruleset", e)
|
||||
try:
|
||||
raw = NFTC.run("list ruleset")
|
||||
logger.debug("raw list ruleset output: %s", str(raw)[:2000])
|
||||
except Exception as e2:
|
||||
logger.error("list ruleset failed: %s", e2)
|
||||
return []
|
||||
|
||||
entries = []
|
||||
if isinstance(out, dict) and "nftables" in out:
|
||||
entries = out["nftables"]
|
||||
elif isinstance(out, dict) and "out" in out and isinstance(out["out"], str):
|
||||
try:
|
||||
parsed = json.loads(out["out"])
|
||||
if isinstance(parsed, dict) and "nftables" in parsed:
|
||||
entries = parsed["nftables"]
|
||||
elif isinstance(parsed, list):
|
||||
entries = parsed
|
||||
except Exception:
|
||||
logger.debug("could not parse textual nft output")
|
||||
entries = []
|
||||
elif isinstance(out, list):
|
||||
entries = out
|
||||
elif isinstance(out, dict) and out:
|
||||
entries = [out]
|
||||
else:
|
||||
entries = []
|
||||
|
||||
entries = _normalize_nft_output(raw)
|
||||
results: List[Dict[str, Any]] = []
|
||||
|
||||
for item in entries:
|
||||
if not item:
|
||||
continue
|
||||
# textual full-dump fallback
|
||||
if isinstance(item, dict) and "text" in item and isinstance(item["text"], str):
|
||||
results.extend(_parse_textual_chain_output(item["text"], fam, table, chain))
|
||||
continue
|
||||
|
||||
# structured rule entry
|
||||
if "rule" in item:
|
||||
rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")}
|
||||
else:
|
||||
rule_entry = item
|
||||
reconstructed = reconstruct_rule_from_rule_entry(rule_entry)
|
||||
results.append(reconstructed)
|
||||
continue
|
||||
|
||||
# some outputs embed table/chain with nested lists (from list ruleset)
|
||||
# find nested 'rule' if present
|
||||
if isinstance(item, dict):
|
||||
# item may be { 'table': {...} } or { 'chain': {...} } etc.
|
||||
# attempt to find inner rule keys recursively
|
||||
if "nftables" in item and isinstance(item["nftables"], list):
|
||||
for sub in item["nftables"]:
|
||||
if isinstance(sub, dict) and "rule" in sub:
|
||||
reconstructed = reconstruct_rule_from_rule_entry(sub)
|
||||
results.append(reconstructed)
|
||||
elif item.get("type") == "rule" or "expr" in item or "expressions" in item:
|
||||
reconstructed = reconstruct_rule_from_rule_entry(item)
|
||||
results.append(reconstructed)
|
||||
else:
|
||||
# nothing rule-like found here; skip
|
||||
pass
|
||||
|
||||
return results
|
||||
|
||||
@@ -486,12 +622,11 @@ def add_rules_replace_all(rules: List[RuleModel], family: Union[str, Family], ta
|
||||
except Exception as e:
|
||||
logger.debug("flush chain may have returned error: %s", e)
|
||||
|
||||
# add rules in order
|
||||
for r in rules:
|
||||
fam_r = r.family.value if isinstance(r.family, Family) else r.family
|
||||
frag = nft_rule_fragment_from_model(r)
|
||||
try:
|
||||
NFTC.add_rule(fam_r, r.table, r.chain, frag)
|
||||
NFTC.run(f"add rule {fam_r} {r.table} {r.chain} {frag}")
|
||||
logger.info("added rule frag=%s", frag)
|
||||
except Exception as e:
|
||||
logger.exception("failed to add rule: %s", e)
|
||||
@@ -500,7 +635,9 @@ def add_rules_replace_all(rules: List[RuleModel], family: Union[str, Family], ta
|
||||
|
||||
# ---------------------- API endpoints ----------------------
|
||||
@router.get("/rules")
|
||||
def get_rules(family: Optional[Union[str, Family]] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
|
||||
def get_rules(family: Optional[Union[str, Family]] = DEFAULT_FAMILY,
|
||||
table: Optional[str] = DEFAULT_TABLE,
|
||||
chain: Optional[str] = DEFAULT_CHAIN):
|
||||
fam = family.value if isinstance(family, Family) else family
|
||||
logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", fam, table, chain)
|
||||
try:
|
||||
@@ -509,7 +646,7 @@ def get_rules(family: Optional[Union[str, Family]] = DEFAULT_FAMILY, table: Opti
|
||||
logger.error("failed to ensure table/chain: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
rules = list_rules_from_nft(fam, table)
|
||||
rules = list_rules_from_nft(fam, table, chain)
|
||||
return {"count": len(rules), "rules": rules, "version": _current_version}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user