diff --git a/netx_api/sql_guard.py b/netx_api/sql_guard.py index 13e193b..3da0cbd 100644 --- a/netx_api/sql_guard.py +++ b/netx_api/sql_guard.py @@ -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() diff --git a/netx_api/topology_fabric.py b/netx_api/topology_fabric.py index bd9e27d..908df60 100644 --- a/netx_api/topology_fabric.py +++ b/netx_api/topology_fabric.py @@ -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", diff --git a/netx_api/topology_fabric_nodes.py b/netx_api/topology_fabric_nodes.py index 8b2bd86..e3a7797 100644 --- a/netx_api/topology_fabric_nodes.py +++ b/netx_api/topology_fabric_nodes.py @@ -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 + ], + } + + diff --git a/netx_api/topology_router.py b/netx_api/topology_router.py index 097afd9..d996f83 100644 --- a/netx_api/topology_router.py +++ b/netx_api/topology_router.py @@ -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, diff --git a/netx_api/ume_alarms_router.py b/netx_api/ume_alarms_router.py index 28583eb..702fa07 100644 --- a/netx_api/ume_alarms_router.py +++ b/netx_api/ume_alarms_router.py @@ -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], } diff --git a/packages/netx-mcp/src/netx_mcp/http_tools.py b/packages/netx-mcp/src/netx_mcp/http_tools.py index dd9d457..bc8986b 100644 --- a/packages/netx-mcp/src/netx_mcp/http_tools.py +++ b/packages/netx-mcp/src/netx_mcp/http_tools.py @@ -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", } diff --git a/packages/netx-mcp/tests/test_mcp_http.py b/packages/netx-mcp/tests/test_mcp_http.py index 4d26df8..6a881f3 100644 --- a/packages/netx-mcp/tests/test_mcp_http.py +++ b/packages/netx-mcp/tests/test_mcp_http.py @@ -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) diff --git a/tests/test_rbac_scopes.py b/tests/test_rbac_scopes.py index 5d44f20..6e38cd2 100644 --- a/tests/test_rbac_scopes.py +++ b/tests/test_rbac_scopes.py @@ -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: