diff --git a/tests/unit/test_zpa_app_segment_server_group_ids.py b/tests/unit/test_zpa_app_segment_server_group_ids.py new file mode 100644 index 00000000..be520330 --- /dev/null +++ b/tests/unit/test_zpa_app_segment_server_group_ids.py @@ -0,0 +1,84 @@ +""" +Unit tests for optional ``server_group_ids`` handling in the ZPA application +segment resources. + +Passing ``server_group_ids=None`` -- the natural result of forwarding an +optional keyword argument whose default is ``None``, i.e. a partial update that +does not intend to change server groups -- used to raise +``TypeError: 'NoneType' object is not iterable`` inside the SDK before any +request was built. ``None`` must mean "not supplied", while an explicitly +supplied list (including an empty one) must still be sent as ``serverGroups``. +""" + +import pytest + +from zscaler.zpa.app_segments_ba import ApplicationSegmentBAAPI +from zscaler.zpa.app_segments_ba_v2 import AppSegmentsBAV2API +from zscaler.zpa.app_segments_inspection import AppSegmentsInspectionAPI +from zscaler.zpa.app_segments_pra import AppSegmentsPRAAPI +from zscaler.zpa.application_segment import ApplicationSegmentAPI + +CONFIG = {"client": {"customerId": "1234567890"}} + +SEGMENT_ID = "72058304855089379" +SERVER_GROUP_ID = "72058304855090128" + + +class CapturingExecutor: + """Captures the body the SDK built, then short-circuits before any HTTP.""" + + def __init__(self): + self.body = None + + def create_request(self, method, endpoint, body=None, headers=None, params=None, **kwargs): + self.body = body + return None, "short-circuit" + + +# Every add/update entry point that reformats ``server_group_ids`` into ``serverGroups``. +SEGMENT_METHODS = [ + (ApplicationSegmentAPI, "add_segment", ()), + (ApplicationSegmentAPI, "update_segment", (SEGMENT_ID,)), + (ApplicationSegmentAPI, "add_segment_provision", ()), + (AppSegmentsBAV2API, "add_segment_ba", ()), + (AppSegmentsBAV2API, "update_segment_ba", (SEGMENT_ID,)), + (ApplicationSegmentBAAPI, "add_segment_ba", ()), + (ApplicationSegmentBAAPI, "update_segment_ba", (SEGMENT_ID,)), + (AppSegmentsInspectionAPI, "add_segment_inspection", ()), + (AppSegmentsInspectionAPI, "update_segment_inspection", (SEGMENT_ID,)), + (AppSegmentsPRAAPI, "add_segment_pra", ()), + (AppSegmentsPRAAPI, "update_segment_pra", (SEGMENT_ID,)), +] + +METHOD_IDS = [f"{api_cls.__name__}.{method}" for api_cls, method, _ in SEGMENT_METHODS] + + +def build_body(api_cls, method, args, **kwargs): + """Invoke a segment method and return the request body it assembled.""" + executor = CapturingExecutor() + api = api_cls(executor, CONFIG) + getattr(api, method)(*args, **kwargs) + return executor.body + + +@pytest.mark.parametrize(("api_cls", "method", "args"), SEGMENT_METHODS, ids=METHOD_IDS) +def test_none_server_group_ids_is_treated_as_not_supplied(api_cls, method, args): + body = build_body(api_cls, method, args, name="app.example.com", server_group_ids=None) + + assert "serverGroups" not in body + assert "server_group_ids" not in body + + +@pytest.mark.parametrize(("api_cls", "method", "args"), SEGMENT_METHODS, ids=METHOD_IDS) +def test_server_group_ids_list_is_reformatted(api_cls, method, args): + body = build_body(api_cls, method, args, name="app.example.com", server_group_ids=[SERVER_GROUP_ID]) + + assert body["serverGroups"] == [{"id": SERVER_GROUP_ID}] + assert "server_group_ids" not in body + + +@pytest.mark.parametrize(("api_cls", "method", "args"), SEGMENT_METHODS, ids=METHOD_IDS) +def test_empty_server_group_ids_is_still_sent(api_cls, method, args): + body = build_body(api_cls, method, args, name="app.example.com", server_group_ids=[]) + + assert body["serverGroups"] == [] diff --git a/zscaler/zpa/app_segments_ba.py b/zscaler/zpa/app_segments_ba.py index 81ab53ee..201dc63c 100644 --- a/zscaler/zpa/app_segments_ba.py +++ b/zscaler/zpa/app_segments_ba.py @@ -272,8 +272,9 @@ def add_segment_ba(self, **kwargs) -> APIResult[dict]: microtenant_id = kwargs.get("microtenant_id") or body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] # --- Prevent mixed legacy + structured port range usage --- if "tcp_port_ranges" in body and "tcp_port_range" in body: @@ -409,8 +410,9 @@ def update_segment_ba(self, segment_id: str, **kwargs) -> APIResult[dict]: microtenant_id = body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] if "clientless_app_ids" in body: clientless_apps = body.pop("clientless_app_ids") diff --git a/zscaler/zpa/app_segments_ba_v2.py b/zscaler/zpa/app_segments_ba_v2.py index c5c76da6..66eaa017 100644 --- a/zscaler/zpa/app_segments_ba_v2.py +++ b/zscaler/zpa/app_segments_ba_v2.py @@ -254,8 +254,9 @@ def add_segment_ba(self, **kwargs) -> APIResult[dict]: params = {"microtenantId": microtenant_id} if microtenant_id else {} # Reformat server_group_ids to match the expected API format (serverGroups) - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] # Auto-add `"app_types": ["BROWSER_ACCESS"]` if missing common_apps_dto = kwargs.get("common_apps_dto") @@ -357,8 +358,9 @@ def update_segment_ba(self, segment_id: str, **kwargs) -> APIResult[dict]: microtenant_id = body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": gid} for gid in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": gid} for gid in server_group_ids] # Auto-add `"app_types": ["BROWSER_ACCESS"]` if missing common_apps_dto = kwargs.get("common_apps_dto") diff --git a/zscaler/zpa/app_segments_inspection.py b/zscaler/zpa/app_segments_inspection.py index 2a4b0ae6..87d034e8 100644 --- a/zscaler/zpa/app_segments_inspection.py +++ b/zscaler/zpa/app_segments_inspection.py @@ -241,8 +241,9 @@ def add_segment_inspection(self, **kwargs) -> APIResult[dict]: return None, None, ValueError("Cannot use both 'udp_port_ranges' and 'udp_port_range' in the same request.") # Reformat server_group_ids to match the expected API format (serverGroups) - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] # Auto-add `"app_types": ["INSPECT"]` if missing common_apps_dto = kwargs.get("common_apps_dto") @@ -393,8 +394,9 @@ def update_segment_inspection(self, segment_id: str, **kwargs) -> APIResult[dict microtenant_id = body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": gid} for gid in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": gid} for gid in server_group_ids] # Auto-add `"app_types": ["INSPECT"]` if missing common_apps_dto = kwargs.get("common_apps_dto") diff --git a/zscaler/zpa/app_segments_pra.py b/zscaler/zpa/app_segments_pra.py index 7ba845e2..5b9f8d1e 100644 --- a/zscaler/zpa/app_segments_pra.py +++ b/zscaler/zpa/app_segments_pra.py @@ -249,8 +249,9 @@ def add_segment_pra(self, **kwargs) -> APIResult[dict]: params = {"microtenantId": microtenant_id} if microtenant_id else {} # Reformat server_group_ids to match the expected API format (serverGroups) - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] # Auto-add `"app_types": ["SECURE_REMOTE_ACCESS"]` if missing common_apps_dto = kwargs.get("common_apps_dto") @@ -353,8 +354,9 @@ def update_segment_pra(self, segment_id: str, **kwargs) -> APIResult[dict]: microtenant_id = body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": gid} for gid in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": gid} for gid in server_group_ids] common_apps_dto = kwargs.get("common_apps_dto") if common_apps_dto and "apps_config" in common_apps_dto: diff --git a/zscaler/zpa/application_segment.py b/zscaler/zpa/application_segment.py index 416f4bcf..ca692b48 100644 --- a/zscaler/zpa/application_segment.py +++ b/zscaler/zpa/application_segment.py @@ -270,8 +270,9 @@ def add_segment(self, **kwargs) -> APIResult[ApplicationSegments]: microtenant_id = kwargs.get("microtenant_id") or body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] # --- Prevent mixed legacy + structured port range usage --- if "tcp_port_ranges" in body and "tcp_port_range" in body: @@ -405,8 +406,9 @@ def update_segment(self, segment_id: str, **kwargs) -> APIResult[ApplicationSegm microtenant_id = body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] if "clientless_app_ids" in body: clientless_apps = body.pop("clientless_app_ids") @@ -766,8 +768,9 @@ def add_segment_provision(self, **kwargs) -> APIResult[dict]: microtenant_id = kwargs.get("microtenant_id") or body.get("microtenant_id", None) params = {"microtenantId": microtenant_id} if microtenant_id else {} - if "server_group_ids" in body: - body["serverGroups"] = [{"id": group_id} for group_id in body.pop("server_group_ids")] + server_group_ids = body.pop("server_group_ids", None) + if server_group_ids is not None: + body["serverGroups"] = [{"id": group_id} for group_id in server_group_ids] if "tcp_port_ranges" in body: body["tcpPortRanges"] = body.pop("tcp_port_ranges")