From 88dba5415c93475205ead881594afe34be210fe7 Mon Sep 17 00:00:00 2001 From: Ashwin Upadhyay Date: Sun, 17 May 2026 00:04:00 +0530 Subject: [PATCH] test(proxy): cover nested access group resolution + validation 11 unit tests for the nested-groups feature: - resolve_nested_groups: flat compat, single/multi-hop nesting, self- and indirect cycles (warning logged + safe degradation), propagation, pure-composition groups, diamond inheritance - _classify_member_names: router model takes precedence over group name when both match; mixed input splits into the three buckets - validate_models_exist: accepts known group names; reports truly unknown; backwards compatible when no groups passed All pure-function tests - no DB / no fake Prisma client needed. Refs #28032 --- .../proxy/auth/test_nested_access_groups.py | 266 ++++++++++++++++++ 1 file changed, 266 insertions(+) create mode 100644 tests/test_litellm/proxy/auth/test_nested_access_groups.py diff --git a/tests/test_litellm/proxy/auth/test_nested_access_groups.py b/tests/test_litellm/proxy/auth/test_nested_access_groups.py new file mode 100644 index 00000000000..64061958698 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_nested_access_groups.py @@ -0,0 +1,266 @@ +""" +Unit tests for nested access group resolution (#28032). + +Covers: +- resolve_nested_groups (DFS + cycle detection) +- _get_models_from_access_groups with the new group_memberships kwarg +- _classify_member_names precedence +- validate_models_exist with known_access_groups +""" +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +import pytest + +from litellm.proxy.auth.model_checks import ( + _get_models_from_access_groups, + resolve_nested_groups, +) +from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + _classify_member_names, + validate_models_exist, +) + + +# --------------------------------------------------------------------------- +# resolve_nested_groups + _get_models_from_access_groups +# --------------------------------------------------------------------------- + + +def test_resolve_flat_group_unchanged(): + """When group_memberships is empty, behavior must match today's flat path.""" + model_access_groups = {"prod": ["gpt-4", "claude-3-opus"]} + all_models = ["prod", "gemini-pro"] + + result = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=list(all_models), + ) + + assert sorted(result) == sorted(["gpt-4", "claude-3-opus", "gemini-pro"]) + + +def test_resolve_single_level_nested_group(): + """A parent group whose child has direct models expands to those models.""" + model_access_groups = {"image": ["dall-e-3", "stable-diffusion-xl"]} + group_memberships = {"project-x": ["image"]} + all_models = ["project-x"] + + result = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=list(all_models), + group_memberships=group_memberships, + ) + + assert sorted(result) == sorted(["dall-e-3", "stable-diffusion-xl"]) + + +def test_resolve_three_level_deep_nested_group(): + """A -> B -> C -> [models] resolves through three hops.""" + model_access_groups = {"C": ["gpt-4"]} + group_memberships = {"A": ["B"], "B": ["C"]} + + result = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=["A"], + group_memberships=group_memberships, + ) + + assert result == ["gpt-4"] + + +def test_detect_and_skip_self_referential_cycle(caplog): + """A -> A logs a warning and returns A's direct models without exploding.""" + model_access_groups = {"A": ["gpt-4"]} + group_memberships = {"A": ["A"]} + + with caplog.at_level("WARNING"): + result = resolve_nested_groups( + group_name="A", + model_access_groups=model_access_groups, + group_memberships=group_memberships, + visited=set(), + ) + + assert result == ["gpt-4"] + assert any("cycle detected" in m.lower() for m in caplog.messages) + + +def test_detect_and_skip_indirect_cycle(caplog): + """A -> B -> A logs a warning and returns the union of both direct sets.""" + model_access_groups = {"A": ["gpt-4"], "B": ["claude-3"]} + group_memberships = {"A": ["B"], "B": ["A"]} + + with caplog.at_level("WARNING"): + result = resolve_nested_groups( + group_name="A", + model_access_groups=model_access_groups, + group_memberships=group_memberships, + visited=set(), + ) + + assert sorted(set(result)) == sorted(["gpt-4", "claude-3"]) + assert any("cycle detected" in m.lower() for m in caplog.messages) + + +def test_propagation_to_parent(): + """Adding a model to a child group surfaces on the parent's resolved list.""" + # Initial: parent resolves to whatever the child has at the moment of call + model_access_groups = {"image": ["dall-e-3"]} + group_memberships = {"project-x": ["image"]} + + before = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=["project-x"], + group_memberships=group_memberships, + ) + assert before == ["dall-e-3"] + + # Add a new model to the child group; parent must see it on the next resolve + model_access_groups["image"].append("stable-diffusion-xl") + after = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=["project-x"], + group_memberships=group_memberships, + ) + assert sorted(after) == sorted(["dall-e-3", "stable-diffusion-xl"]) + + +def test_pure_composition_group_resolves(): + """A group present only in memberships (no deployment tags) still resolves.""" + model_access_groups = {"image": ["dall-e-3"], "reasoning": ["o1"]} + group_memberships = {"project-x": ["image", "reasoning"]} + + result = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=["project-x"], + group_memberships=group_memberships, + ) + + assert sorted(result) == sorted(["dall-e-3", "o1"]) + + +def test_include_model_access_groups_keeps_parent_name(): + """When include_model_access_groups=True, the group name is preserved alongside expansion.""" + model_access_groups = {"image": ["dall-e-3"]} + group_memberships = {"project-x": ["image"]} + + result = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=["project-x"], + include_model_access_groups=True, + group_memberships=group_memberships, + ) + + assert "project-x" in result + assert "dall-e-3" in result + + +# --------------------------------------------------------------------------- +# _classify_member_names +# --------------------------------------------------------------------------- + + +def test_classify_router_model_takes_precedence_over_group(): + """A name registered as both a router model and a group is classified as a model.""" + real, child, unknown = _classify_member_names( + names=["shared-name"], + router_model_names={"shared-name"}, + known_access_groups={"shared-name"}, + ) + assert real == ["shared-name"] + assert child == [] + assert unknown == [] + + +def test_classify_splits_mixed_input(): + """Mixed input is split across the three buckets in order.""" + real, child, unknown = _classify_member_names( + names=["gpt-4", "image-group", "mystery"], + router_model_names={"gpt-4"}, + known_access_groups={"image-group"}, + ) + assert real == ["gpt-4"] + assert child == ["image-group"] + assert unknown == ["mystery"] + + +# --------------------------------------------------------------------------- +# validate_models_exist +# --------------------------------------------------------------------------- + + +def test_validate_models_exist_accepts_known_group_names(): + """Group names are valid members when passed via known_access_groups.""" + + class FakeRouter: + def get_model_names(self): + return ["gpt-4"] + + all_valid, missing = validate_models_exist( + model_names=["gpt-4", "image-group"], + llm_router=FakeRouter(), + known_access_groups={"image-group"}, + ) + assert all_valid is True + assert missing == [] + + +def test_validate_models_exist_reports_unknown_names(): + """Names that match neither a router model nor a known group are reported.""" + + class FakeRouter: + def get_model_names(self): + return ["gpt-4"] + + all_valid, missing = validate_models_exist( + model_names=["gpt-4", "image-group", "ghost"], + llm_router=FakeRouter(), + known_access_groups={"image-group"}, + ) + assert all_valid is False + assert missing == ["ghost"] + + +def test_validate_models_exist_backwards_compatible_without_groups(): + """Without known_access_groups, behavior is identical to today's pure-model check.""" + + class FakeRouter: + def get_model_names(self): + return ["gpt-4"] + + all_valid, missing = validate_models_exist( + model_names=["gpt-4", "anything-else"], + llm_router=FakeRouter(), + ) + assert all_valid is False + assert missing == ["anything-else"] + + +# --------------------------------------------------------------------------- +# Smoke: chained scenarios +# --------------------------------------------------------------------------- + + +def test_diamond_inheritance_resolves_once(): + """ + Diamond shape: + project-x -> [image, reasoning] + image -> [dall-e-3] + reasoning -> [dall-e-3, o1] (dall-e-3 deliberately shared) + + Resolution must include all models. Duplicates are caller's job to dedup. + """ + model_access_groups = {"image": ["dall-e-3"], "reasoning": ["dall-e-3", "o1"]} + group_memberships = {"project-x": ["image", "reasoning"]} + + result = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=["project-x"], + group_memberships=group_memberships, + ) + + # set() removes the legitimate duplicate, all 2 unique models should be present + assert sorted(set(result)) == sorted(["dall-e-3", "o1"])