Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
84 changes: 84 additions & 0 deletions tests/unit/test_zpa_app_segment_server_group_ids.py
Original file line number Diff line number Diff line change
@@ -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"] == []
10 changes: 6 additions & 4 deletions zscaler/zpa/app_segments_ba.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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")
Expand Down
10 changes: 6 additions & 4 deletions zscaler/zpa/app_segments_ba_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand Down
10 changes: 6 additions & 4 deletions zscaler/zpa/app_segments_inspection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand Down
10 changes: 6 additions & 4 deletions zscaler/zpa/app_segments_pra.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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:
Expand Down
15 changes: 9 additions & 6 deletions zscaler/zpa/application_segment.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand Down