mirror of
https://github.com/hansjone/netx.git
synced 2026-10-08 22:20:58 +08:00
add topology tool
This commit is contained in:
parent
99a262e348
commit
298859e608
8 changed files with 216 additions and 22 deletions
|
|
@ -19,9 +19,10 @@ _SQL_FORBIDDEN_RE = re.compile(
|
|||
r"set\s+role|set\s+session|into\s+outfile|pg_read_file|lo_import|lo_export)\b",
|
||||
flags=re.IGNORECASE,
|
||||
)
|
||||
_WITH_RE = re.compile(r"^\s*with\b", flags=re.IGNORECASE)
|
||||
_COMMENT_RE = re.compile(r"/\*.*?\*/|--.*?$", flags=re.IGNORECASE | re.DOTALL | re.MULTILINE)
|
||||
_FROM_JOIN_RE = re.compile(r"\b(?:from|join)\s+([a-zA-Z0-9_\"\.]+)", flags=re.IGNORECASE)
|
||||
# CTE / subquery alias names: "name AS (" defines an inline virtual table, not a DB table.
|
||||
_ALIAS_NAME_RE = re.compile(r"\b(\w+)\s+AS\s*\(", flags=re.IGNORECASE)
|
||||
|
||||
_readonly_engine: Engine | None = None
|
||||
_ReadonlySession: sessionmaker | None = None
|
||||
|
|
@ -43,10 +44,8 @@ def validate_select_sql(sql: str, *, allowed_tables: set[str] | None = None) ->
|
|||
raise HTTPException(status_code=400, detail="sql_required")
|
||||
if ";" in cleaned:
|
||||
raise HTTPException(status_code=400, detail="single_statement_only")
|
||||
if _WITH_RE.search(cleaned):
|
||||
raise HTTPException(status_code=400, detail="with_cte_not_allowed")
|
||||
low = cleaned.lower().lstrip()
|
||||
if not low.startswith("select"):
|
||||
if not low.startswith(("select", "with")):
|
||||
raise HTTPException(status_code=400, detail="select_only")
|
||||
if _SQL_FORBIDDEN_RE.search(cleaned):
|
||||
raise HTTPException(status_code=400, detail="forbidden_keyword")
|
||||
|
|
@ -57,10 +56,13 @@ def validate_select_sql(sql: str, *, allowed_tables: set[str] | None = None) ->
|
|||
refs = _FROM_JOIN_RE.findall(cleaned)
|
||||
if not refs:
|
||||
raise HTTPException(status_code=400, detail="from_required")
|
||||
alias_names = {m.group(1).lower() for m in _ALIAS_NAME_RE.finditer(cleaned)}
|
||||
for ref in refs:
|
||||
normalized = str(ref).strip().strip('"')
|
||||
if "." in normalized:
|
||||
normalized = normalized.split(".")[-1]
|
||||
if normalized.lower() in alias_names:
|
||||
continue # CTE / subquery alias, not a DB table
|
||||
if normalized.lower() not in allowed_tables:
|
||||
raise HTTPException(status_code=400, detail=f"ume_table_not_allowed:{normalized}")
|
||||
return cleaned
|
||||
|
|
@ -109,13 +111,11 @@ def run_select(
|
|||
if statement_timeout_ms > 0:
|
||||
try:
|
||||
if str(getattr(getattr(session, "bind", None), "dialect", None).name).lower().startswith("postgres"):
|
||||
session.execute(
|
||||
sql_text("SET LOCAL statement_timeout = :ms"),
|
||||
{"ms": int(statement_timeout_ms)},
|
||||
)
|
||||
ms_int = int(statement_timeout_ms)
|
||||
session.execute(sql_text(f"SET LOCAL statement_timeout = {ms_int}"))
|
||||
session.execute(sql_text("SET LOCAL search_path TO public"))
|
||||
except Exception:
|
||||
pass
|
||||
session.rollback()
|
||||
res = session.execute(sql_text(wrapped), bind_params)
|
||||
cols = list(res.keys())
|
||||
raw_rows = res.fetchall()
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from .topology_fabric_nodes import (
|
|||
_nodes_by_ids,
|
||||
ensure_fabric_node_for_managed,
|
||||
ensure_fabric_node_for_ume,
|
||||
find_fabric_paths,
|
||||
get_fabric_neighborhood,
|
||||
get_fabric_summary,
|
||||
list_fabric_edges,
|
||||
|
|
@ -44,6 +45,7 @@ __all__ = [
|
|||
"ensure_fabric_node_for_managed",
|
||||
"ensure_fabric_node_for_ume",
|
||||
"ensure_lldp_discovered_managed_ne",
|
||||
"find_fabric_paths",
|
||||
"get_fabric_neighborhood",
|
||||
"get_fabric_summary",
|
||||
"list_fabric_edges",
|
||||
|
|
|
|||
|
|
@ -457,3 +457,112 @@ def get_fabric_neighborhood(
|
|||
)
|
||||
|
||||
|
||||
def find_fabric_paths(
|
||||
db: Session,
|
||||
*,
|
||||
from_ume_ne_id: str = "",
|
||||
from_managed_ne_id: str = "",
|
||||
to_ume_ne_id: str = "",
|
||||
to_managed_ne_id: str = "",
|
||||
max_paths: int = 3,
|
||||
max_hops: int = 6,
|
||||
layer: str = "physical",
|
||||
) -> dict[str, Any]:
|
||||
"""Find up to max_paths simple paths between two fabric nodes.
|
||||
|
||||
Accepts ume_ne_id (from UME alarms) or managed_ne_id (from managed NE) — resolved
|
||||
to fabric_node_id internally so agents can use alarm ne_id directly.
|
||||
"""
|
||||
from_uid = str(from_ume_ne_id or "").strip()
|
||||
from_mid = str(from_managed_ne_id or "").strip()
|
||||
to_uid = str(to_ume_ne_id or "").strip()
|
||||
to_mid = str(to_managed_ne_id or "").strip()
|
||||
if bool(from_uid) == bool(from_mid):
|
||||
raise HTTPException(400, detail="exactly_one_of_from_ume_ne_id_or_from_managed_ne_id_required")
|
||||
if bool(to_uid) == bool(to_mid):
|
||||
raise HTTPException(400, detail="exactly_one_of_to_ume_ne_id_or_to_managed_ne_id_required")
|
||||
|
||||
def _resolve(uid: str, mid: str) -> str:
|
||||
q = db.query(TopoFabricNode)
|
||||
if uid:
|
||||
q = q.filter(TopoFabricNode.ume_ne_id == uid)
|
||||
else:
|
||||
q = q.filter(TopoFabricNode.managed_ne_id == mid)
|
||||
row = q.first()
|
||||
if not row:
|
||||
raise HTTPException(404, detail="fabric_node_not_found_for_ne_id")
|
||||
return row.id
|
||||
|
||||
from_id = _resolve(from_uid, from_mid)
|
||||
to_id = _resolve(to_uid, to_mid)
|
||||
if from_id == to_id:
|
||||
raise HTTPException(400, detail="from_and_to_are_same_node")
|
||||
|
||||
max_paths = max(1, min(10, int(max_paths or 3)))
|
||||
max_hops = max(1, min(12, int(max_hops or 6)))
|
||||
layer_v = str(layer or "physical").strip() or "physical"
|
||||
|
||||
edges = db.query(TopoFabricEdge).filter(TopoFabricEdge.layer == layer_v).all()
|
||||
edge_map: dict[str, TopoFabricEdge] = {e.id: e for e in edges}
|
||||
adj: dict[str, list[tuple[str, str]]] = {}
|
||||
for e in edges:
|
||||
adj.setdefault(e.a_node_id, []).append((e.b_node_id, e.id))
|
||||
adj.setdefault(e.b_node_id, []).append((e.a_node_id, e.id))
|
||||
|
||||
# DFS for simple paths (no repeated nodes), capped to avoid blow-up on large graphs.
|
||||
_EXPLORE_CAP = 100
|
||||
all_paths: list[list[str]] = []
|
||||
stack: list[tuple[str, list[str], set[str]]] = [(from_id, [], {from_id})]
|
||||
while stack and len(all_paths) < _EXPLORE_CAP:
|
||||
node, edge_path, visited = stack.pop()
|
||||
if len(edge_path) >= max_hops:
|
||||
continue
|
||||
for nbr, eid in adj.get(node, []):
|
||||
if nbr in visited:
|
||||
continue
|
||||
new_path = edge_path + [eid]
|
||||
if nbr == to_id:
|
||||
all_paths.append(new_path)
|
||||
continue
|
||||
stack.append((nbr, new_path, visited | {nbr}))
|
||||
|
||||
all_paths.sort(key=len)
|
||||
found = all_paths[:max_paths]
|
||||
|
||||
node_ids = {from_id, to_id}
|
||||
for p in found:
|
||||
for eid in p:
|
||||
e = edge_map.get(eid)
|
||||
if e:
|
||||
node_ids.add(e.a_node_id)
|
||||
node_ids.add(e.b_node_id)
|
||||
node_map = _nodes_by_ids(db, node_ids)
|
||||
|
||||
def _path_nodes(edge_ids: list[str]) -> list[dict]:
|
||||
ids = [from_id]
|
||||
cur = from_id
|
||||
for eid in edge_ids:
|
||||
e = edge_map.get(eid)
|
||||
if not e:
|
||||
break
|
||||
nxt = e.b_node_id if e.a_node_id == cur else e.a_node_id
|
||||
ids.append(nxt)
|
||||
cur = nxt
|
||||
return [_node_out(node_map[nid]).model_dump() for nid in ids if nid in node_map]
|
||||
|
||||
return {
|
||||
"from_node_id": from_id,
|
||||
"to_node_id": to_id,
|
||||
"layer": layer_v,
|
||||
"path_count": len(found),
|
||||
"paths": [
|
||||
{
|
||||
"hops": len(p),
|
||||
"nodes": _path_nodes(p),
|
||||
"edges": [_edge_out(edge_map[eid], nodes_by_id=node_map).model_dump() for eid in p if eid in edge_map],
|
||||
}
|
||||
for p in found
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from __future__ import annotations
|
|||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .db import get_db
|
||||
|
|
@ -49,6 +49,7 @@ from .topology_discover import get_discover_job, start_discover_job
|
|||
from .topology_fabric import (
|
||||
delete_fabric_edge,
|
||||
delete_fabric_edges,
|
||||
find_fabric_paths,
|
||||
get_fabric_neighborhood,
|
||||
get_fabric_summary,
|
||||
list_fabric_edges,
|
||||
|
|
@ -147,6 +148,21 @@ def api_fabric_neighborhood(
|
|||
return get_fabric_neighborhood(db, node_id, depth=depth, layer=layer).model_dump()
|
||||
|
||||
|
||||
@router.post("/fabric/paths")
|
||||
def api_fabric_paths(body: dict[str, Any] = Body(...), db: Session = Depends(get_db)) -> dict[str, Any]:
|
||||
"""Find up to max_paths simple paths between two fabric nodes (by ne_id)."""
|
||||
return find_fabric_paths(
|
||||
db,
|
||||
from_ume_ne_id=str(body.get("from_ume_ne_id") or ""),
|
||||
from_managed_ne_id=str(body.get("from_managed_ne_id") or ""),
|
||||
to_ume_ne_id=str(body.get("to_ume_ne_id") or ""),
|
||||
to_managed_ne_id=str(body.get("to_managed_ne_id") or ""),
|
||||
max_paths=int(body.get("max_paths") or 3),
|
||||
max_hops=int(body.get("max_hops") or 6),
|
||||
layer=str(body.get("layer") or "physical"),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/fabric/edges")
|
||||
def api_fabric_manual_edge(
|
||||
body: FabricManualEdgeIn,
|
||||
|
|
|
|||
|
|
@ -437,9 +437,9 @@ def ume_diagnostics(
|
|||
return {
|
||||
"source": "ume_alarms_current",
|
||||
"total_alarms": len(rows),
|
||||
"severity_summary": [{"key": k, "count": v} for k, v in by_severity],
|
||||
"top_alarm_codes": [{"key": k, "count": v} for k, v in by_alarm_code],
|
||||
"top_ne": [{"key": k, "count": v} for k, v in by_ne],
|
||||
"severity_summary": by_severity,
|
||||
"top_alarm_codes": by_alarm_code,
|
||||
"top_ne": by_ne,
|
||||
"protocol_summary": [{"key": k, "count": v} for k, v in protocol_summary],
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -285,6 +285,31 @@ def _list_cli_targets(args: dict[str, Any]) -> dict[str, Any]:
|
|||
return http_json("GET", "/v1/cli/targets", params=params)
|
||||
|
||||
|
||||
def _find_topology_paths(args: dict[str, Any]) -> dict[str, Any]:
|
||||
from_uid = str(args.get("from_ume_ne_id") or "").strip()
|
||||
from_mid = str(args.get("from_managed_ne_id") or "").strip()
|
||||
to_uid = str(args.get("to_ume_ne_id") or "").strip()
|
||||
to_mid = str(args.get("to_managed_ne_id") or "").strip()
|
||||
if bool(from_uid) == bool(from_mid):
|
||||
return {"ok": False, "error": "exactly_one_of_from_ume_ne_id_or_from_managed_ne_id_required"}
|
||||
if bool(to_uid) == bool(to_mid):
|
||||
return {"ok": False, "error": "exactly_one_of_to_ume_ne_id_or_to_managed_ne_id_required"}
|
||||
body: dict[str, Any] = {
|
||||
"max_paths": max(1, min(10, int(args.get("max_paths") or 3))),
|
||||
"max_hops": max(1, min(12, int(args.get("max_hops") or 6))),
|
||||
"layer": str(args.get("layer") or "physical").strip() or "physical",
|
||||
}
|
||||
if from_uid:
|
||||
body["from_ume_ne_id"] = from_uid
|
||||
else:
|
||||
body["from_managed_ne_id"] = from_mid
|
||||
if to_uid:
|
||||
body["to_ume_ne_id"] = to_uid
|
||||
else:
|
||||
body["to_managed_ne_id"] = to_mid
|
||||
return http_post_json("/v1/topology/fabric/paths", body, timeout=30.0)
|
||||
|
||||
|
||||
HTTP_MCP_TOOLS: list[dict[str, Any]] = [
|
||||
{
|
||||
"name": "queryUmeAlarms",
|
||||
|
|
@ -469,6 +494,28 @@ HTTP_MCP_TOOLS: list[dict[str, Any]] = [
|
|||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "findTopologyPaths",
|
||||
"description": (
|
||||
"Find up to max_paths simple paths between two fabric nodes for troubleshooting. "
|
||||
"Accepts ume_ne_id (from UME alarms) or managed_ne_id — resolved to fabric node "
|
||||
"internally. Returns paths with node sequence + edge status (up/down)."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"from_ume_ne_id": {"type": "string", "description": "Source UME ne_id (from alarm ne_id)"},
|
||||
"from_managed_ne_id": {"type": "string", "description": "Source managed NE id"},
|
||||
"to_ume_ne_id": {"type": "string", "description": "Target UME ne_id"},
|
||||
"to_managed_ne_id": {"type": "string", "description": "Target managed NE id"},
|
||||
"max_paths": {"type": "integer", "minimum": 1, "maximum": 10, "default": 3},
|
||||
"max_hops": {"type": "integer", "minimum": 1, "maximum": 12, "default": 6},
|
||||
"layer": {"type": "string", "default": "physical"},
|
||||
},
|
||||
"required": ["from_ume_ne_id", "to_ume_ne_id"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
_HANDLERS: dict[str, Callable[[dict[str, Any]], dict[str, Any]]] = {
|
||||
|
|
@ -485,6 +532,7 @@ _HANDLERS: dict[str, Callable[[dict[str, Any]], dict[str, Any]]] = {
|
|||
"getManagedNe": _get_managed_ne,
|
||||
"execManagedNe": _exec_managed_ne,
|
||||
"listCliTargets": _list_cli_targets,
|
||||
"findTopologyPaths": _find_topology_paths,
|
||||
}
|
||||
|
||||
# Minimum scope required to advertise / invoke each tool (matches netx API RBAC).
|
||||
|
|
@ -502,6 +550,7 @@ TOOL_REQUIRED_SCOPE: dict[str, str] = {
|
|||
"getManagedNe": "ne:read",
|
||||
"execManagedNe": "ne:exec",
|
||||
"listCliTargets": "ne:read",
|
||||
"findTopologyPaths": "ne:read",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -15,11 +15,12 @@ from netx_mcp.server import _fetch_scopes
|
|||
|
||||
def test_http_mcp_tool_list_has_expected_tools() -> None:
|
||||
names = [str(t.get("name") or "") for t in HTTP_MCP_TOOLS]
|
||||
assert len(names) == 13
|
||||
assert len(names) == 14
|
||||
assert "queryUmeAlarms" in names
|
||||
assert "queryUmeAlarmsRaw" in names
|
||||
assert "execManagedNe" in names
|
||||
assert "listCliTargets" in names
|
||||
assert "findTopologyPaths" in names
|
||||
assert "queryTopologyEdges" not in names
|
||||
exec_tool = next(t for t in HTTP_MCP_TOOLS if t.get("name") == "execManagedNe")
|
||||
assert exec_tool["inputSchema"]["properties"]["commands"]["maxItems"] >= 5
|
||||
|
|
@ -136,7 +137,7 @@ def test_stdio_initialize_and_tools_list() -> None:
|
|||
list_resp = json.loads(list_line)
|
||||
assert "error" not in list_resp, list_resp
|
||||
tools = list_resp["result"]["tools"]
|
||||
assert len(tools) == 13
|
||||
assert len(tools) == 14
|
||||
|
||||
proc.terminate()
|
||||
proc.wait(timeout=5)
|
||||
|
|
|
|||
|
|
@ -54,13 +54,30 @@ class ScopeUnitTests(unittest.TestCase):
|
|||
|
||||
|
||||
class SqlGuardTests(unittest.TestCase):
|
||||
def test_rejects_cte(self) -> None:
|
||||
with self.assertRaises(HTTPException) as ctx:
|
||||
validate_select_sql(
|
||||
"WITH x AS (SELECT * FROM app_user) SELECT * FROM x",
|
||||
allowed_tables={"ume_alarms_current", "ume_inventory_ne"},
|
||||
)
|
||||
self.assertEqual(ctx.exception.detail, "with_cte_not_allowed")
|
||||
def test_allows_cte(self) -> None:
|
||||
# CTE (including WITH RECURSIVE) is now allowed; table whitelist still applies.
|
||||
cleaned = validate_select_sql(
|
||||
"WITH x AS (SELECT * FROM ume_alarms_current) SELECT * FROM x",
|
||||
allowed_tables={"ume_alarms_current", "ume_inventory_ne"},
|
||||
)
|
||||
self.assertTrue(cleaned.lower().startswith("with"))
|
||||
|
||||
def test_allows_cte_referencing_alias_in_from(self) -> None:
|
||||
# CTE name appears in FROM but should not be flagged as unauthorized table.
|
||||
cleaned = validate_select_sql(
|
||||
"WITH stats AS (SELECT host_name, COUNT(*) AS cnt FROM ume_alarms_current GROUP BY host_name) "
|
||||
"SELECT * FROM stats WHERE cnt > 5",
|
||||
allowed_tables={"ume_alarms_current", "ume_inventory_ne"},
|
||||
)
|
||||
self.assertIn("stats", cleaned.lower())
|
||||
|
||||
def test_allows_subquery_alias_in_from(self) -> None:
|
||||
# Subquery alias in FROM should not be flagged as unauthorized table.
|
||||
cleaned = validate_select_sql(
|
||||
"SELECT * FROM (SELECT host_name, COUNT(*) AS cnt FROM ume_alarms_current GROUP BY host_name) AS stats WHERE cnt > 5",
|
||||
allowed_tables={"ume_alarms_current", "ume_inventory_ne"},
|
||||
)
|
||||
self.assertIn("stats", cleaned.lower())
|
||||
|
||||
def test_rejects_disallowed_table(self) -> None:
|
||||
with self.assertRaises(HTTPException) as ctx:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue