diff --git a/netx_api/ne_service.py b/netx_api/ne_service.py index 5fe8287..3ba404c 100644 --- a/netx_api/ne_service.py +++ b/netx_api/ne_service.py @@ -236,6 +236,12 @@ def list_managed_ne( ) -> dict[str, Any]: stmt = db.query(ManagedNE) kw = str(keyword or "").strip() + v = str(vendor or "").strip() + cs = str(connect_status or "").strip() + if not (kw or v or cs): + raise HTTPException(status_code=400, detail="managed_ne_filter_required") + if kw and len(kw) < 2: + raise HTTPException(status_code=400, detail="managed_ne_keyword_too_short") if kw: like = f"%{kw}%" stmt = stmt.filter( @@ -248,10 +254,8 @@ def list_managed_ne( ManagedNE.device_type.ilike(like), ) ) - v = str(vendor or "").strip() if v: stmt = stmt.filter(ManagedNE.vendor == v) - cs = str(connect_status or "").strip() if cs: stmt = stmt.filter(ManagedNE.connect_status == cs) total = int(stmt.count()) diff --git a/packages/netx-mcp/src/netx_mcp/http_tools.py b/packages/netx-mcp/src/netx_mcp/http_tools.py index d2578da..e610c81 100644 --- a/packages/netx-mcp/src/netx_mcp/http_tools.py +++ b/packages/netx-mcp/src/netx_mcp/http_tools.py @@ -201,14 +201,21 @@ def _sql_query_ume(args: dict[str, Any]) -> dict[str, Any]: def _list_managed_ne(args: dict[str, Any]) -> dict[str, Any]: page = max(1, int(args.get("page") or 1)) - page_size = min(500, max(1, int(args.get("page_size") or 50))) + page_size = min(100, max(1, int(args.get("page_size") or 20))) + keyword = str(args.get("keyword") or "").strip() + vendor = str(args.get("vendor") or "").strip() + connect_status = str(args.get("connect_status") or "").strip() + if not (keyword or vendor or connect_status): + return {"ok": False, "error": "managed_ne_filter_required", "error_code": "managed_ne_filter_required"} + if keyword and len(keyword) < 2: + return {"ok": False, "error": "managed_ne_keyword_too_short", "error_code": "managed_ne_keyword_too_short"} params: dict[str, Any] = {"page": page, "page_size": page_size} - if str(args.get("keyword") or "").strip(): - params["keyword"] = str(args.get("keyword")).strip() - if str(args.get("vendor") or "").strip(): - params["vendor"] = str(args.get("vendor")).strip() - if str(args.get("connect_status") or "").strip(): - params["connect_status"] = str(args.get("connect_status")).strip() + if keyword: + params["keyword"] = keyword + if vendor: + params["vendor"] = vendor + if connect_status: + params["connect_status"] = connect_status return http_json("GET", "/v1/managed-ne", params=params) @@ -385,7 +392,7 @@ HTTP_MCP_TOOLS: list[dict[str, Any]] = [ }, { "name": "listManagedNe", - "description": "List netx managed NEs (SSH/Telnet inventory); use before execManagedNe.", + "description": "List filtered netx managed NEs (keyword/vendor/connect_status required); use before execManagedNe.", "inputSchema": { "type": "object", "properties": { @@ -393,7 +400,7 @@ HTTP_MCP_TOOLS: list[dict[str, Any]] = [ "vendor": {"type": "string"}, "connect_status": {"type": "string", "enum": ["unknown", "testing", "pass", "fail"]}, "page": {"type": "integer", "minimum": 1, "default": 1}, - "page_size": {"type": "integer", "minimum": 1, "maximum": 500, "default": 50}, + "page_size": {"type": "integer", "minimum": 1, "maximum": 100, "default": 20}, }, "required": [], "additionalProperties": False, diff --git a/tests/test_managed_ne.py b/tests/test_managed_ne.py index 55bfcc0..d12f51c 100644 --- a/tests/test_managed_ne.py +++ b/tests/test_managed_ne.py @@ -145,6 +145,35 @@ class ManagedNeApiTests(unittest.TestCase): self.assertEqual(listed.json()["total"], 1, listed.text) self.client.delete(f"/v1/managed-ne/{ne_id}") + def test_list_requires_filter(self): + listed = self.client.get("/v1/managed-ne") + self.assertEqual(listed.status_code, 400) + self.assertEqual(listed.json()["detail"], "managed_ne_filter_required") + + def test_list_rejects_short_keyword(self): + listed = self.client.get("/v1/managed-ne", params={"keyword": "r"}) + self.assertEqual(listed.status_code, 400) + self.assertEqual(listed.json()["detail"], "managed_ne_keyword_too_short") + + def test_list_allows_vendor_only_filter(self): + created = self.client.post( + "/v1/managed-ne", + json={ + "name": "Agg-R2", + "vendor": "Cisco", + "device_type": "cisco_ios", + "ip_address": "192.168.0.12", + "username": "admin", + "password": "pass123", + }, + ) + self.assertEqual(created.status_code, 200, created.text) + ne_id = created.json()["id"] + listed = self.client.get("/v1/managed-ne", params={"vendor": "Cisco"}) + self.assertEqual(listed.status_code, 200) + self.assertEqual(listed.json()["total"], 1, listed.text) + self.client.delete(f"/v1/managed-ne/{ne_id}") + def test_create_without_crypto_key(self): settings.credential_secret_key = "" r = self.client.post(