This commit is contained in:
@@ -404,7 +404,6 @@ def nft_list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Di
|
||||
add_command = f"add rule {table_name} {chain_name} {nft_rule_text}".strip() if nft_rule_text else None
|
||||
logger.debug("Attached textual rule for handle %s", handle)
|
||||
else:
|
||||
# If mapping missing, per your instruction do not attempt to reconstruct — leave textual fields None
|
||||
logger.debug("No textual mapping for handle %s — textual fields will be None", handle)
|
||||
|
||||
results.append({
|
||||
|
||||
303
backend/src/api/nft_manager.py
Normal file
303
backend/src/api/nft_manager.py
Normal file
@@ -0,0 +1,303 @@
|
||||
# app.py
|
||||
"""
|
||||
Stateless nftables FastAPI service (libnftables only, no CLI fallback, no annotations).
|
||||
Endpoints:
|
||||
GET /firewall/rules -> list rules (kernel-provided rules with handles)
|
||||
POST /firewall/rules -> add a constrained rule (returns 201)
|
||||
DELETE /firewall/rules/{handle} -> delete rule by handle
|
||||
Requirements:
|
||||
python-nftables must be installed and usable:
|
||||
pip install python-nftables fastapi uvicorn pydantic
|
||||
The process must have CAP_NET_ADMIN (or be run as root) to modify nftables.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from fastapi import FastAPI, APIRouter, HTTPException, status
|
||||
from pydantic import BaseModel, Field
|
||||
import ipaddress
|
||||
import logging
|
||||
|
||||
# libnftables import (must be present)
|
||||
from nftables import Nftables # type: ignore
|
||||
|
||||
# ---------- logging ----------
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("nft_api")
|
||||
|
||||
# ---------- NftManager (libnftables only) ----------
|
||||
class NftError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class NftManager:
|
||||
"""
|
||||
Minimal libnftables wrapper that performs JSON transactions only via python-nftables.
|
||||
No CLI fallback, no comments/annotations.
|
||||
"""
|
||||
|
||||
# Tighten these to your environment
|
||||
ALLOWED_FAMILIES = {"inet"}
|
||||
ALLOWED_TABLES = {"filter"}
|
||||
ALLOWED_CHAINS = {"input", "output", "forward"}
|
||||
ALLOWED_PROTOS = {"tcp", "udp"}
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.nft = Nftables()
|
||||
# Request JSON output where applicable
|
||||
try:
|
||||
self.nft.set_json_output(True)
|
||||
except Exception:
|
||||
# Some libnftables builds may ignore this; proceed anyway
|
||||
logger.debug("set_json_output not available or failed")
|
||||
|
||||
def _json_cmd(self, cmd: Dict[str, Any]) -> Any:
|
||||
rc, out, err = self.nft.json_cmd(cmd)
|
||||
if rc != 0:
|
||||
logger.error("libnftables error: %s", err)
|
||||
raise NftError(err)
|
||||
return out
|
||||
|
||||
def _validate_family_table_chain(self, family: str, table: str, chain: str) -> None:
|
||||
if family not in self.ALLOWED_FAMILIES:
|
||||
raise ValueError(f"family '{family}' not allowed")
|
||||
if table not in self.ALLOWED_TABLES:
|
||||
raise ValueError(f"table '{table}' not allowed")
|
||||
if chain not in self.ALLOWED_CHAINS:
|
||||
raise ValueError(f"chain '{chain}' not allowed")
|
||||
|
||||
def _validate_proto(self, proto: Optional[str]) -> None:
|
||||
if proto is None:
|
||||
return
|
||||
if proto not in self.ALLOWED_PROTOS:
|
||||
raise ValueError(f"protocol '{proto}' not allowed")
|
||||
|
||||
def _validate_src(self, src: Optional[str]) -> None:
|
||||
if not src:
|
||||
return
|
||||
try:
|
||||
# allow host or network
|
||||
ipaddress.ip_network(src, strict=False)
|
||||
except Exception as e:
|
||||
raise ValueError(f"invalid src '{src}': {e}")
|
||||
|
||||
def _validate_port(self, port: Optional[int]) -> None:
|
||||
if port is None:
|
||||
return
|
||||
if not (1 <= port <= 65535):
|
||||
raise ValueError(f"invalid port: {port}")
|
||||
|
||||
def list_rules(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Return raw rule dicts extracted from libnftables JSON output.
|
||||
Each item is the inner 'rule' dict returned by nftables JSON.
|
||||
"""
|
||||
out = self._json_cmd({"nftables": [{"list": {"ruleset": None}}]})
|
||||
rules: List[Dict[str, Any]] = []
|
||||
for item in out.get("nftables", []):
|
||||
if "rule" in item:
|
||||
rules.append(item["rule"])
|
||||
return rules
|
||||
|
||||
def add_rule(
|
||||
self,
|
||||
family: str,
|
||||
table: str,
|
||||
chain: str,
|
||||
proto: Optional[str],
|
||||
src: Optional[str],
|
||||
dst_port: Optional[int],
|
||||
verdict: str,
|
||||
) -> None:
|
||||
"""
|
||||
Add a constrained rule via libnftables JSON.
|
||||
This constructs an 'expr' list that covers:
|
||||
- ip saddr match (if src provided)
|
||||
- l4 proto match (if proto provided)
|
||||
- dport match (if dst_port provided)
|
||||
- final verdict (accept/drop)
|
||||
"""
|
||||
# validate
|
||||
self._validate_family_table_chain(family, table, chain)
|
||||
self._validate_proto(proto)
|
||||
self._validate_src(src)
|
||||
self._validate_port(dst_port)
|
||||
if verdict not in {"accept", "drop"}:
|
||||
raise ValueError("verdict must be 'accept' or 'drop'")
|
||||
|
||||
expr: List[Dict[str, Any]] = []
|
||||
|
||||
# match source address (IP)
|
||||
if src:
|
||||
# Construct prefix match using nft JSON structure (ip payload -> saddr)
|
||||
# The json shape used here is compatible with python-nftables expectations.
|
||||
# We use ip payload field and prefix match for networks.
|
||||
try:
|
||||
net = ipaddress.ip_network(src, strict=False)
|
||||
addr = str(net.network_address)
|
||||
prefix_len = net.prefixlen
|
||||
except Exception:
|
||||
# fallback to treating as single IP
|
||||
addr = src
|
||||
prefix_len = 32 if ":" not in src else 128
|
||||
expr.append({
|
||||
"match": {
|
||||
"left": {"payload": {"protocol": "ip", "field": "saddr"}},
|
||||
"op": "==",
|
||||
"right": {"prefix": {"addr": addr, "len": prefix_len}}
|
||||
}
|
||||
})
|
||||
|
||||
# match L4 protocol
|
||||
if proto:
|
||||
# meta.l4proto match is a common approach in JSON exprs
|
||||
expr.append({
|
||||
"match": {
|
||||
"left": {"meta": {"key": "l4proto"}},
|
||||
"op": "==",
|
||||
"right": proto
|
||||
}
|
||||
})
|
||||
|
||||
# match destination port (only meaningful if proto provided)
|
||||
if dst_port:
|
||||
if not proto:
|
||||
raise ValueError("dst_port requires proto to be set")
|
||||
expr.append({
|
||||
"match": {
|
||||
"left": {"payload": {"protocol": proto, "field": "dport"}},
|
||||
"op": "==",
|
||||
"right": dst_port
|
||||
}
|
||||
})
|
||||
|
||||
# final verdict
|
||||
expr.append({"verdict": verdict})
|
||||
|
||||
payload = {
|
||||
"nftables": [{
|
||||
"add": {
|
||||
"rule": {
|
||||
"family": family,
|
||||
"table": table,
|
||||
"chain": chain,
|
||||
"expr": expr
|
||||
}
|
||||
}
|
||||
}]
|
||||
}
|
||||
|
||||
self._json_cmd(payload)
|
||||
|
||||
def delete_rule(self, family: str, table: str, chain: str, handle: int) -> None:
|
||||
"""
|
||||
Delete a rule by handle (authoritative).
|
||||
"""
|
||||
self._validate_family_table_chain(family, table, chain)
|
||||
if not isinstance(handle, int) or handle <= 0:
|
||||
raise ValueError("handle must be a positive integer")
|
||||
|
||||
payload = {
|
||||
"nftables": [{
|
||||
"delete": {
|
||||
"rule": {
|
||||
"family": family,
|
||||
"table": table,
|
||||
"chain": chain,
|
||||
"handle": handle
|
||||
}
|
||||
}
|
||||
}]
|
||||
}
|
||||
self._json_cmd(payload)
|
||||
|
||||
|
||||
# ---------- FastAPI application (router combined) ----------
|
||||
router = APIRouter(prefix="/firewall", tags=["firewall"])
|
||||
manager = NftManager()
|
||||
|
||||
|
||||
class CreateRuleRequest(BaseModel):
|
||||
family: str = Field("inet", description="family (whitelisted)")
|
||||
table: str = Field("filter", description="table (whitelisted)")
|
||||
chain: str = Field("input", description="chain (whitelisted)")
|
||||
proto: Optional[str] = Field(None, description="tcp|udp")
|
||||
src: Optional[str] = Field(None, description="source IP or CIDR")
|
||||
dst_port: Optional[int] = Field(None, description="destination port")
|
||||
verdict: str = Field("accept", description="accept|drop")
|
||||
|
||||
|
||||
class RuleOut(BaseModel):
|
||||
family: Optional[str]
|
||||
table: Optional[str]
|
||||
chain: Optional[str]
|
||||
handle: Optional[int]
|
||||
expr: Optional[Any]
|
||||
|
||||
|
||||
@router.get("/rules", response_model=List[RuleOut])
|
||||
def list_rules():
|
||||
"""
|
||||
Returns a list of rules (raw rule dicts from nftables) that include kernel handles.
|
||||
Clients should save handles if they want to delete rules later.
|
||||
"""
|
||||
try:
|
||||
rules = manager.list_rules()
|
||||
# Normalize fields for the response model: ensure expected keys exist
|
||||
out = []
|
||||
for r in rules:
|
||||
out.append({
|
||||
"family": r.get("family"),
|
||||
"table": r.get("table"),
|
||||
"chain": r.get("chain"),
|
||||
"handle": r.get("handle"),
|
||||
"expr": r.get("expr")
|
||||
})
|
||||
return out
|
||||
except NftError as e:
|
||||
logger.exception("list_rules failed")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post("/rules", status_code=status.HTTP_201_CREATED)
|
||||
def create_rule(req: CreateRuleRequest):
|
||||
"""
|
||||
Create a constrained rule. Returns 201 on success.
|
||||
Use GET /firewall/rules to obtain the kernel-assigned handle.
|
||||
"""
|
||||
try:
|
||||
manager.add_rule(
|
||||
family=req.family,
|
||||
table=req.table,
|
||||
chain=req.chain,
|
||||
proto=req.proto,
|
||||
src=req.src,
|
||||
dst_port=req.dst_port,
|
||||
verdict=req.verdict,
|
||||
)
|
||||
return {"status": "created"}
|
||||
except (ValueError, NftError) as e:
|
||||
logger.warning("create_rule bad request: %s", e)
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
logger.exception("create_rule internal error")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.delete("/rules/{handle}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def delete_rule(
|
||||
handle: int,
|
||||
family: str = "inet",
|
||||
table: str = "filter",
|
||||
chain: str = "input",
|
||||
):
|
||||
"""
|
||||
Delete a rule by handle. family/table/chain default to common values but must match.
|
||||
"""
|
||||
try:
|
||||
manager.delete_rule(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))
|
||||
@@ -6,6 +6,7 @@ from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
|
||||
from src.api import nft_manager
|
||||
from src.utilities.packet_broadcaster import PacketBroadcaster
|
||||
import src.shared_objects as shared_objects
|
||||
from src.utilities.database import DatabasePool
|
||||
@@ -137,4 +138,5 @@ app.include_router(network_api.router, prefix="/network", tags=["network"])
|
||||
app.include_router(sniffer_api.router, prefix="/sniffer", tags=["sniffer"])
|
||||
app.include_router(packet_api.router, prefix="/packets", tags=["packets"])
|
||||
app.include_router(nftables_api.router, prefix="/nftables", tags=["nftables"])
|
||||
app.include_router(nft_api.router, prefix="/nft", tags=["nft"])
|
||||
app.include_router(nft_api.router, prefix="/nft", tags=["nft"])
|
||||
app.include_router(nft_manager.router, prefix="/firewall", tags=["firewall"])
|
||||
Reference in New Issue
Block a user