diff --git a/netx_api/auth_scopes.py b/netx_api/auth_scopes.py index 6210e3f..7c3dc27 100644 --- a/netx_api/auth_scopes.py +++ b/netx_api/auth_scopes.py @@ -114,7 +114,8 @@ def required_scope_for_request(method: str, path: str) -> str | None: return SCOPE_WEBCRT if p.startswith("/v1/managed-ne"): - if p.rstrip("/").endswith("/exec") and m == "POST": + path_tail = p.rstrip("/") + if m == "POST" and (path_tail.endswith("/exec") or path_tail.endswith("/exec-batch")): return SCOPE_NE_EXEC if m in ("POST", "PUT", "PATCH", "DELETE"): return SCOPE_NE_WRITE diff --git a/netx_api/managed_ne_router.py b/netx_api/managed_ne_router.py index dfcae43..be14851 100644 --- a/netx_api/managed_ne_router.py +++ b/netx_api/managed_ne_router.py @@ -8,12 +8,13 @@ from .db import get_db from .device_types import SUPPORTED_VENDORS from .ne_connect import schedule_connect_tests from .ne_crypto import credentials_configured -from .ne_exec import execute_managed_ne_commands +from .ne_exec import execute_managed_ne_commands, execute_managed_ne_commands_batch from .ne_schemas import ( BatchAccountApplyRequest, BatchHopApplyRequest, ConnectTestRequest, ManagedNeCreate, + ManagedNeExecBatchRequest, ManagedNeExecRequest, ManagedNeUpdate, ) @@ -141,6 +142,22 @@ def api_exec_managed_ne(body: ManagedNeExecRequest, db: Session = Depends(get_db ) +@router.post("/exec-batch") +def api_exec_managed_ne_batch(body: ManagedNeExecBatchRequest): + """Run read-only CLI on many NEs concurrently (field multi-NE sweeps).""" + targets = None + if body.targets: + targets = [t.model_dump() for t in body.targets] + return execute_managed_ne_commands_batch( + targets=targets, + ne_ids=body.ne_ids, + ume_ne_ids=body.ume_ne_ids, + commands=body.commands, + read_timeout_sec=body.read_timeout_sec, + concurrency=body.concurrency, + ) + + @router.post("/connect-test") def api_connect_test(body: ConnectTestRequest, db: Session = Depends(get_db)): ids = [str(x).strip() for x in body.ids if str(x).strip()] diff --git a/netx_api/ne_exec.py b/netx_api/ne_exec.py index 97fd106..cde728b 100644 --- a/netx_api/ne_exec.py +++ b/netx_api/ne_exec.py @@ -2,6 +2,7 @@ from __future__ import annotations +from concurrent.futures import ThreadPoolExecutor, as_completed from typing import Any from fastapi import HTTPException @@ -9,6 +10,7 @@ from sqlalchemy.orm import Session from .cli_resolve import resolve_cli_target from .config import settings +from .db import SessionLocal from .ne_collect_runner import _collect_on_device from .ne_crypto import credentials_configured from .ne_exec_guard import _validate_command, validate_ne_exec_command @@ -17,10 +19,14 @@ _EXEC_MAX_COMMANDS_CAP = 50 _EXEC_MAX_OUTPUT = 32_000 _EXEC_READ_TIMEOUT_DEFAULT = 60 _EXEC_READ_TIMEOUT_MAX = 120 +_EXEC_BATCH_MAX_TARGETS = 20 +_EXEC_BATCH_DEFAULT_CONCURRENCY = 4 +_EXEC_BATCH_MAX_CONCURRENCY = 8 __all__ = [ "_validate_command", "execute_managed_ne_commands", + "execute_managed_ne_commands_batch", "validate_ne_exec_command", ] @@ -83,3 +89,124 @@ def execute_managed_ne_commands( "read_timeout_sec": read_timeout, "output": output, } + + +def _normalize_batch_targets( + *, + targets: list[dict[str, Any]] | None, + ne_ids: list[str] | None, + ume_ne_ids: list[str] | None, + shared_commands: list[str] | None, +) -> list[dict[str, Any]]: + out: list[dict[str, Any]] = [] + for raw in targets or []: + if not isinstance(raw, dict): + continue + mid = str(raw.get("ne_id") or "").strip() + uid = str(raw.get("ume_ne_id") or "").strip() + cmds_raw = raw.get("commands") + cmds = ( + [str(c).strip() for c in cmds_raw if str(c).strip()] + if isinstance(cmds_raw, list) + else list(shared_commands or []) + ) + out.append({"ne_id": mid or None, "ume_ne_id": uid or None, "commands": cmds}) + for mid in ne_ids or []: + s = str(mid or "").strip() + if s: + out.append({"ne_id": s, "ume_ne_id": None, "commands": list(shared_commands or [])}) + for uid in ume_ne_ids or []: + s = str(uid or "").strip() + if s: + out.append({"ne_id": None, "ume_ne_id": s, "commands": list(shared_commands or [])}) + if not out: + raise HTTPException(status_code=400, detail="targets_required") + if len(out) > _EXEC_BATCH_MAX_TARGETS: + raise HTTPException( + status_code=400, + detail=f"too_many_targets (max {_EXEC_BATCH_MAX_TARGETS})", + ) + return out + + +def execute_managed_ne_commands_batch( + *, + targets: list[dict[str, Any]] | None = None, + ne_ids: list[str] | None = None, + ume_ne_ids: list[str] | None = None, + commands: list[str] | None = None, + read_timeout_sec: int | None = None, + concurrency: int | None = None, +) -> dict[str, Any]: + """Fan out read-only CLI across many NEs with bounded concurrency. + + Each worker opens its own DB session. Partial failures stay in ``results``; + top-level ``ok`` is true when the batch request itself is valid. + """ + shared = [str(c).strip() for c in (commands or []) if str(c).strip()] + normalized = _normalize_batch_targets( + targets=targets, + ne_ids=ne_ids, + ume_ne_ids=ume_ne_ids, + shared_commands=shared, + ) + workers = int(concurrency if concurrency is not None else _EXEC_BATCH_DEFAULT_CONCURRENCY) + workers = max(1, min(_EXEC_BATCH_MAX_CONCURRENCY, workers, len(normalized))) + + def _run_one(idx: int, item: dict[str, Any]) -> tuple[int, dict[str, Any]]: + db = SessionLocal() + try: + try: + row = execute_managed_ne_commands( + db, + list(item.get("commands") or []), + ne_id=item.get("ne_id"), + ume_ne_id=item.get("ume_ne_id"), + read_timeout_sec=read_timeout_sec, + ) + except HTTPException as exc: + row = { + "ok": False, + "ne_id": item.get("ne_id"), + "ume_ne_id": item.get("ume_ne_id"), + "commands": list(item.get("commands") or []), + "error": str(exc.detail), + "http_status": int(exc.status_code), + } + except Exception as exc: + row = { + "ok": False, + "ne_id": item.get("ne_id"), + "ume_ne_id": item.get("ume_ne_id"), + "commands": list(item.get("commands") or []), + "error": type(exc).__name__, + "detail": str(exc)[:2000], + } + if isinstance(row, dict): + row.setdefault("ne_id", item.get("ne_id")) + row.setdefault("ume_ne_id", item.get("ume_ne_id")) + row["target_index"] = idx + return idx, row if isinstance(row, dict) else {"ok": False, "error": "invalid_result", "target_index": idx} + finally: + db.close() + + ordered: list[dict[str, Any] | None] = [None] * len(normalized) + with ThreadPoolExecutor(max_workers=workers) as ex: + futs = [ex.submit(_run_one, i, item) for i, item in enumerate(normalized)] + for fut in as_completed(futs): + idx, row = fut.result() + ordered[idx] = row + + results = [r if isinstance(r, dict) else {"ok": False, "error": "missing_result"} for r in ordered] + ok_n = sum(1 for r in results if r.get("ok") is True) + fail_n = len(results) - ok_n + return { + "ok": True, + "concurrency": workers, + "summary": {"total": len(results), "ok": ok_n, "failed": fail_n}, + "results": results, + "hint": ( + "Multi-NE CLI finished in one batch. Summarize ok/failed counts; " + "do not re-loop execManagedNe per NE for the same commands." + ), + } diff --git a/netx_api/ne_schemas.py b/netx_api/ne_schemas.py index 2ee5800..84ead33 100644 --- a/netx_api/ne_schemas.py +++ b/netx_api/ne_schemas.py @@ -127,6 +127,25 @@ class ManagedNeExecRequest(BaseModel): read_timeout_sec: int | None = Field(default=None, ge=10, le=120) +class ManagedNeExecBatchTarget(BaseModel): + """One NE in a concurrent exec batch (exactly one of ne_id / ume_ne_id).""" + + ne_id: str | None = None + ume_ne_id: str | None = None + commands: list[str] | None = Field(default=None, max_length=50) + + +class ManagedNeExecBatchRequest(BaseModel): + """Run the same (or per-target) read-only CLI on many NEs concurrently.""" + + targets: list[ManagedNeExecBatchTarget] | None = Field(default=None, max_length=20) + ne_ids: list[str] | None = Field(default=None, max_length=20) + ume_ne_ids: list[str] | None = Field(default=None, max_length=20) + commands: list[str] | None = Field(default=None, max_length=50) + read_timeout_sec: int | None = Field(default=None, ge=10, le=120) + concurrency: int | None = Field(default=4, ge=1, le=8) + + class HopProxyConfig(BaseModel): """Shared jump-host (proxy) settings applied to one or many NEs.""" diff --git a/packages/netx-mcp/src/netx_mcp/http_tools.py b/packages/netx-mcp/src/netx_mcp/http_tools.py index 2917b5f..47d539c 100644 --- a/packages/netx-mcp/src/netx_mcp/http_tools.py +++ b/packages/netx-mcp/src/netx_mcp/http_tools.py @@ -282,6 +282,60 @@ def _get_managed_ne(args: dict[str, Any]) -> dict[str, Any]: def _exec_managed_ne(args: dict[str, Any]) -> dict[str, Any]: + targets_raw = args.get("targets") + ne_ids_raw = args.get("ne_ids") + ume_ne_ids_raw = args.get("ume_ne_ids") + shared_cmds_raw = args.get("commands") + shared_commands = ( + [str(c).strip() for c in shared_cmds_raw if str(c).strip()] + if isinstance(shared_cmds_raw, list) + else [] + ) + multi = bool( + (isinstance(targets_raw, list) and targets_raw) + or (isinstance(ne_ids_raw, list) and ne_ids_raw) + or (isinstance(ume_ne_ids_raw, list) and ume_ne_ids_raw) + ) + if multi: + body: dict[str, Any] = {} + if isinstance(targets_raw, list) and targets_raw: + cleaned_targets: list[dict[str, Any]] = [] + for t in targets_raw: + if not isinstance(t, dict): + continue + row: dict[str, Any] = {} + if str(t.get("ne_id") or "").strip(): + row["ne_id"] = str(t.get("ne_id")).strip() + if str(t.get("ume_ne_id") or "").strip(): + row["ume_ne_id"] = str(t.get("ume_ne_id")).strip() + cmds = t.get("commands") + if isinstance(cmds, list) and cmds: + row["commands"] = [str(c).strip() for c in cmds if str(c).strip()] + if row: + cleaned_targets.append(row) + body["targets"] = cleaned_targets + if isinstance(ne_ids_raw, list) and ne_ids_raw: + body["ne_ids"] = [str(x).strip() for x in ne_ids_raw if str(x).strip()] + if isinstance(ume_ne_ids_raw, list) and ume_ne_ids_raw: + body["ume_ne_ids"] = [str(x).strip() for x in ume_ne_ids_raw if str(x).strip()] + if shared_commands: + if len(shared_commands) > exec_max_commands(): + return {"ok": False, "error": "too_many_commands", "error_code": "too_many_commands"} + body["commands"] = shared_commands + rts = args.get("read_timeout_sec") + body["read_timeout_sec"] = int(rts) if rts is not None else 60 + conc = args.get("concurrency") + if conc is not None: + body["concurrency"] = int(conc) + # Wall clock: many NEs × per-cmd timeout; keep below oclaw MCP override. + out = http_post_json("/v1/managed-ne/exec-batch", body, timeout=600.0) + if not out.get("ok"): + return out + data = out.get("data") or {} + if isinstance(data, dict) and data.get("ok") is False: + return {"ok": False, "data": data, "error": str(data.get("error") or "exec_batch_failed")} + return {"ok": True, "data": data} + ne_id = str(args.get("ne_id") or "").strip() ume_ne_id = str(args.get("ume_ne_id") or "").strip() if bool(ne_id) == bool(ume_ne_id): @@ -289,16 +343,16 @@ def _exec_managed_ne(args: dict[str, Any]) -> dict[str, Any]: "ok": False, "error": "exactly_one_of_ne_id_or_ume_ne_id_required", "error_code": "exactly_one_of_ne_id_or_ume_ne_id_required", + "hint": ( + "For one NE pass ne_id OR ume_ne_id. For many NEs pass ne_ids / ume_ne_ids " + "(or targets[]) with shared commands — one call, concurrent on server." + ), } - raw_cmds = args.get("commands") - if not isinstance(raw_cmds, list) or not raw_cmds: + if not shared_commands: return {"ok": False, "error": "commands_required", "error_code": "commands_required"} - commands = [str(c).strip() for c in raw_cmds if str(c).strip()] - if not commands: - return {"ok": False, "error": "commands_required", "error_code": "commands_required"} - if len(commands) > exec_max_commands(): + if len(shared_commands) > exec_max_commands(): return {"ok": False, "error": "too_many_commands", "error_code": "too_many_commands"} - body: dict[str, Any] = {"commands": commands} + body = {"commands": shared_commands} if ne_id: body["ne_id"] = ne_id if ume_ne_id: @@ -592,21 +646,55 @@ HTTP_MCP_TOOLS: list[dict[str, Any]] = [ "name": "execManagedNe", "description": ( f"Run read-only CLI via netx (show/display/ping/traceroute; " - f"max {exec_max_commands()} commands per call, NETX_NE_EXEC_MAX_COMMANDS). " - "Use ne_id (managed NE) OR ume_ne_id (UME inventory). " - "Batch multiple show commands in one call instead of looping. " - "Default read_timeout_sec=60; on timeout raise to 90–120 or shrink commands — do not blind-retry." + f"max {exec_max_commands()} commands per NE, NETX_NE_EXEC_MAX_COMMANDS). " + "Single NE: ne_id OR ume_ne_id + commands. " + "Many NEs (preferred for sweeps): ne_ids[] or ume_ne_ids[] or targets[] with shared commands " + "— server runs them concurrently (default concurrency 4, max 20 targets). " + "Do NOT loop one-NE execManagedNe for the same show commands. " + "Default read_timeout_sec=60; on timeout raise to 90–120 — do not blind-retry." ), "inputSchema": { "type": "object", "properties": { "ne_id": {"type": "string"}, "ume_ne_id": {"type": "string"}, + "ne_ids": { + "type": "array", + "items": {"type": "string"}, + "maxItems": 20, + "description": "Managed NE ids for concurrent batch (shared commands).", + }, + "ume_ne_ids": { + "type": "array", + "items": {"type": "string"}, + "maxItems": 20, + "description": "UME inventory ne_ids for concurrent batch (shared commands).", + }, + "targets": { + "type": "array", + "maxItems": 20, + "items": { + "type": "object", + "properties": { + "ne_id": {"type": "string"}, + "ume_ne_id": {"type": "string"}, + "commands": { + "type": "array", + "items": {"type": "string"}, + "minItems": 1, + "maxItems": exec_max_commands(), + }, + }, + "additionalProperties": False, + }, + "description": "Explicit per-NE targets; optional per-target commands override shared commands.", + }, "commands": { "type": "array", "items": {"type": "string"}, "minItems": 1, "maxItems": exec_max_commands(), + "description": "Commands for single NE, or shared commands for ne_ids/ume_ne_ids/targets.", }, "read_timeout_sec": { "type": "integer", @@ -615,8 +703,15 @@ HTTP_MCP_TOOLS: list[dict[str, Any]] = [ "default": 60, "description": "Per-command read timeout (default 60; use 90–120 for slow show).", }, + "concurrency": { + "type": "integer", + "minimum": 1, + "maximum": 8, + "default": 4, + "description": "Parallel NEs for batch mode (ignored for single-NE).", + }, }, - "required": ["commands"], + "required": [], "additionalProperties": False, }, }, diff --git a/packages/netx-mcp/tests/test_mcp_http.py b/packages/netx-mcp/tests/test_mcp_http.py index 182b371..5835d50 100644 --- a/packages/netx-mcp/tests/test_mcp_http.py +++ b/packages/netx-mcp/tests/test_mcp_http.py @@ -156,6 +156,41 @@ def test_call_exec_managed_ne_posts_body() -> None: assert payload["ok"] is True +def test_call_exec_managed_ne_batch_ne_ids() -> None: + with patch("netx_mcp.http_tools.http_post_json") as mock_post: + mock_post.return_value = { + "ok": True, + "data": { + "ok": True, + "summary": {"total": 2, "ok": 2, "failed": 0}, + "results": [{"ok": True}, {"ok": True}], + }, + } + out = call_http_tool( + "execManagedNe", + {"ne_ids": ["a", "b"], "commands": ["show version"], "concurrency": 2}, + ) + mock_post.assert_called_once() + assert mock_post.call_args[0][0] == "/v1/managed-ne/exec-batch" + body = mock_post.call_args[0][1] + assert body["ne_ids"] == ["a", "b"] + assert body["commands"] == ["show version"] + assert body["concurrency"] == 2 + payload = json.loads(out["content"][0]["text"]) + assert payload["ok"] is True + assert payload["data"]["summary"]["total"] == 2 + + +def test_exec_managed_ne_schema_documents_batch() -> None: + tool = next(t for t in HTTP_MCP_TOOLS if t.get("name") == "execManagedNe") + props = tool["inputSchema"]["properties"] + assert "ne_ids" in props + assert "ume_ne_ids" in props + assert "targets" in props + assert "concurrency" in props + assert "Many NEs" in tool["description"] or "ne_ids" in tool["description"] + + def test_call_exec_managed_ne_respects_max_commands_env(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("NETX_NE_EXEC_MAX_COMMANDS", "10") cmds = [f"show version {i}" for i in range(6)] diff --git a/tests/test_ne_exec.py b/tests/test_ne_exec.py index 650b111..ebcef1c 100644 --- a/tests/test_ne_exec.py +++ b/tests/test_ne_exec.py @@ -205,5 +205,51 @@ class NeExecRunTests(unittest.TestCase): collect.assert_called_once() +class NeExecBatchTests(unittest.TestCase): + @patch("netx_api.ne_exec.credentials_configured", return_value=True) + @patch("netx_api.ne_exec._collect_on_device", return_value="ok-output") + @patch("netx_api.ne_exec.resolve_cli_target") + @patch("netx_api.ne_exec.SessionLocal") + def test_batch_ne_ids_runs_concurrently( + self, + session_local: MagicMock, + resolve: MagicMock, + collect: MagicMock, + _creds: MagicMock, + ) -> None: + from netx_api.ne_exec import execute_managed_ne_commands_batch + + session_local.return_value = MagicMock() + resolve.return_value = ( + {"ip_address": "1.1.1.1"}, + { + "source": "managed", + "id": "ne-1", + "name": "R1", + "ip_address": "1.1.1.1", + }, + ) + with patch("netx_api.ne_exec.settings") as mock_settings: + mock_settings.ne_exec_max_commands = 5 + out = execute_managed_ne_commands_batch( + ne_ids=["a", "b", "c"], + commands=["show version"], + concurrency=3, + ) + self.assertTrue(out["ok"]) + self.assertEqual(out["summary"]["total"], 3) + self.assertEqual(out["summary"]["ok"], 3) + self.assertEqual(out["summary"]["failed"], 0) + self.assertEqual(collect.call_count, 3) + self.assertEqual(len(out["results"]), 3) + + def test_batch_requires_targets(self) -> None: + from netx_api.ne_exec import execute_managed_ne_commands_batch + + with self.assertRaises(HTTPException) as ctx: + execute_managed_ne_commands_batch(commands=["show version"]) + self.assertEqual(ctx.exception.detail, "targets_required") + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_rbac_scopes.py b/tests/test_rbac_scopes.py index 7257923..f89cc8c 100644 --- a/tests/test_rbac_scopes.py +++ b/tests/test_rbac_scopes.py @@ -50,6 +50,7 @@ class ScopeUnitTests(unittest.TestCase): self.assertEqual(required_scope_for_request("POST", "/v1/sql/ume_query"), SCOPE_SQL) self.assertEqual(required_scope_for_request("GET", "/v1/webcrt/sessions"), SCOPE_WEBCRT) self.assertEqual(required_scope_for_request("POST", "/v1/managed-ne/exec"), "ne:exec") + self.assertEqual(required_scope_for_request("POST", "/v1/managed-ne/exec-batch"), "ne:exec") self.assertEqual(required_scope_for_request("POST", "/v1/managed-ne"), "ne:write") self.assertEqual(required_scope_for_request("POST", "/v1/topology/fabric/paths"), "ne:read") self.assertEqual(required_scope_for_request("POST", "/v1/topology/fabric/edges"), "ne:write")