add topology tool

This commit is contained in:
oliver 2026-08-10 20:29:05 +08:00
parent 99a262e348
commit 298859e608
8 changed files with 216 additions and 22 deletions

View file

@ -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()

View file

@ -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",

View file

@ -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
],
}

View file

@ -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,

View file

@ -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],
}

View file

@ -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",
}

View file

@ -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)

View file

@ -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: