diff --git a/tests/recipes/test_recipes.py b/tests/recipes/test_recipes.py index 172d9fe5..6230817b 100644 --- a/tests/recipes/test_recipes.py +++ b/tests/recipes/test_recipes.py @@ -149,6 +149,7 @@ def test_recipe_by_production_semantic_version(mocker): "production_version_id": "production-version-id" } ) + mocker.patch("wrangles.data.model_claim", return_value={}) model_content = mocker.patch( "wrangles.data.model_content", return_value={"recipe": "{}"} @@ -173,6 +174,7 @@ def test_recipe_by_production_semantic_version_falls_back_to_latest( "wrangles.data.model", return_value={"purpose": "recipe"} ) + mocker.patch("wrangles.data.model_claim", return_value={}) model_content = mocker.patch( "wrangles.data.model_content", return_value={"recipe": "{}"} diff --git a/tests/recipes/test_variables.py b/tests/recipes/test_variables.py index d4cebfcb..32ec379e 100644 --- a/tests/recipes/test_variables.py +++ b/tests/recipes/test_variables.py @@ -367,17 +367,12 @@ def test_variables_variable_overwrite(): assert isinstance(df['vars'][0], dict) -def test_applied_permission_group_variable(monkeypatch): +def test_applied_permission_group_variable_explicit(monkeypatch): """ - Test that the authenticated user's effective permission group is available as a recipe variable. + Test that an explicitly-passed applied_permission_group is available as + a recipe variable for a non-model_id recipe (there is no server-side + default to fall back to in that case). """ - token = wrangles.auth._jwt.encode( - {"applied_permission_group": "enterprise"}, - "test-secret", - algorithm="HS256" - ) - monkeypatch.setattr(wrangles.auth, "get_access_token", lambda: token) - df = wrangles.recipe.run( """ read: @@ -385,7 +380,8 @@ def test_applied_permission_group_variable(monkeypatch): rows: 1 values: group: ${applied_permission_group} - """ + """, + variables={"applied_permission_group": "enterprise"} ) assert df['group'][0] == 'enterprise' @@ -395,8 +391,6 @@ def test_applied_permission_group_variable_if(monkeypatch): """ Test that applied_permission_group can be used in Python-style if conditions. """ - monkeypatch.setattr(wrangles.auth, "get_applied_permission_group", lambda: "enterprise") - df = wrangles.recipe.run( """ read: @@ -409,46 +403,68 @@ def test_applied_permission_group_variable_if(monkeypatch): output: allowed value: true if: applied_permission_group == 'enterprise' - """ + """, + variables={"applied_permission_group": "enterprise"} ) assert df['allowed'][0] == True -def test_applied_permission_group_variable_user_override(monkeypatch): +def test_applied_permission_group_variable_from_recipe_metadata(monkeypatch): """ - Test that explicit variables still override the authenticated permission group. + Test that recipe metadata permission group is preferred for remote recipes. """ - monkeypatch.setattr(wrangles.auth, "get_applied_permission_group", lambda: "enterprise") - - df = wrangles.recipe.run( - """ - read: - - test: - rows: 1 - values: - group: ${applied_permission_group} - """, - variables={"applied_permission_group": "manual"} + monkeypatch.setattr( + wrangles.recipe._data, + "model", + lambda model_id: { + "purpose": "recipe", + "production_version_id": "v1", + "applied_permission_group": "metadata-group", + } ) + # No model claim available - the metadata-derived value above should be + # left untouched rather than overridden. + monkeypatch.setattr(wrangles.recipe._data, "model_claim", lambda model_id: {}) + monkeypatch.setattr( + wrangles.recipe._data, + "model_content", + lambda model_id, version_id=None: { + "recipe": """ + read: + - test: + rows: 1 + values: + group: ${applied_permission_group} + """ + } + ) + + df = wrangles.recipe.run("12345678-1234-1234") - assert df['group'][0] == 'manual' + assert df["group"][0] == "metadata-group" -def test_applied_permission_group_variable_from_recipe_metadata(monkeypatch): +def test_applied_permission_group_variable_metadata_overrides_explicit(monkeypatch, caplog): """ - Test that recipe metadata permission group is preferred for remote recipes. + A model_id-addressed recipe's real permission group (resolved + server-side from the model's metadata) must override an explicit + variables={"applied_permission_group": ...} too - otherwise a caller + could simply claim a higher role than the model's database actually + grants them. """ - monkeypatch.setattr(wrangles.auth, "get_applied_permission_group", lambda: "token-group") monkeypatch.setattr( wrangles.recipe._data, "model", lambda model_id: { "purpose": "recipe", "production_version_id": "v1", - "applied_permission_group": "metadata-group", + "applied_permission_group": "editor", } ) + # No model claim available - only the metadata-derived override (from + # data.model above) is exercised by this test. + monkeypatch.setattr(wrangles.recipe._data, "model_claim", lambda model_id: {}) monkeypatch.setattr( wrangles.recipe._data, "model_content", @@ -463,6 +479,133 @@ def test_applied_permission_group_variable_from_recipe_metadata(monkeypatch): } ) + with caplog.at_level("WARNING"): + df = wrangles.recipe.run( + "12345678-1234-1234", + variables={"applied_permission_group": "admin"} + ) + + assert df["group"][0] == "editor" + assert "does not match this model's actual permission group" in caplog.text + + +def test_applied_permission_level_variable_from_model_claim(monkeypatch): + """ + Test that applied_permission_level is filled from the model claim's role + when running a model_id directly, e.g. from Python. + """ + monkeypatch.setattr( + wrangles.recipe._data, + "model", + lambda model_id: {"purpose": "recipe", "production_version_id": "v1"} + ) + monkeypatch.setattr( + wrangles.recipe._data, + "model_claim", + lambda model_id: { + "model_id": model_id, + "role": "viewer", + "applied_group": "Dev (WrangleWorks)", + } + ) + monkeypatch.setattr( + wrangles.recipe._data, + "model_content", + lambda model_id, version_id=None: { + "recipe": """ + read: + - test: + rows: 1 + values: + level: ${applied_permission_level} + group: ${applied_permission_group} + """ + } + ) + df = wrangles.recipe.run("12345678-1234-1234") - assert df["group"][0] == "metadata-group" + assert df["level"][0] == "viewer" + assert df["group"][0] == "Dev (WrangleWorks)" + + +def test_applied_permission_level_variable_claim_overrides_explicit(monkeypatch, caplog): + """ + Like applied_permission_group, an explicit + variables={"applied_permission_level": ...} must not let a caller claim + a higher role than the model claim actually grants - run(model, + variables={"applied_permission_level": "admin"}) when the real claim + says "viewer" must use "viewer". + """ + monkeypatch.setattr( + wrangles.recipe._data, + "model", + lambda model_id: {"purpose": "recipe", "production_version_id": "v1"} + ) + monkeypatch.setattr( + wrangles.recipe._data, + "model_claim", + lambda model_id: { + "model_id": model_id, + "role": "viewer", + "applied_group": "Dev (WrangleWorks)", + } + ) + monkeypatch.setattr( + wrangles.recipe._data, + "model_content", + lambda model_id, version_id=None: { + "recipe": """ + read: + - test: + rows: 1 + values: + level: ${applied_permission_level} + """ + } + ) + + with caplog.at_level("WARNING"): + df = wrangles.recipe.run( + "12345678-1234-1234", + variables={"applied_permission_level": "admin"} + ) + + assert df["level"][0] == "viewer" + assert "does not match this model's actual permission level" in caplog.text + + +def test_model_claim_failure_does_not_block_recipe_load(monkeypatch, caplog): + """ + A failure resolving the model claim (e.g. network issue) must not block + the recipe from loading - applied_permission_level is simply left unset. + """ + monkeypatch.setattr( + wrangles.recipe._data, + "model", + lambda model_id: {"purpose": "recipe", "production_version_id": "v1"} + ) + + def _raise(model_id): + raise RuntimeError("boom") + + monkeypatch.setattr(wrangles.recipe._data, "model_claim", _raise) + monkeypatch.setattr( + wrangles.recipe._data, + "model_content", + lambda model_id, version_id=None: { + "recipe": """ + read: + - test: + rows: 1 + values: + result: kept + """ + } + ) + + with caplog.at_level("WARNING"): + df = wrangles.recipe.run("12345678-1234-1234") + + assert df["result"][0] == "kept" + assert "Could not resolve model claim" in caplog.text diff --git a/tests/recipes/wrangles/test_extract.py b/tests/recipes/wrangles/test_extract.py index d6a6f5cd..22adb157 100644 --- a/tests/recipes/wrangles/test_extract.py +++ b/tests/recipes/wrangles/test_extract.py @@ -4413,8 +4413,11 @@ def test_ai_invalid_model_per_row_error(self): "data": ["wrench 25mm", "6m cable"], }) ) + # OpenAI may report an unknown model as any 4xx client error + # (e.g. 400 invalid_request_error or 404 model_not_found) + # depending on the API version, so don't pin to one exact code. assert all( - "OpenAI API error" in value and "status=400" in value + "OpenAI API error" in value and "status=4" in value for value in df['length'] ) diff --git a/tests/recipes/wrangles/test_main.py b/tests/recipes/wrangles/test_main.py index 10ce1f70..b691f9ee 100644 --- a/tests/recipes/wrangles/test_main.py +++ b/tests/recipes/wrangles/test_main.py @@ -4049,7 +4049,8 @@ def fake_model_content(model_id, version_id=None): """ with patch('wrangles.recipe._data.model', side_effect=fake_model), \ - patch('wrangles.recipe._data.model_content', side_effect=fake_model_content): + patch('wrangles.recipe._data.model_content', side_effect=fake_model_content), \ + patch('wrangles.recipe._data.model_claim', return_value={}): with pytest.raises(Exception) as info: wrangles.recipe.run(outer_recipe) diff --git a/tests/test_data.py b/tests/test_data.py index dafb26ba..37405b8b 100644 --- a/tests/test_data.py +++ b/tests/test_data.py @@ -27,6 +27,7 @@ def _mock_model_response(monkeypatch, response): lambda: data.model(MODEL_ID), lambda: data.model_update(MODEL_ID, {"name": "Updated model"}), lambda: data.model_content(MODEL_ID), + lambda: data.model_claim(MODEL_ID), ], ) def test_model_endpoints_raise_authentication_error_for_401(monkeypatch, call_model_endpoint): @@ -48,6 +49,7 @@ def test_model_endpoints_raise_authentication_error_for_401(monkeypatch, call_mo lambda: data.model(MODEL_ID), lambda: data.model_update(MODEL_ID, {"name": "Updated model"}), lambda: data.model_content(MODEL_ID), + lambda: data.model_claim(MODEL_ID), ], ) def test_model_endpoints_raise_authorization_error_for_403(monkeypatch, call_model_endpoint): @@ -80,3 +82,17 @@ def test_model_content_success_returns_content(monkeypatch): _mock_model_response(monkeypatch, FakeResponse(200, content)) assert data.model_content(MODEL_ID) == content + + +def test_model_claim_success_returns_claim(monkeypatch): + claim = { + "model_id": MODEL_ID, + "role": "admin", + "organization_id": "team-id", + "applied_group": "Dev (WrangleWorks)", + "applied_group_type": "group", + "applied_group_id": "team-id", + } + _mock_model_response(monkeypatch, FakeResponse(200, claim)) + + assert data.model_claim(MODEL_ID) == claim diff --git a/wrangles/auth.py b/wrangles/auth.py index ffda8510..9ce3025c 100644 --- a/wrangles/auth.py +++ b/wrangles/auth.py @@ -88,7 +88,7 @@ def get_access_token(): def extract_applied_permission_group(source: dict): """ - Extract the effective permission group from a metadata or token payload. + Extract the effective permission group from a model metadata payload. """ if not isinstance(source, dict): return None @@ -96,21 +96,23 @@ def extract_applied_permission_group(source: dict): return source.get("applied_permission_group") -def get_applied_permission_group(): +def extract_applied_permission_group_from_claim(claim: dict): """ - Return the authenticated user's effective permission group from the current access token. - - If no user is authenticated or the token does not contain the claim, - return None so recipes can still run without backend credentials. + Extract the applied permission group (the group/org/user display name a + model claim is granted through) from a /model/claim response. """ - try: - token = get_access_token() - except Exception: + if not isinstance(claim, dict): return None - try: - claims = _jwt.decode(token, options={"verify_signature": False}) - except Exception: + return claim.get("applied_group") + + +def extract_applied_permission_level(claim: dict): + """ + Extract the applied permission level (the user's role on a model - e.g. + admin, editor, viewer) from a /model/claim response. + """ + if not isinstance(claim, dict): return None - return extract_applied_permission_group(claims) + return claim.get("role") diff --git a/wrangles/data.py b/wrangles/data.py index 25b1eb37..60f31b6f 100644 --- a/wrangles/data.py +++ b/wrangles/data.py @@ -93,6 +93,29 @@ def model_update(id: str, metadata: dict) -> None: _raise_model_response_error(response, id, 'update') +def model_claim(id: str) -> dict: + """ + Get the current user's highest effective claim on a model - their role + and the group/org/user that claim is applied through. + + :param id: Model ID + :returns: Dict with model_id, role, organization_id, applied_group, \ + applied_group_type, applied_group_id + """ + response = _utils.request_retries( + request_type='GET', + url=f'{_config.api_host}/model/claim', + **{ + 'params': {'model_id': id}, + 'headers': {'Authorization': f'Bearer {_auth.get_access_token()}'} + } + ) + if response.ok: + return response.json() + else: + _raise_model_response_error(response, id, 'access') + + def model_content(id: str, version_id: str = None) -> list: """ Get the training data for a model diff --git a/wrangles/recipe.py b/wrangles/recipe.py index c35f19e4..73d53717 100644 --- a/wrangles/recipe.py +++ b/wrangles/recipe.py @@ -79,9 +79,6 @@ def _load_recipe( user_variable_keys = set(variables.keys()) - if "applied_permission_group" not in variables: - variables["applied_permission_group"] = _auth.get_applied_permission_group() - # Accept path-like objects (e.g. pathlib.Path) by converting to str if isinstance(recipe, _os.PathLike): recipe = str(recipe) @@ -130,8 +127,68 @@ def _load_recipe( metadata_applied_permission_group = _auth.extract_applied_permission_group(metadata) if metadata_applied_permission_group is not None: - if "applied_permission_group" not in user_variable_keys: - variables["applied_permission_group"] = metadata_applied_permission_group + + # Authoritative for a model_id-addressed recipe: this is the + # caller's real, freshly-verified role on this specific model, + # resolved server-side from their access token - not something + # a caller should be able to bypass by simply passing a + # different value. This intentionally overrides an explicit + # variables={"applied_permission_group": ...} too, unlike the + # generic JWT-derived fallback above, which has no independent + # source to check a caller-supplied value against. + if ( + "applied_permission_group" in user_variable_keys + and variables["applied_permission_group"] != metadata_applied_permission_group + ): + _logging.warning( + "applied_permission_group=" + f"'{variables['applied_permission_group']}' was passed in but does not " + f"match this model's actual permission group '{metadata_applied_permission_group}' " + "- using the actual value." + ) + variables["applied_permission_group"] = metadata_applied_permission_group + + # The model's real, freshly-verified claim (role + applied group) is + # authoritative for a model_id-addressed recipe, and is the source + # for applied_permission_level (there's no generic, model-less + # notion of a user's "role" the way there is for a permission + # group, so unlike applied_permission_group there is no JWT-derived + # fallback for this variable). Best-effort: a failure here should + # not block the recipe from loading. + try: + claim = _data.model_claim(model_id) + except Exception as e: + _logging.warning(f": Could not resolve model claim for {model_id}: {e}") + claim = None + + if claim: + claim_applied_permission_group = _auth.extract_applied_permission_group_from_claim(claim) + if claim_applied_permission_group is not None: + if ( + "applied_permission_group" in user_variable_keys + and variables["applied_permission_group"] != claim_applied_permission_group + ): + _logging.warning( + "applied_permission_group=" + f"'{variables['applied_permission_group']}' was passed in but does not " + f"match this model's actual permission group '{claim_applied_permission_group}' " + "- using the actual value." + ) + variables["applied_permission_group"] = claim_applied_permission_group + + claim_applied_permission_level = _auth.extract_applied_permission_level(claim) + if claim_applied_permission_level is not None: + if ( + "applied_permission_level" in user_variable_keys + and variables.get("applied_permission_level") != claim_applied_permission_level + ): + _logging.warning( + "applied_permission_level=" + f"'{variables['applied_permission_level']}' was passed in but does not " + f"match this model's actual permission level '{claim_applied_permission_level}' " + "- using the actual value." + ) + variables["applied_permission_level"] = claim_applied_permission_level # Using model_id in wrong function purpose = metadata['purpose']