fix(managed-ne): require filters for listManagedNe and tighten pagination

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-07-02 11:07:47 +08:00
parent 572b0cfd9d
commit 987b501e53
3 changed files with 51 additions and 11 deletions

View file

@ -236,6 +236,12 @@ def list_managed_ne(
) -> dict[str, Any]: ) -> dict[str, Any]:
stmt = db.query(ManagedNE) stmt = db.query(ManagedNE)
kw = str(keyword or "").strip() 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: if kw:
like = f"%{kw}%" like = f"%{kw}%"
stmt = stmt.filter( stmt = stmt.filter(
@ -248,10 +254,8 @@ def list_managed_ne(
ManagedNE.device_type.ilike(like), ManagedNE.device_type.ilike(like),
) )
) )
v = str(vendor or "").strip()
if v: if v:
stmt = stmt.filter(ManagedNE.vendor == v) stmt = stmt.filter(ManagedNE.vendor == v)
cs = str(connect_status or "").strip()
if cs: if cs:
stmt = stmt.filter(ManagedNE.connect_status == cs) stmt = stmt.filter(ManagedNE.connect_status == cs)
total = int(stmt.count()) total = int(stmt.count())

View file

@ -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]: def _list_managed_ne(args: dict[str, Any]) -> dict[str, Any]:
page = max(1, int(args.get("page") or 1)) 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} params: dict[str, Any] = {"page": page, "page_size": page_size}
if str(args.get("keyword") or "").strip(): if keyword:
params["keyword"] = str(args.get("keyword")).strip() params["keyword"] = keyword
if str(args.get("vendor") or "").strip(): if vendor:
params["vendor"] = str(args.get("vendor")).strip() params["vendor"] = vendor
if str(args.get("connect_status") or "").strip(): if connect_status:
params["connect_status"] = str(args.get("connect_status")).strip() params["connect_status"] = connect_status
return http_json("GET", "/v1/managed-ne", params=params) return http_json("GET", "/v1/managed-ne", params=params)
@ -385,7 +392,7 @@ HTTP_MCP_TOOLS: list[dict[str, Any]] = [
}, },
{ {
"name": "listManagedNe", "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": { "inputSchema": {
"type": "object", "type": "object",
"properties": { "properties": {
@ -393,7 +400,7 @@ HTTP_MCP_TOOLS: list[dict[str, Any]] = [
"vendor": {"type": "string"}, "vendor": {"type": "string"},
"connect_status": {"type": "string", "enum": ["unknown", "testing", "pass", "fail"]}, "connect_status": {"type": "string", "enum": ["unknown", "testing", "pass", "fail"]},
"page": {"type": "integer", "minimum": 1, "default": 1}, "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": [], "required": [],
"additionalProperties": False, "additionalProperties": False,

View file

@ -145,6 +145,35 @@ class ManagedNeApiTests(unittest.TestCase):
self.assertEqual(listed.json()["total"], 1, listed.text) self.assertEqual(listed.json()["total"], 1, listed.text)
self.client.delete(f"/v1/managed-ne/{ne_id}") 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): def test_create_without_crypto_key(self):
settings.credential_secret_key = "" settings.credential_secret_key = ""
r = self.client.post( r = self.client.post(