From e3c076f4e5d72b91b8e492f92a2055ab26e4c83c Mon Sep 17 00:00:00 2001 From: yaswanth-pula-skyflow Date: Mon, 3 Aug 2026 15:42:28 +0530 Subject: [PATCH 1/7] SK-3039:Fix missing roles & ctx in bearerToken. --- skyflow/utils/validations/_validations.py | 28 ++++- skyflow/vault/client/client.py | 13 +- tests/client/test_skyflow.py | 49 ++++++++ tests/utils/validations/test__validations.py | 108 +++++++++++++++- tests/vault/client/test__client.py | 123 ++++++++++++++++++- 5 files changed, 307 insertions(+), 14 deletions(-) diff --git a/skyflow/utils/validations/_validations.py b/skyflow/utils/validations/_validations.py index 42abe188..ac31ebbc 100644 --- a/skyflow/utils/validations/_validations.py +++ b/skyflow/utils/validations/_validations.py @@ -2,6 +2,7 @@ import json import os from skyflow.service_account import is_expired +from skyflow.service_account._utils import _validate_and_resolve_ctx from skyflow.utils.enums import LogLevel, Env, RedactionType, TokenMode, DetectEntities, DetectOutputTranscriptions, \ MaskingMethod from skyflow.error import SkyflowError @@ -104,20 +105,35 @@ def validate_credentials(logger, credentials, config_id_type=None, config_id=Non if CredentialField.ROLES in credentials: validate_required_field( logger, credentials, CredentialField.ROLES, list, - SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE_IN_CONFIG.value.format(config_id_type, config_id) - if config_id_type and config_id else SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE.value, SkyflowMessages.Error.EMPTY_ROLES_IN_CONFIG.value.format(config_id_type, config_id) - if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value + if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value, + SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE_IN_CONFIG.value.format(config_id_type, config_id) + if config_id_type and config_id else SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE.value ) + if not credentials.get(CredentialField.ROLES): + raise SkyflowError( + SkyflowMessages.Error.EMPTY_ROLES_IN_CONFIG.value.format(config_id_type, config_id) + if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value, + invalid_input_error_code + ) if CredentialField.CONTEXT in credentials: - validate_required_field( - logger, credentials, CredentialField.CONTEXT, str, + empty_context_error = ( SkyflowMessages.Error.EMPTY_CONTEXT_IN_CONFIG.value.format(config_id_type, config_id) - if config_id_type and config_id else SkyflowMessages.Error.EMPTY_CONTEXT.value, + if config_id_type and config_id else SkyflowMessages.Error.EMPTY_CONTEXT.value + ) + validate_required_field( + logger, credentials, CredentialField.CONTEXT, (str, dict), + empty_context_error, SkyflowMessages.Error.INVALID_CONTEXT_IN_CONFIG.value.format(config_id_type, config_id) if config_id_type and config_id else SkyflowMessages.Error.INVALID_CONTEXT.value ) + context = credentials.get(CredentialField.CONTEXT) + if isinstance(context, dict): + if not context: + raise SkyflowError(empty_context_error, invalid_input_error_code) + # Surface invalid ctx keys at config time instead of on the first API call. + _validate_and_resolve_ctx(context) if CredentialField.CREDENTIALS_STRING in credentials: validate_required_field( diff --git a/skyflow/vault/client/client.py b/skyflow/vault/client/client.py index 8023646c..e7255cb1 100644 --- a/skyflow/vault/client/client.py +++ b/skyflow/vault/client/client.py @@ -74,10 +74,15 @@ def get_bearer_token(self, credentials): elif CredentialField.TOKEN in credentials: return credentials.get(CredentialField.TOKEN) - options = { - OptionField.ROLE_IDS: self.__config.get(OptionField.ROLES), - OptionField.CTX: self.__config.get(OptionField.CTX) - } + options = {} + roles = credentials.get(CredentialField.ROLES) + if roles: + options[OptionField.ROLE_IDS] = roles + + context = credentials.get(CredentialField.CONTEXT) + if context is not None: + options[OptionField.CTX] = context + if CredentialField.TOKEN_URI_OPTION in credentials and credentials.get(CredentialField.TOKEN_URI_OPTION): options[CredentialField.TOKEN_URI_OPTION] = credentials.get(CredentialField.TOKEN_URI_OPTION) diff --git a/tests/client/test_skyflow.py b/tests/client/test_skyflow.py index 5b7ea675..d2c07392 100644 --- a/tests/client/test_skyflow.py +++ b/tests/client/test_skyflow.py @@ -420,6 +420,55 @@ def test_get_bearer_token_passes_token_uri_option(self, _mock_expired, mock_gen) self.assertEqual(options_passed["token_uri"], "https://custom-token-uri.com/token") +class TestBearerTokenContextAndRolesEndToEnd(unittest.TestCase): + """roles/context configured via add_vault_config() must reach the token engine.""" + + CREDENTIALS_STRING = '{"clientID":"id","privateKey":"pk","keyID":"kid","tokenURI":"https://token.uri"}' + + def _build_and_generate(self, credentials, mock_gen): + config = { + "vault_id": "VAULT_ID", + "cluster_id": "CLUSTER_ID", + "env": Env.DEV, + "credentials": credentials, + } + client = ( + Skyflow.builder() + .add_vault_config(config) + .set_log_level(LogLevel.OFF) + .build() + ) + vault_client = client._Skyflow__builder.get_vault_config("VAULT_ID").get("vault_client") + vault_client.initialize_client_configuration() + return mock_gen.call_args[0][1] + + @patch("skyflow.vault.client.client.Skyflow") + @patch("skyflow.vault.client.client.generate_bearer_token_from_creds", return_value=("token", "bearer")) + def test_string_context_and_roles_reach_token_engine(self, mock_gen, _mock_api): + options = self._build_and_generate( + { + "credentials_string": self.CREDENTIALS_STRING, + "roles": ["role_id_1", "role_id_2"], + "context": "user_12345", + }, + mock_gen, + ) + self.assertEqual(options["role_ids"], ["role_id_1", "role_id_2"]) + self.assertEqual(options["ctx"], "user_12345") + + @patch("skyflow.vault.client.client.Skyflow") + @patch("skyflow.vault.client.client.generate_bearer_token_from_creds", return_value=("token", "bearer")) + def test_dict_context_reaches_token_engine(self, mock_gen, _mock_api): + options = self._build_and_generate( + { + "credentials_string": self.CREDENTIALS_STRING, + "context": {"role": "admin", "department": "finance"}, + }, + mock_gen, + ) + self.assertEqual(options["ctx"], {"role": "admin", "department": "finance"}) + + class TestUpdateLogLevelDeprecation(unittest.TestCase): def _build_client(self): return Skyflow.builder().add_vault_config(VALID_VAULT_CONFIG).build() diff --git a/tests/utils/validations/test__validations.py b/tests/utils/validations/test__validations.py index 9ee05cef..c9573084 100644 --- a/tests/utils/validations/test__validations.py +++ b/tests/utils/validations/test__validations.py @@ -145,7 +145,113 @@ def test_validate_credentials_with_empty_context(self): with self.assertRaises(SkyflowError) as context: validate_credentials(self.logger, credentials) self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_CONTEXT.value) - + + # ------------------------------------------------------------------ # + # credentials.context — string and dict/JSON object # + # ------------------------------------------------------------------ # + + def test_validate_credentials_with_string_context(self): + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "context": "user_12345" + } + validate_credentials(self.logger, credentials) + + def test_validate_credentials_with_dict_context(self): + """A dict/JSON object context is accepted — the token engine supports it.""" + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "context": {"role": "admin", "department": "finance"} + } + validate_credentials(self.logger, credentials) + + def test_validate_credentials_with_empty_dict_context(self): + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "context": {} + } + with self.assertRaises(SkyflowError) as context: + validate_credentials(self.logger, credentials) + self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_CONTEXT.value) + + def test_validate_credentials_with_invalid_dict_context_key(self): + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "context": {"invalid key": "value"} + } + with self.assertRaises(SkyflowError) as context: + validate_credentials(self.logger, credentials) + self.assertEqual( + context.exception.message, + SkyflowMessages.Error.INVALID_CTX_MAP_KEY.value.format("invalid key") + ) + + def test_validate_credentials_with_invalid_context_type(self): + for invalid_context in [123, ["user_12345"], True]: + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "context": invalid_context + } + with self.assertRaises(SkyflowError) as context: + validate_credentials(self.logger, credentials) + self.assertEqual(context.exception.message, SkyflowMessages.Error.INVALID_CONTEXT.value) + + def test_validate_credentials_with_dict_context_in_config(self): + """Config-scoped messages are used when a config id is available.""" + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "context": {} + } + with self.assertRaises(SkyflowError) as context: + validate_credentials(self.logger, credentials, "vault", "vault123") + self.assertEqual( + context.exception.message, + SkyflowMessages.Error.EMPTY_CONTEXT_IN_CONFIG.value.format("vault", "vault123") + ) + + def test_validate_vault_config_with_dict_context_and_roles(self): + from skyflow.utils.enums import Env + config = { + "vault_id": "vault123", + "cluster_id": "cluster123", + "credentials": { + "credentials_string": '{"clientID": "x"}', + "roles": ["role_id_1", "role_id_2"], + "context": {"role": "admin", "department": "finance"} + }, + "env": Env.DEV + } + self.assertTrue(validate_vault_config(self.logger, config)) + + # ------------------------------------------------------------------ # + # credentials.roles # + # ------------------------------------------------------------------ # + + def test_validate_credentials_with_roles(self): + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "roles": ["role_id_1", "role_id_2"] + } + validate_credentials(self.logger, credentials) + + def test_validate_credentials_with_invalid_roles_type(self): + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "roles": "role_id_1" + } + with self.assertRaises(SkyflowError) as context: + validate_credentials(self.logger, credentials) + self.assertEqual(context.exception.message, SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE.value) + + def test_validate_credentials_with_empty_roles(self): + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "roles": [] + } + with self.assertRaises(SkyflowError) as context: + validate_credentials(self.logger, credentials) + self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_ROLES.value) + def test_validate_log_level_valid(self): from skyflow.utils.enums import LogLevel log_level = LogLevel.ERROR diff --git a/tests/vault/client/test__client.py b/tests/vault/client/test__client.py index 4df508c7..1eb2435e 100644 --- a/tests/vault/client/test__client.py +++ b/tests/vault/client/test__client.py @@ -9,9 +9,7 @@ "credentials": "some_credentials", "cluster_id": "test_cluster_id", "env": "test_env", - "vault_id": "test_vault_id", - "roles": ["role_id_1", "role_id_2"], - "ctx": "context" + "vault_id": "test_vault_id" } CREDENTIALS_WITH_API_KEY = {"api_key": "dummy_api_key"} @@ -19,6 +17,25 @@ CREDENTIALS_WITH_PATH = {"path": "/some/path/credentials.json"} CREDENTIALS_WITH_STRING = {"credentials_string": '{"clientID": "x"}'} +ROLES = ["role_id_1", "role_id_2"] +STRING_CONTEXT = "user_12345" +DICT_CONTEXT = {"role": "admin", "department": "finance"} + +CREDENTIALS_WITH_PATH_ROLES_AND_STRING_CONTEXT = { + "path": "/some/path/credentials.json", + "roles": ROLES, + "context": STRING_CONTEXT, +} +CREDENTIALS_WITH_PATH_AND_DICT_CONTEXT = { + "path": "/some/path/credentials.json", + "context": DICT_CONTEXT, +} +CREDENTIALS_WITH_STRING_ROLES_AND_CONTEXT = { + "credentials_string": '{"clientID": "x"}', + "roles": ROLES, + "context": STRING_CONTEXT, +} + class TestVaultClient(unittest.TestCase): def setUp(self): @@ -297,6 +314,106 @@ def test_get_bearer_token_reuses_valid_cached_token(self, mock_log, mock_is_expi mock_generate.assert_not_called() self.assertEqual(result, "valid_token") + # ------------------------------------------------------------------ # + # get_bearer_token — roles / context forwarded to the token engine # + # ------------------------------------------------------------------ # + + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("sa_token", None)) + def test_get_bearer_token_forwards_roles_and_string_context_from_path_credentials(self, mock_generate): + """roles/context nested under credentials must reach the token engine options.""" + self.vault_client.get_bearer_token(CREDENTIALS_WITH_PATH_ROLES_AND_STRING_CONTEXT) + + _, options, _ = mock_generate.call_args[0] + self.assertEqual(options["role_ids"], ROLES) + self.assertEqual(options["ctx"], STRING_CONTEXT) + + @patch("skyflow.vault.client.client.generate_bearer_token_from_creds", return_value=("sa_token_str", None)) + @patch("skyflow.vault.client.client.log_info") + def test_get_bearer_token_forwards_roles_and_context_from_credentials_string(self, mock_log, mock_generate): + self.vault_client.get_bearer_token(CREDENTIALS_WITH_STRING_ROLES_AND_CONTEXT) + + _, options, _ = mock_generate.call_args[0] + self.assertEqual(options["role_ids"], ROLES) + self.assertEqual(options["ctx"], STRING_CONTEXT) + + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("sa_token", None)) + def test_get_bearer_token_forwards_dict_context_unmodified(self, mock_generate): + """A dict/JSON object context is passed through as-is.""" + self.vault_client.get_bearer_token(CREDENTIALS_WITH_PATH_AND_DICT_CONTEXT) + + _, options, _ = mock_generate.call_args[0] + self.assertEqual(options["ctx"], DICT_CONTEXT) + self.assertNotIn("role_ids", options) + + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("sa_token", None)) + def test_get_bearer_token_omits_options_when_roles_and_context_absent(self, mock_generate): + self.vault_client.get_bearer_token(CREDENTIALS_WITH_PATH) + + _, options, _ = mock_generate.call_args[0] + self.assertNotIn("role_ids", options) + self.assertNotIn("ctx", options) + + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("sa_token", None)) + def test_get_bearer_token_ignores_top_level_config_roles_and_ctx(self, mock_generate): + """roles/ctx are only read from credentials — never from the top-level config.""" + vault_client = VaultClient({**CONFIG, "roles": ["stale_role"], "ctx": "stale_ctx"}) + vault_client.get_bearer_token(CREDENTIALS_WITH_PATH) + + _, options, _ = mock_generate.call_args[0] + self.assertNotIn("role_ids", options) + self.assertNotIn("ctx", options) + + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("sa_token", None)) + def test_get_bearer_token_forwards_token_uri_alongside_roles_and_context(self, mock_generate): + credentials = { + **CREDENTIALS_WITH_PATH_ROLES_AND_STRING_CONTEXT, + "token_uri": "https://manage.skyflowapis.com/v1/auth/sa/oauth/token", + } + self.vault_client.get_bearer_token(credentials) + + _, options, _ = mock_generate.call_args[0] + self.assertEqual(options["role_ids"], ROLES) + self.assertEqual(options["ctx"], STRING_CONTEXT) + self.assertEqual(options["token_uri"], credentials["token_uri"]) + + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("new_token", None)) + @patch("skyflow.vault.client.client.is_expired", return_value=True) + @patch("skyflow.vault.client.client.log_info") + def test_get_bearer_token_forwards_roles_and_context_on_refresh(self, mock_log, mock_is_expired, mock_generate): + """Auto-refresh of an expired token must carry roles/context too, not just the first call.""" + self.vault_client._VaultClient__bearer_token = "expired_token" + result = self.vault_client.get_bearer_token(CREDENTIALS_WITH_PATH_ROLES_AND_STRING_CONTEXT) + + self.assertEqual(result, "new_token") + _, options, _ = mock_generate.call_args[0] + self.assertEqual(options["role_ids"], ROLES) + self.assertEqual(options["ctx"], STRING_CONTEXT) + + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("refreshed_token", None)) + @patch("skyflow.vault.client.client.is_expired", return_value=True) + @patch("skyflow.vault.client.client.get_credentials") + @patch("skyflow.vault.client.client.get_vault_url") + @patch("skyflow.vault.client.client.VaultClient.initialize_api_client") + def test_initialize_client_configuration_refresh_preserves_roles_and_context( + self, mock_init_api_client, mock_get_vault_url, mock_get_credentials, + mock_is_expired, mock_generate + ): + """End-to-end refresh path through initialize_client_configuration keeps the options.""" + mock_get_credentials.return_value = CREDENTIALS_WITH_PATH_ROLES_AND_STRING_CONTEXT + mock_get_vault_url.return_value = "https://test-vault-url.com" + + # Warm client with an expired service account token + self.vault_client._VaultClient__api_client = MagicMock() + self.vault_client._VaultClient__is_static_token = False + self.vault_client._VaultClient__bearer_token = "expired_sa_token" + self.vault_client._VaultClient__credentials = CREDENTIALS_WITH_PATH_ROLES_AND_STRING_CONTEXT + + self.vault_client.initialize_client_configuration() + + _, options, _ = mock_generate.call_args[0] + self.assertEqual(options["role_ids"], ROLES) + self.assertEqual(options["ctx"], STRING_CONTEXT) + # ------------------------------------------------------------------ # # update_config # # ------------------------------------------------------------------ # From ca13208494dfcea50007e304207d0c55be050311 Mon Sep 17 00:00:00 2001 From: yaswanth-pula-skyflow Date: Mon, 3 Aug 2026 22:29:00 +0530 Subject: [PATCH 2/7] SK-3039:Update validations in ctx & scope in skyflow client. --- skyflow/utils/validations/__init__.py | 1 + skyflow/utils/validations/_validations.py | 59 +++++++++++--------- skyflow/vault/client/client.py | 13 +++-- tests/utils/validations/test__validations.py | 31 +++++----- tests/vault/client/test__client.py | 32 ++++++----- 5 files changed, 76 insertions(+), 60 deletions(-) diff --git a/skyflow/utils/validations/__init__.py b/skyflow/utils/validations/__init__.py index 2f0bc710..ede551d8 100644 --- a/skyflow/utils/validations/__init__.py +++ b/skyflow/utils/validations/__init__.py @@ -5,6 +5,7 @@ validate_update_vault_config, validate_update_connection_config, validate_credentials, + validate_token_options, validate_log_level, validate_delete_request, validate_query_request, diff --git a/skyflow/utils/validations/_validations.py b/skyflow/utils/validations/_validations.py index ac31ebbc..3e53bfc4 100644 --- a/skyflow/utils/validations/_validations.py +++ b/skyflow/utils/validations/_validations.py @@ -82,40 +82,25 @@ def validate_api_key(api_key: str, logger = None) -> bool: return True -def validate_credentials(logger, credentials, config_id_type=None, config_id=None): - key_present = [k for k in [CredentialField.PATH, CredentialField.TOKEN, CredentialField.CREDENTIALS_STRING, CredentialField.API_KEY] if credentials.get(k)] - - if len(key_present) == 0: - error_message = ( - SkyflowMessages.Error.INVALID_CREDENTIALS_IN_CONFIG.value.format(config_id_type, config_id) - if config_id_type and config_id else - SkyflowMessages.Error.INVALID_CREDENTIALS.value - ) - log_error_log(error_message, logger) - raise SkyflowError(error_message, invalid_input_error_code) - elif len(key_present) > 1: - error_message = ( - SkyflowMessages.Error.MULTIPLE_CREDENTIALS_PASSED_IN_CONFIG.value.format(config_id_type, config_id) - if config_id_type and config_id else - SkyflowMessages.Error.MULTIPLE_CREDENTIALS_PASSED.value - ) - log_error_log(error_message, logger) - raise SkyflowError(error_message, invalid_input_error_code) +def validate_token_options(logger, credentials, config_id_type=None, config_id=None): + """Validate the roles/context token options nested under credentials. + Shared by config validation and VaultClient, so a directly constructed client + cannot silently generate an unscoped or context-less token. + """ if CredentialField.ROLES in credentials: + empty_roles_error = ( + SkyflowMessages.Error.EMPTY_ROLES_IN_CONFIG.value.format(config_id_type, config_id) + if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value + ) validate_required_field( logger, credentials, CredentialField.ROLES, list, - SkyflowMessages.Error.EMPTY_ROLES_IN_CONFIG.value.format(config_id_type, config_id) - if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value, + empty_roles_error, SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE_IN_CONFIG.value.format(config_id_type, config_id) if config_id_type and config_id else SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE.value ) if not credentials.get(CredentialField.ROLES): - raise SkyflowError( - SkyflowMessages.Error.EMPTY_ROLES_IN_CONFIG.value.format(config_id_type, config_id) - if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value, - invalid_input_error_code - ) + raise SkyflowError(empty_roles_error, invalid_input_error_code) if CredentialField.CONTEXT in credentials: empty_context_error = ( @@ -135,6 +120,28 @@ def validate_credentials(logger, credentials, config_id_type=None, config_id=Non # Surface invalid ctx keys at config time instead of on the first API call. _validate_and_resolve_ctx(context) +def validate_credentials(logger, credentials, config_id_type=None, config_id=None): + key_present = [k for k in [CredentialField.PATH, CredentialField.TOKEN, CredentialField.CREDENTIALS_STRING, CredentialField.API_KEY] if credentials.get(k)] + + if len(key_present) == 0: + error_message = ( + SkyflowMessages.Error.INVALID_CREDENTIALS_IN_CONFIG.value.format(config_id_type, config_id) + if config_id_type and config_id else + SkyflowMessages.Error.INVALID_CREDENTIALS.value + ) + log_error_log(error_message, logger) + raise SkyflowError(error_message, invalid_input_error_code) + elif len(key_present) > 1: + error_message = ( + SkyflowMessages.Error.MULTIPLE_CREDENTIALS_PASSED_IN_CONFIG.value.format(config_id_type, config_id) + if config_id_type and config_id else + SkyflowMessages.Error.MULTIPLE_CREDENTIALS_PASSED.value + ) + log_error_log(error_message, logger) + raise SkyflowError(error_message, invalid_input_error_code) + + validate_token_options(logger, credentials, config_id_type, config_id) + if CredentialField.CREDENTIALS_STRING in credentials: validate_required_field( logger, credentials, CredentialField.CREDENTIALS_STRING, str, diff --git a/skyflow/vault/client/client.py b/skyflow/vault/client/client.py index e7255cb1..5872d7d7 100644 --- a/skyflow/vault/client/client.py +++ b/skyflow/vault/client/client.py @@ -4,6 +4,7 @@ from skyflow.utils import get_vault_url, get_credentials, SkyflowMessages from skyflow.utils.logger import log_info from skyflow.utils.constants import OptionField, CredentialField, ConfigField +from skyflow.utils.validations import validate_token_options class VaultClient: @@ -74,14 +75,14 @@ def get_bearer_token(self, credentials): elif CredentialField.TOKEN in credentials: return credentials.get(CredentialField.TOKEN) + validate_token_options(self.__logger, credentials) + options = {} - roles = credentials.get(CredentialField.ROLES) - if roles: - options[OptionField.ROLE_IDS] = roles + if CredentialField.ROLES in credentials: + options[OptionField.ROLE_IDS] = credentials.get(CredentialField.ROLES) - context = credentials.get(CredentialField.CONTEXT) - if context is not None: - options[OptionField.CTX] = context + if CredentialField.CONTEXT in credentials: + options[OptionField.CTX] = credentials.get(CredentialField.CONTEXT) if CredentialField.TOKEN_URI_OPTION in credentials and credentials.get(CredentialField.TOKEN_URI_OPTION): options[CredentialField.TOKEN_URI_OPTION] = credentials.get(CredentialField.TOKEN_URI_OPTION) diff --git a/tests/utils/validations/test__validations.py b/tests/utils/validations/test__validations.py index c9573084..c1ba76bb 100644 --- a/tests/utils/validations/test__validations.py +++ b/tests/utils/validations/test__validations.py @@ -13,7 +13,7 @@ validate_get_detect_run_request, validate_get_request, validate_update_request, validate_detokenize_request, validate_tokenize_request, validate_invoke_connection_params, validate_deidentify_text_request, validate_reidentify_text_request, validate_deidentify_file_request, - validate_file_upload_request + validate_file_upload_request, validate_token_options ) from skyflow.utils import SkyflowMessages from skyflow.utils.enums import DetectEntities, RedactionType @@ -209,20 +209,6 @@ def test_validate_credentials_with_dict_context_in_config(self): SkyflowMessages.Error.EMPTY_CONTEXT_IN_CONFIG.value.format("vault", "vault123") ) - def test_validate_vault_config_with_dict_context_and_roles(self): - from skyflow.utils.enums import Env - config = { - "vault_id": "vault123", - "cluster_id": "cluster123", - "credentials": { - "credentials_string": '{"clientID": "x"}', - "roles": ["role_id_1", "role_id_2"], - "context": {"role": "admin", "department": "finance"} - }, - "env": Env.DEV - } - self.assertTrue(validate_vault_config(self.logger, config)) - # ------------------------------------------------------------------ # # credentials.roles # # ------------------------------------------------------------------ # @@ -252,6 +238,21 @@ def test_validate_credentials_with_empty_roles(self): validate_credentials(self.logger, credentials) self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_ROLES.value) + # ------------------------------------------------------------------ # + # validate_token_options — shared by config validation and VaultClient # + # ------------------------------------------------------------------ # + + def test_validate_token_options_rejects_empty_values(self): + """The contract VaultClient relies on; the cases themselves are covered above + via validate_credentials, which delegates here.""" + with self.assertRaises(SkyflowError) as context: + validate_token_options(self.logger, {"roles": []}) + self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_ROLES.value) + + with self.assertRaises(SkyflowError) as context: + validate_token_options(self.logger, {"context": {}}) + self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_CONTEXT.value) + def test_validate_log_level_valid(self): from skyflow.utils.enums import LogLevel log_level = LogLevel.ERROR diff --git a/tests/vault/client/test__client.py b/tests/vault/client/test__client.py index 1eb2435e..bc9931ae 100644 --- a/tests/vault/client/test__client.py +++ b/tests/vault/client/test__client.py @@ -353,6 +353,25 @@ def test_get_bearer_token_omits_options_when_roles_and_context_absent(self, mock self.assertNotIn("role_ids", options) self.assertNotIn("ctx", options) + @patch("skyflow.vault.client.client.generate_bearer_token") + def test_get_bearer_token_rejects_empty_roles(self, mock_generate): + """An empty roles list must not silently produce an unscoped token.""" + credentials = {**CREDENTIALS_WITH_PATH, "roles": []} + with self.assertRaises(SkyflowError) as ctx: + self.vault_client.get_bearer_token(credentials) + self.assertEqual(ctx.exception.message, SkyflowMessages.Error.EMPTY_ROLES.value) + mock_generate.assert_not_called() + + @patch("skyflow.vault.client.client.generate_bearer_token") + def test_get_bearer_token_rejects_empty_context(self, mock_generate): + """An empty context must not silently produce a context-less token.""" + for empty_context in [{}, "", " "]: + credentials = {**CREDENTIALS_WITH_PATH, "context": empty_context} + with self.assertRaises(SkyflowError) as ctx: + self.vault_client.get_bearer_token(credentials) + self.assertEqual(ctx.exception.message, SkyflowMessages.Error.EMPTY_CONTEXT.value) + mock_generate.assert_not_called() + @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("sa_token", None)) def test_get_bearer_token_ignores_top_level_config_roles_and_ctx(self, mock_generate): """roles/ctx are only read from credentials — never from the top-level config.""" @@ -376,19 +395,6 @@ def test_get_bearer_token_forwards_token_uri_alongside_roles_and_context(self, m self.assertEqual(options["ctx"], STRING_CONTEXT) self.assertEqual(options["token_uri"], credentials["token_uri"]) - @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("new_token", None)) - @patch("skyflow.vault.client.client.is_expired", return_value=True) - @patch("skyflow.vault.client.client.log_info") - def test_get_bearer_token_forwards_roles_and_context_on_refresh(self, mock_log, mock_is_expired, mock_generate): - """Auto-refresh of an expired token must carry roles/context too, not just the first call.""" - self.vault_client._VaultClient__bearer_token = "expired_token" - result = self.vault_client.get_bearer_token(CREDENTIALS_WITH_PATH_ROLES_AND_STRING_CONTEXT) - - self.assertEqual(result, "new_token") - _, options, _ = mock_generate.call_args[0] - self.assertEqual(options["role_ids"], ROLES) - self.assertEqual(options["ctx"], STRING_CONTEXT) - @patch("skyflow.vault.client.client.generate_bearer_token", return_value=("refreshed_token", None)) @patch("skyflow.vault.client.client.is_expired", return_value=True) @patch("skyflow.vault.client.client.get_credentials") From 501aa2f27647231caf0f355169ba54e6e6f454ac Mon Sep 17 00:00:00 2001 From: yaswanth-pula-skyflow Date: Mon, 3 Aug 2026 22:41:56 +0530 Subject: [PATCH 3/7] SK-3039:Add delete interface SkyflowID input validation. --- skyflow/utils/validations/_validations.py | 4 ++++ tests/utils/validations/test__validations.py | 13 +++++++++++++ 2 files changed, 17 insertions(+) diff --git a/skyflow/utils/validations/_validations.py b/skyflow/utils/validations/_validations.py index 3e53bfc4..2a76c2bf 100644 --- a/skyflow/utils/validations/_validations.py +++ b/skyflow/utils/validations/_validations.py @@ -561,6 +561,10 @@ def validate_delete_request(logger, request): log_error_log(SkyflowMessages.ErrorLogs.EMPTY_IDS.value.format(RequestOperation.DELETE), logger=logger) raise SkyflowError(SkyflowMessages.Error.EMPTY_RECORD_IDS_IN_DELETE.value, invalid_input_error_code) + if not isinstance(request.ids, list): + log_error_log(SkyflowMessages.ErrorLogs.EMPTY_IDS.value.format(RequestOperation.DELETE), logger=logger) + raise SkyflowError(SkyflowMessages.Error.INVALID_IDS_TYPE.value.format(type(request.ids)), invalid_input_error_code) + def validate_query_request(logger, request): if not isinstance(request.query, str): query_type = str(type(request.query)) diff --git a/tests/utils/validations/test__validations.py b/tests/utils/validations/test__validations.py index c1ba76bb..c680201b 100644 --- a/tests/utils/validations/test__validations.py +++ b/tests/utils/validations/test__validations.py @@ -506,6 +506,19 @@ def test_validate_delete_request_missing_ids(self): validate_delete_request(self.logger, request) self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_RECORD_IDS_IN_DELETE.value) + def test_validate_delete_request_non_list_ids(self): + for ids in ["id1", 123, {"id": "id1"}, ("id1", "id2")]: + with self.subTest(ids=ids): + request = MagicMock() + request.table = "test_table" + request.ids = ids + with self.assertRaises(SkyflowError) as context: + validate_delete_request(self.logger, request) + self.assertEqual( + context.exception.message, + SkyflowMessages.Error.INVALID_IDS_TYPE.value.format(type(ids)) + ) + def test_validate_query_request_valid(self): request = MagicMock() request.query = "SELECT * FROM test_table" From f03510f50817294419efb6cc346685893538ada0 Mon Sep 17 00:00:00 2001 From: yaswanth-pula-skyflow Date: Mon, 3 Aug 2026 23:08:50 +0530 Subject: [PATCH 4/7] SK-3039:Update fileUpload & update interfaces input validations. --- skyflow/utils/_skyflow_messages.py | 1 + skyflow/utils/validations/_validations.py | 12 ++++++- tests/utils/validations/test__validations.py | 37 ++++++++++++++++++++ 3 files changed, 49 insertions(+), 1 deletion(-) diff --git a/skyflow/utils/_skyflow_messages.py b/skyflow/utils/_skyflow_messages.py index 232bd8b0..ee4ebd54 100644 --- a/skyflow/utils/_skyflow_messages.py +++ b/skyflow/utils/_skyflow_messages.py @@ -109,6 +109,7 @@ class Error(Enum): EMPTY_RECORD_IDS_IN_DELETE = f"{error_prefix} Validation error. 'record ids' array can't be empty. Specify one or more record ids." BULK_DELETE_FAILURE = f"{error_prefix} Delete operation failed." EMPTY_SKYFLOW_ID= f"{error_prefix} Validation error. skyflow_id can't be empty." + INVALID_SKYFLOW_ID_TYPE = f"{error_prefix} Validation error. 'skyflow_id' has a value of type {{}}. Specify 'skyflow_id' as a string." INVALID_FILE_COLUMN_NAME= f"{error_prefix} Validation error. 'column_name' can't be empty." INVALID_QUERY_TYPE = f"{error_prefix} Validation error. Query parameter is of type {{}}. Specify as a string." diff --git a/skyflow/utils/validations/_validations.py b/skyflow/utils/validations/_validations.py index 2a76c2bf..43d8bff4 100644 --- a/skyflow/utils/validations/_validations.py +++ b/skyflow/utils/validations/_validations.py @@ -677,6 +677,8 @@ def validate_update_request(logger, request): skyflow_id = request.data.get(ResponseField.SKYFLOW_ID) if skyflow_id is None: log_error_log(SkyflowMessages.ErrorLogs.SKYFLOW_ID_IS_REQUIRED.value.format(RequestOperation.UPDATE), logger=logger) + elif not isinstance(skyflow_id, str): + raise SkyflowError(SkyflowMessages.Error.INVALID_SKYFLOW_ID_TYPE.value.format(type(skyflow_id)), invalid_input_error_code) elif not skyflow_id.strip(): log_error_log(SkyflowMessages.ErrorLogs.EMPTY_SKYFLOW_ID.value.format(RequestOperation.UPDATE), logger=logger) @@ -793,11 +795,16 @@ def validate_file_upload_request(logger, request): table = getattr(request, FileUploadField.TABLE, None) if table is None: raise SkyflowError(SkyflowMessages.Error.INVALID_TABLE_VALUE.value, invalid_input_error_code) + elif not isinstance(table, str): + log_error_log(SkyflowMessages.ErrorLogs.TABLE_IS_REQUIRED.value.format(RequestOperation.FILE_UPLOAD), logger=logger) + raise SkyflowError(SkyflowMessages.Error.INVALID_TABLE_VALUE.value, invalid_input_error_code) elif table.strip() == "": raise SkyflowError(SkyflowMessages.Error.EMPTY_TABLE_VALUE.value, invalid_input_error_code) # Skyflow ID skyflow_id = getattr(request, FileUploadField.SKYFLOW_ID, None) + if skyflow_id is not None and not isinstance(skyflow_id, str): + raise SkyflowError(SkyflowMessages.Error.INVALID_SKYFLOW_ID_TYPE.value.format(type(skyflow_id)), invalid_input_error_code) if skyflow_id is not None and skyflow_id.strip() == "": raise SkyflowError(SkyflowMessages.Error.EMPTY_SKYFLOW_ID.value.format(RequestOperation.FILE_UPLOAD), invalid_input_error_code) @@ -805,6 +812,9 @@ def validate_file_upload_request(logger, request): column_name = getattr(request, FileUploadField.COLUMN_NAME, None) if column_name is None: raise SkyflowError(SkyflowMessages.Error.INVALID_FILE_COLUMN_NAME.value.format(type(column_name)), invalid_input_error_code) + elif not isinstance(column_name, str): + log_error_log(SkyflowMessages.ErrorLogs.EMPTY_FILE_COLUMN_NAME.value, logger) + raise SkyflowError(SkyflowMessages.Error.INVALID_COLUMN_NAME.value.format(type(column_name)), invalid_input_error_code) elif column_name.strip() == "": log_error_log(SkyflowMessages.ErrorLogs.EMPTY_FILE_COLUMN_NAME.value, logger) raise SkyflowError(SkyflowMessages.Error.INVALID_FILE_COLUMN_NAME.value.format(type(column_name)), invalid_input_error_code) @@ -823,7 +833,7 @@ def validate_file_upload_request(logger, request): # Check base64 if present if not is_none_or_empty(base64_str): - if is_none_or_empty(file_name): + if not isinstance(file_name, str) or is_none_or_empty(file_name): raise SkyflowError(SkyflowMessages.Error.INVALID_FILE_NAME.value, invalid_input_error_code) try: base64.b64decode(base64_str) diff --git a/tests/utils/validations/test__validations.py b/tests/utils/validations/test__validations.py index c680201b..47ca531b 100644 --- a/tests/utils/validations/test__validations.py +++ b/tests/utils/validations/test__validations.py @@ -684,6 +684,20 @@ def test_validate_update_request_invalid_table_type(self): validate_update_request(self.logger, request) self.assertEqual(context.exception.message, SkyflowMessages.Error.INVALID_TABLE_VALUE.value) + def test_validate_update_request_non_string_skyflow_id(self): + for skyflow_id in [123, ["id123"], {"id": "id123"}]: + with self.subTest(skyflow_id=skyflow_id): + request = UpdateRequest( + table="test_table", + data={"skyflow_id": skyflow_id, "field1": "value1"} + ) + with self.assertRaises(SkyflowError) as context: + validate_update_request(self.logger, request) + self.assertEqual( + context.exception.message, + SkyflowMessages.Error.INVALID_SKYFLOW_ID_TYPE.value.format(type(skyflow_id)) + ) + def test_validate_update_request_invalid_token_mode(self): request = UpdateRequest( table="test_table", @@ -1588,6 +1602,29 @@ def test_validate_file_upload_request_missing_file_source(self): validate_file_upload_request(self.logger, request) self.assertEqual(context.exception.message, SkyflowMessages.Error.MISSING_FILE_SOURCE.value) + def test_validate_file_upload_request_non_string_fields(self): + cases = [ + ("table", 123, SkyflowMessages.Error.INVALID_TABLE_VALUE.value), + ("table", ["test_table"], SkyflowMessages.Error.INVALID_TABLE_VALUE.value), + ("column_name", 123, SkyflowMessages.Error.INVALID_COLUMN_NAME.value.format(type(123))), + ("column_name", ["file_col"], SkyflowMessages.Error.INVALID_COLUMN_NAME.value.format(type([]))), + ("skyflow_id", 123, SkyflowMessages.Error.INVALID_SKYFLOW_ID_TYPE.value.format(type(123))), + ("file_name", 123, SkyflowMessages.Error.INVALID_FILE_NAME.value), + ] + for field, value, expected_message in cases: + with self.subTest(field=field, value=value): + kwargs = { + "table": "test_table", + "column_name": "file_col", + "base64": "ZmlsZSBjb250ZW50", + "file_name": "sample.txt", + field: value, + } + request = FileUploadRequest(**kwargs) + with self.assertRaises(SkyflowError) as context: + validate_file_upload_request(self.logger, request) + self.assertEqual(context.exception.message, expected_message) + # --- validate_deidentify_text_request transformations --- def test_validate_deidentify_text_request_invalid_transformations(self): From a9da74e0680e359407cc0221a0c56277814a3ea7 Mon Sep 17 00:00:00 2001 From: yaswanth-pula-skyflow Date: Tue, 4 Aug 2026 11:15:30 +0530 Subject: [PATCH 5/7] SK-3039:Update typesafe validations. --- skyflow/utils/_skyflow_messages.py | 2 + skyflow/utils/validations/_validations.py | 21 ++++++--- tests/utils/validations/test__validations.py | 45 +++++++++++++++++++- 3 files changed, 61 insertions(+), 7 deletions(-) diff --git a/skyflow/utils/_skyflow_messages.py b/skyflow/utils/_skyflow_messages.py index ee4ebd54..5182cd53 100644 --- a/skyflow/utils/_skyflow_messages.py +++ b/skyflow/utils/_skyflow_messages.py @@ -58,6 +58,8 @@ class Error(Enum): INVALID_ROLES_KEY_TYPE = f"{error_prefix} Validation error. Invalid roles. Specify roles as an array." EMPTY_ROLES_IN_CONFIG = f"{error_prefix} Validation error. Invalid roles for {{}} with id {{}}. Specify at least one role." EMPTY_ROLES = f"{error_prefix} Validation error. Invalid roles. Specify at least one role." + INVALID_ROLE_ELEMENT_TYPE_IN_CONFIG = f"{error_prefix} Validation error. Invalid roles for {{}} with id {{}}. Each role must be a non-empty string." + INVALID_ROLE_ELEMENT_TYPE = f"{error_prefix} Validation error. Invalid roles. Each role must be a non-empty string." EMPTY_CONTEXT_IN_CONFIG = f"{error_prefix} Initialization failed. Invalid context provided for {{}} with id {{}}. Specify context as type Context." EMPTY_CONTEXT = f"{error_prefix} Initialization failed. Invalid context provided. Specify context as type Context." INVALID_CONTEXT_IN_CONFIG = f"{error_prefix} Initialization failed. Invalid context for {{}} with id {{}}. Specify a valid context." diff --git a/skyflow/utils/validations/_validations.py b/skyflow/utils/validations/_validations.py index 43d8bff4..bc2aa1d8 100644 --- a/skyflow/utils/validations/_validations.py +++ b/skyflow/utils/validations/_validations.py @@ -102,6 +102,14 @@ def validate_token_options(logger, credentials, config_id_type=None, config_id=N if not credentials.get(CredentialField.ROLES): raise SkyflowError(empty_roles_error, invalid_input_error_code) + invalid_role_element_error = ( + SkyflowMessages.Error.INVALID_ROLE_ELEMENT_TYPE_IN_CONFIG.value.format(config_id_type, config_id) + if config_id_type and config_id else SkyflowMessages.Error.INVALID_ROLE_ELEMENT_TYPE.value + ) + for role in credentials.get(CredentialField.ROLES): + if not isinstance(role, str) or not role.strip(): + raise SkyflowError(invalid_role_element_error, invalid_input_error_code) + if CredentialField.CONTEXT in credentials: empty_context_error = ( SkyflowMessages.Error.EMPTY_CONTEXT_IN_CONFIG.value.format(config_id_type, config_id) @@ -310,7 +318,7 @@ def validate_update_connection_config(logger, config): if ConfigField.CREDENTIALS not in config: raise SkyflowError(SkyflowMessages.Error.EMPTY_CREDENTIALS.value.format(ConfigType.CONNECTION, connection_id), invalid_input_error_code) - validate_credentials(logger, config.get(ConfigField.CREDENTIALS)) + validate_credentials(logger, config.get(ConfigField.CREDENTIALS), ConfigType.CONNECTION, connection_id) return True @@ -519,6 +527,11 @@ def validate_insert_request(logger, request): raise SkyflowError(SkyflowMessages.Error.INVALID_CONTINUE_ON_ERROR_TYPE.value, invalid_input_error_code) if request.tokens: + if not isinstance(request.tokens, list) or not request.tokens or not all( + isinstance(t, dict) for t in request.tokens): + log_error_log(SkyflowMessages.ErrorLogs.EMPTY_TOKENS.value.format(RequestOperation.INSERT), logger=logger) + raise SkyflowError(SkyflowMessages.Error.INVALID_TYPE_OF_DATA_IN_INSERT.value, invalid_input_error_code) + for i, item in enumerate(request.tokens, start=1): for key, value in item.items(): if key is None or key == "": @@ -528,10 +541,6 @@ def validate_insert_request(logger, request): if value is None or value == "": log_error_log(SkyflowMessages.ErrorLogs.EMPTY_OR_NULL_KEY_IN_TOKENS.value.format(RequestOperation.INSERT, key), logger=logger) - if not isinstance(request.tokens, list) or not request.tokens or not all( - isinstance(t, dict) for t in request.tokens): - log_error_log(SkyflowMessages.ErrorLogs.EMPTY_TOKENS.value.format(RequestOperation.INSERT), logger=logger) - raise SkyflowError(SkyflowMessages.Error.INVALID_TYPE_OF_DATA_IN_INSERT.value, invalid_input_error_code) if request.token_mode == TokenMode.ENABLE and not request.tokens: raise SkyflowError(SkyflowMessages.Error.NO_TOKENS_IN_INSERT.value.format(request.token_mode), invalid_input_error_code) @@ -737,7 +746,7 @@ def validate_detokenize_request(logger, request): raise SkyflowError(SkyflowMessages.Error.EMPTY_TOKENS_LIST_VALUE.value, invalid_input_error_code) for item in request.data: - if ResponseField.TOKEN not in item: + if not isinstance(item, dict) or ResponseField.TOKEN not in item: raise SkyflowError(SkyflowMessages.Error.INVALID_TOKENS_LIST_VALUE.value.format(type(request.data)), invalid_input_error_code) diff --git a/tests/utils/validations/test__validations.py b/tests/utils/validations/test__validations.py index 47ca531b..096c4d5e 100644 --- a/tests/utils/validations/test__validations.py +++ b/tests/utils/validations/test__validations.py @@ -238,6 +238,15 @@ def test_validate_credentials_with_empty_roles(self): validate_credentials(self.logger, credentials) self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_ROLES.value) + def test_validate_credentials_with_non_string_role_elements(self): + credentials = { + "api_key": "sky-abc12-1234567890abcdef1234567890abcdef", + "roles": [None, 123, {"id": "x"}] + } + with self.assertRaises(SkyflowError) as context: + validate_credentials(self.logger, credentials) + self.assertEqual(context.exception.message, SkyflowMessages.Error.INVALID_ROLE_ELEMENT_TYPE.value) + # ------------------------------------------------------------------ # # validate_token_options — shared by config validation and VaultClient # # ------------------------------------------------------------------ # @@ -414,6 +423,17 @@ def test_validate_update_connection_config_empty_url(self): validate_update_connection_config(self.logger, config) self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_CONNECTION_URL.value.format("conn123")) + def test_validate_update_connection_config_invalid_credentials_includes_connection_id(self): + config = { + "connection_id": "conn123", + "connection_url": "https://example.com", + "credentials": {} + } + with self.assertRaises(SkyflowError) as context: + validate_update_connection_config(self.logger, config) + self.assertEqual(context.exception.message, + SkyflowMessages.Error.INVALID_CREDENTIALS_IN_CONFIG.value.format("connection", "conn123")) + def test_validate_file_from_request_valid_file(self): file_obj = MagicMock() file_obj.name = "test.txt" @@ -483,6 +503,20 @@ def test_validate_insert_request_empty_values(self): validate_insert_request(self.logger, request) self.assertEqual(context.exception.message, SkyflowMessages.Error.EMPTY_DATA_IN_INSERT.value) + def test_validate_insert_request_tokens_as_dict(self): + request = MagicMock() + request.table = "test_table" + request.values = [{"field1": "value1"}] + request.upsert = None + request.homogeneous = None + request.token_mode = None + request.return_tokens = False + request.continue_on_error = False + request.tokens = {"field1": "token1"} + with self.assertRaises(SkyflowError) as context: + validate_insert_request(self.logger, request) + self.assertEqual(context.exception.message, SkyflowMessages.Error.INVALID_TYPE_OF_DATA_IN_INSERT.value) + def test_validate_delete_request_valid(self): request = MagicMock() @@ -728,9 +762,18 @@ def test_validate_detokenize_request_invalid_token(self): request.continue_on_error = False with self.assertRaises(SkyflowError) as context: validate_detokenize_request(self.logger, request) - self.assertEqual(context.exception.message, + self.assertEqual(context.exception.message, SkyflowMessages.Error.INVALID_TOKEN_TYPE.value.format("DETOKENIZE")) + def test_validate_detokenize_request_bare_string_data(self): + request = MagicMock() + request.data = ["mytoken123"] # bare string containing 'token' as a substring + request.continue_on_error = False + with self.assertRaises(SkyflowError) as context: + validate_detokenize_request(self.logger, request) + self.assertEqual(context.exception.message, + SkyflowMessages.Error.INVALID_TOKENS_LIST_VALUE.value.format(list)) + def test_validate_tokenize_request_valid(self): request = MagicMock() request.values = [{"value": "test", "column_group": "group1"}] From 6d65838874333373af6f0ee5e961c2d813d90ac8 Mon Sep 17 00:00:00 2001 From: yaswanth-pula-skyflow Date: Tue, 4 Aug 2026 12:58:18 +0530 Subject: [PATCH 6/7] SK-3039:Fix unit test assertion gaps. --- tests/vault/controller/_request_assertions.py | 43 +++++ tests/vault/controller/test__connection.py | 52 +++++- tests/vault/controller/test__detect.py | 135 ++++++++++++-- tests/vault/controller/test__vault.py | 176 +++++++++++++++++- 4 files changed, 387 insertions(+), 19 deletions(-) create mode 100644 tests/vault/controller/_request_assertions.py diff --git a/tests/vault/controller/_request_assertions.py b/tests/vault/controller/_request_assertions.py new file mode 100644 index 00000000..d8989bab --- /dev/null +++ b/tests/vault/controller/_request_assertions.py @@ -0,0 +1,43 @@ +"""Shared assertions for what the controllers actually send to the generated REST client. + +The controller tests mock the API layer, so `assert_called_once()` on its own only proves the +controller reached the API — never that it built the right request. These helpers close that +gap: a wrong table name, a dropped field or a flipped flag has to fail a test. +""" +import json + +from skyflow.utils._version import SDK_VERSION +from skyflow.utils.constants import SKY_META_DATA_HEADER, SdkMetricsKey, SdkPrefix + +# Identical on every call, and covered once by assert_sky_metadata instead of in every test. +IGNORED_KWARGS = ("request_options",) + + +def assert_request(test_case, api_mock, expected_kwargs, expected_args=None): + """Assert api_mock was called exactly once, with exactly this request. + + Compares the whole kwargs dict rather than a subset, so a dropped or renamed parameter + fails too - not just a wrong value. + """ + api_mock.assert_called_once() + args, kwargs = api_mock.call_args + if expected_args is not None: + test_case.assertEqual(args, expected_args) + actual_kwargs = {key: value for key, value in kwargs.items() if key not in IGNORED_KWARGS} + test_case.assertEqual(actual_kwargs, expected_kwargs) + + +def assert_sky_metadata(test_case, additional_headers): + """Assert a sky-metadata header is present and names this SDK version.""" + test_case.assertIn(SKY_META_DATA_HEADER, additional_headers) + metrics = json.loads(additional_headers[SKY_META_DATA_HEADER]) + test_case.assertEqual( + metrics[SdkMetricsKey.SDK_NAME_VERSION], + SdkPrefix.SKYFLOW_PYTHON + SDK_VERSION + ) + + +def request_options_headers(api_mock): + """The additional_headers the controller passed in request_options.""" + _, kwargs = api_mock.call_args + return kwargs["request_options"]["additional_headers"] diff --git a/tests/vault/controller/test__connection.py b/tests/vault/controller/test__connection.py index f073264c..b3268534 100644 --- a/tests/vault/controller/test__connection.py +++ b/tests/vault/controller/test__connection.py @@ -6,6 +6,7 @@ from skyflow.utils import SkyflowMessages, parse_invoke_connection_response from skyflow.utils._utils import get_data_from_content_type, construct_invoke_connection_request from skyflow.utils.enums import RequestMethod, ContentType +from skyflow.utils.constants import SKY_META_DATA_HEADER from skyflow.utils._version import SDK_VERSION from skyflow.vault.connection import InvokeConnectionRequest from skyflow.vault.controller import Connection @@ -70,6 +71,13 @@ def test_invoke_success(self, mock_send, mock_get_credentials): self.mock_vault_client.get_bearer_token.assert_called_once() mock_get_credentials.assert_called_once() + # Assert the request actually sent - method, url with query params, body, content type + sent_request = mock_send.call_args[0][0] + self.assertEqual(sent_request.method, "POST") + self.assertEqual(sent_request.url, "https://connection_url/?query_key=value") + self.assertEqual(sent_request.body, json.dumps(VALID_BODY)) + self.assertEqual(sent_request.headers['content-type'], "application/json") + @patch('skyflow.vault.controller._connections.get_credentials') @patch('requests.Session.send') def test_invoke_with_x_skyflow_authorization_already_present(self, mock_send, mock_get_credentials): @@ -83,16 +91,46 @@ def test_invoke_with_x_skyflow_authorization_already_present(self, mock_send, mo mock_send.return_value = mock_response custom_auth = "custom_bearer_token" + # The guard compares against a lowercase literal, so any casing must be honoured + for header_key in ["x-skyflow-authorization", "X-Skyflow-Authorization", "X-SKYFLOW-AUTHORIZATION"]: + with self.subTest(header_key=header_key): + mock_send.reset_mock() + request = InvokeConnectionRequest( + method=RequestMethod.POST, + body=VALID_BODY, + headers={header_key: custom_auth} + ) + + response = self.connection.invoke(request) + + # Verify bearer token from vault_client is NOT used + self.assertIsNotNone(response) + sent_headers = mock_send.call_args[0][0].headers + self.assertEqual(sent_headers['x-skyflow-authorization'], custom_auth) + self.assertNotIn(VALID_BEARER_TOKEN, sent_headers.values()) + + @patch('skyflow.vault.controller._connections.get_credentials') + @patch('requests.Session.send') + def test_invoke_injects_bearer_token_when_header_absent(self, mock_send, mock_get_credentials): + """Test that the SDK-generated bearer token is used when the caller supplies none.""" + mock_get_credentials.return_value = {"api_key": "test_api_key"} + + mock_response = Mock() + mock_response.status_code = SUCCESS_STATUS_CODE + mock_response.content = SUCCESS_RESPONSE_CONTENT + mock_response.headers = {'x-request-id': 'test-request-id'} + mock_send.return_value = mock_response + request = InvokeConnectionRequest( method=RequestMethod.POST, body=VALID_BODY, - headers={"x-skyflow-authorization": custom_auth} + headers=VALID_HEADERS ) - response = self.connection.invoke(request) - - # Verify bearer token from vault_client is NOT used - self.assertIsNotNone(response) + self.connection.invoke(request) + + sent_headers = mock_send.call_args[0][0].headers + self.assertEqual(sent_headers['x-skyflow-authorization'], VALID_BEARER_TOKEN) @patch('skyflow.vault.controller._connections.get_credentials') def test_invoke_invalid_headers(self, mock_get_credentials): @@ -239,9 +277,11 @@ def test_invoke_adds_sky_metadata_header(self, mock_send, mock_get_metrics, mock response = self.connection.invoke(request) - # Verify get_metrics was called + # Verify get_metrics was called and its payload is on the sent request mock_get_metrics.assert_called_once() self.assertIsNotNone(response) + sent_headers = mock_send.call_args[0][0].headers + self.assertEqual(json.loads(sent_headers[SKY_META_DATA_HEADER]), {"sdk_version": SDK_VERSION}) def test_parse_invoke_connection_response_error_from_client(self): mock_response = Mock(spec=requests.Response) diff --git a/tests/vault/controller/test__detect.py b/tests/vault/controller/test__detect.py index f0f2aa87..9569961f 100644 --- a/tests/vault/controller/test__detect.py +++ b/tests/vault/controller/test__detect.py @@ -4,16 +4,17 @@ import os import tempfile from skyflow.error import SkyflowError -from skyflow.generated.rest import WordCharacterCount +from skyflow.generated.rest import WordCharacterCount, FileDataDeidentifyText, FileDataDeidentifyAudio, Format from skyflow.utils import SkyflowMessages from skyflow.vault.controller import Detect from skyflow.vault.detect import DeidentifyTextRequest, ReidentifyTextRequest, \ TokenFormat, DateTransformation, Transformations, DeidentifyFileRequest, GetDetectRunRequest, \ - DeidentifyFileResponse, FileInput + DeidentifyFileResponse, FileInput, Bleep from skyflow.utils.enums import DetectEntities, TokenType import io from skyflow.vault.detect._file import File +from tests.vault.controller._request_assertions import assert_request VAULT_ID = "test_vault_id" @@ -65,7 +66,26 @@ def test_deidentify_text_success(self, mock_parse_response, mock_validate): # Assertions mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response) - detect_api.deidentify_string.assert_called_once() + + # Assert the text, entity list and token format reach the API + assert_request( + self, + detect_api.deidentify_string, + { + "vault_id": VAULT_ID, + "text": "John lives in NYC", + "entity_types": [DetectEntities.NAME], + "token_type": { + "default": TokenType.ENTITY_ONLY, + "entity_unq_counter": None, + "entity_only": None, + "vault_token": None + }, + "allow_regex": None, + "restrict_regex": None, + "transformations": None + } + ) @patch("skyflow.vault.controller._detect.validate_reidentify_text_request") @patch("skyflow.vault.controller._detect.parse_reidentify_text_response") @@ -76,10 +96,12 @@ def test_reidentify_text_success(self, mock_parse_response, mock_validate): 'text': 'John lives in NYC' } - # Create request + # Create request - all three re-identification formats populated request = ReidentifyTextRequest( text="Token1 lives in Token2", - redacted_entities=[DetectEntities.NAME] + redacted_entities=[DetectEntities.NAME], + masked_entities=[DetectEntities.LOCATION], + plain_text_entities=[DetectEntities.EMAIL_ADDRESS] ) # Mock detect API @@ -92,7 +114,21 @@ def test_reidentify_text_success(self, mock_parse_response, mock_validate): # Assertions mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response) - detect_api.reidentify_string.assert_called_once() + + # Assert each entity list lands on the right re-identification format + assert_request( + self, + detect_api.reidentify_string, + { + "vault_id": VAULT_ID, + "text": "Token1 lives in Token2", + "format": Format( + redacted=[DetectEntities.NAME], + masked=[DetectEntities.LOCATION], + plaintext=[DetectEntities.EMAIL_ADDRESS] + ) + } + ) @patch("skyflow.vault.controller._detect.validate_deidentify_text_request") def test_deidentify_text_handles_generic_error(self, mock_validate): @@ -165,10 +201,29 @@ def test_deidentify_file_txt_success(self, mock_open, mock_basename, mock_base64 result = self.detect.deidentify_file(req) mock_validate.assert_called_once() - files_api.deidentify_text.assert_called_once() mock_poll.assert_called_once() mock_parse.assert_called_once() + # Assert the encoded file and detect options reach the txt endpoint + assert_request( + self, + files_api.deidentify_text, + { + "vault_id": VAULT_ID, + "file": FileDataDeidentifyText(base_64="dGVzdCBjb250ZW50", data_format="txt"), + "entity_types": [], + "token_type": { + "default": req.token_format.default, + "entity_unq_counter": req.token_format.entity_unique_counter, + "entity_only": req.token_format.entity_only, + "vault_token": req.token_format.vault_token + }, + "allow_regex": [], + "restrict_regex": [], + "transformations": None + } + ) + self.assertIsInstance(result, DeidentifyFileResponse) self.assertEqual(result.status, "SUCCESS") self.assertEqual(result.run_id, "runid123") @@ -195,7 +250,12 @@ def test_deidentify_file_audio_success(self, mock_base64, mock_validate): file_obj.read.return_value = file_content file_obj.name = "audio.mp3" mock_base64.b64encode.return_value = b"YXVkaW8gYnl0ZXM=" - req = DeidentifyFileRequest(file=FileInput(file=file_obj)) + bleep = Bleep(gain=0.5, frequency=800.0, start_padding=0.1, stop_padding=0.2) + req = DeidentifyFileRequest( + file=FileInput(file=file_obj), + bleep=bleep, + output_processed_audio=True + ) req.entities = [] req.token_format = Mock(default="default", entity_unique_counter=[], entity_only=[]) req.allow_regex_list = [] @@ -226,9 +286,34 @@ def test_deidentify_file_audio_success(self, mock_base64, mock_validate): status="SUCCESS")) as mock_parse: result = self.detect.deidentify_file(req) mock_validate.assert_called_once() - files_api.deidentify_audio.assert_called_once() mock_poll.assert_called_once() mock_parse.assert_called_once() + + # Assert the audio-only options - bleep settings and transcription flags - are sent + assert_request( + self, + files_api.deidentify_audio, + { + "vault_id": VAULT_ID, + "file": FileDataDeidentifyAudio(base_64="YXVkaW8gYnl0ZXM=", data_format="mp3"), + "entity_types": [], + "token_type": { + "default": req.token_format.default, + "entity_unq_counter": req.token_format.entity_unique_counter, + "entity_only": req.token_format.entity_only, + "vault_token": req.token_format.vault_token + }, + "allow_regex": [], + "restrict_regex": [], + "transformations": None, + "output_transcription": None, + "output_processed_audio": True, + "bleep_gain": 0.5, + "bleep_frequency": 800.0, + "bleep_start_padding": 0.1, + "bleep_stop_padding": 0.2 + } + ) self.assertIsInstance(result, DeidentifyFileResponse) self.assertEqual(result.status, "SUCCESS") @@ -267,8 +352,15 @@ def test_get_detect_run_success(self, mock_validate): run_id="runid789", status="SUCCESS")) as mock_parse: result = self.detect.get_detect_run(req) mock_validate.assert_called_once() - files_api.get_run.assert_called_once() mock_parse.assert_called_once() + + # Assert the run id is the one looked up + assert_request( + self, + files_api.get_run, + {"vault_id": VAULT_ID}, + expected_args=("runid789",) + ) self.assertIsInstance(result, DeidentifyFileResponse) self.assertEqual(result.status, "SUCCESS") @@ -685,8 +777,29 @@ def test_deidentify_file_using_file_path(self, mock_open, mock_basename, mock_ba mock_file.read.assert_called_once() mock_validate.assert_called_once() - files_api.deidentify_text.assert_called_once() mock_basename.assert_called_with("/path/to/test.txt") + # The file_path branch encodes the file read off disk, same shape as file_object + assert_request( + self, + files_api.deidentify_text, + { + "vault_id": VAULT_ID, + "file": FileDataDeidentifyText( + base_64="dGVzdCBjb250ZW50IGZyb20gZmlsZSBwYXRo", + data_format="txt" + ), + "entity_types": [], + "token_type": { + "default": req.token_format.default, + "entity_unq_counter": req.token_format.entity_unique_counter, + "entity_only": req.token_format.entity_only, + "vault_token": req.token_format.vault_token + }, + "allow_regex": [], + "restrict_regex": [], + "transformations": None + } + ) mock_poll.assert_called_once() mock_parse.assert_called_once() diff --git a/tests/vault/controller/test__vault.py b/tests/vault/controller/test__vault.py index 5acdf779..d7545fd6 100644 --- a/tests/vault/controller/test__vault.py +++ b/tests/vault/controller/test__vault.py @@ -9,6 +9,7 @@ from skyflow.vault.tokens import DetokenizeRequest, DetokenizeResponse, TokenizeResponse, TokenizeRequest from skyflow.error import SkyflowError from skyflow.utils.validations import validate_file_upload_request +from tests.vault.controller._request_assertions import assert_request, assert_sky_metadata, request_options_headers VAULT_ID = "test_vault_id" TABLE_NAME = "test_table" @@ -78,6 +79,18 @@ def test_insert_with_continue_on_error(self, mock_parse_response, mock_validate) mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response, True) + # Assert the request built for the API matches the expected batch body + assert_request( + self, + records_api.with_raw_response.record_service_batch_operation, + { + "records": expected_body, + "continue_on_error": True, + "byot": request.token_mode.value + }, + expected_args=(VAULT_ID,) + ) + # Assert that the result matches the expected InsertResponse self.assertEqual(result.inserted_fields, expected_inserted_fields) self.assertEqual(result.errors, expected_errors) @@ -123,6 +136,20 @@ def test_insert_with_continue_on_error_false(self, mock_parse_response, mock_val mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response, False) + # Assert the request built for the API matches the expected bulk body + assert_request( + self, + records_api.with_raw_response.record_service_insert_record, + { + "records": expected_body, + "tokenization": request.return_tokens, + "upsert": request.upsert, + "homogeneous": request.homogeneous, + "byot": request.token_mode.value + }, + expected_args=(VAULT_ID, TABLE_NAME) + ) + # Assert that the result matches the expected InsertResponse self.assertEqual(result.inserted_fields, expected_inserted_fields) self.assertEqual(result.errors, None) # No errors expected @@ -181,6 +208,20 @@ def test_insert_with_continue_on_error_false_when_tokens_are_not_none(self, mock mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response, False) + # Assert the request built for the API matches the expected bulk body + assert_request( + self, + records_api.with_raw_response.record_service_insert_record, + { + "records": expected_body, + "tokenization": request.return_tokens, + "upsert": request.upsert, + "homogeneous": request.homogeneous, + "byot": request.token_mode.value + }, + expected_args=(VAULT_ID, TABLE_NAME) + ) + # Assert that the result matches the expected InsertResponse self.assertEqual(result.inserted_fields, expected_inserted_fields) self.assertEqual(result.errors, None) # No errors expected @@ -223,6 +264,19 @@ def test_update_successful(self, mock_parse_response, mock_validate): mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response) + # Assert the request built for the API strips skyflow_id out of the record fields + assert_request( + self, + records_api.record_service_update_record, + { + "id": "12345", + "record": expected_record, + "tokenization": request.return_tokens, + "byot": request.token_mode.value + }, + expected_args=(VAULT_ID, TABLE_NAME) + ) + # Check that the result matches the expected UpdateResponse self.assertEqual(result.updated_field, expected_updated_field) self.assertEqual(result.errors, None) # No errors expected @@ -273,6 +327,14 @@ def test_delete_successful(self, mock_parse_response, mock_validate): mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response) + # Assert the request built for the API carries the ids to delete + assert_request( + self, + records_api.record_service_bulk_delete_record, + {"skyflow_ids": expected_payload}, + expected_args=(VAULT_ID, TABLE_NAME) + ) + # Check that the result matches the expected DeleteResponse self.assertEqual(result.deleted_ids, expected_deleted_ids) self.assertEqual(result.errors, None) # No errors expected @@ -346,6 +408,14 @@ def test_get_successful(self, mock_parse_response, mock_validate): mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response) + # Assert the request built for the API carries every get parameter + assert_request( + self, + records_api.record_service_bulk_get_record, + expected_payload, + expected_args=(VAULT_ID,) + ) + # Check that the result matches the expected GetResponse self.assertEqual(result.data, expected_data) self.assertEqual(result.errors, None) # No errors expected @@ -363,10 +433,16 @@ def test_get_successful_with_column_values(self, mock_parse_response, mock_valid column_name='email' ) - # Expected payload + # Expected payload - ids/fields/paging are absent on a column lookup expected_payload = { "object_name": request.table, + "skyflow_ids": None, + "redaction": request.redaction_type.value, "tokenization": request.return_tokens, + "fields": None, + "offset": None, + "limit": None, + "download_url": None, "column_name": request.column_name, "column_values": request.column_values } @@ -395,7 +471,13 @@ def test_get_successful_with_column_values(self, mock_parse_response, mock_valid # Assertions mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) - records_api.record_service_bulk_get_record.assert_called_once() + # Assert the request built for the API looks the record up by column, not by id + assert_request( + self, + records_api.record_service_bulk_get_record, + expected_payload, + expected_args=(VAULT_ID,) + ) # Check that the result matches the expected GetResponse self.assertEqual(result.data, expected_data) @@ -446,10 +528,29 @@ def test_query_successful(self, mock_parse_response, mock_validate): # Assertions mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + # Assert the query string reaches the API unaltered + assert_request( + self, + query_api.query_service_execute_query, + {"query": "SELECT * FROM test_table"}, + expected_args=(VAULT_ID,) + ) + # Check that the result matches the expected QueryResponse self.assertEqual(result.fields, expected_fields) self.assertEqual(result.errors, None) # No errors expected + @patch("skyflow.vault.controller._vault.validate_query_request") + @patch("skyflow.vault.controller._vault.parse_query_response") + def test_request_carries_sky_metadata_header(self, mock_parse_response, mock_validate): + """Every controller call attaches the sky-metadata header naming this SDK version.""" + query_api = self.vault_client.get_query_api.return_value + query_api.query_service_execute_query.return_value = Mock() + + self.vault.query(QueryRequest(query="SELECT * FROM test_table")) + + assert_sky_metadata(self, request_options_headers(query_api.query_service_execute_query)) + @patch("skyflow.vault.controller._vault.validate_query_request") def test_query_handles_generic_error(self, mock_validate): request = QueryRequest(query="SELECT * from table_name") @@ -511,6 +612,17 @@ def test_detokenize_successful(self, mock_parse_response, mock_validate): mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response) + # Assert every token and its redaction reach the API + assert_request( + self, + tokens_api.with_raw_response.record_service_detokenize, + { + "detokenization_parameters": expected_tokens_list, + "continue_on_error": False + }, + expected_args=(VAULT_ID,) + ) + # Check that the result matches the expected DetokenizeResponse self.assertEqual(result.detokenized_fields, expected_fields) self.assertEqual(result.errors, None) # No errors expected @@ -582,6 +694,14 @@ def test_tokenize_successful(self, mock_parse_response, mock_validate): mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mock_parse_response.assert_called_once_with(mock_api_response) + # Assert every value and its column group reach the API + assert_request( + self, + tokens_api.record_service_tokenize, + {"tokenization_parameters": expected_records_list}, + expected_args=(VAULT_ID,) + ) + # Check that the result matches the expected TokenizeResponse self.assertEqual(result.tokenized_fields, expected_fields) @@ -626,6 +746,19 @@ def test_upload_file_with_file_path_successful(self, mock_validate): result = self.vault.upload_file(request) mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) mocked_open.assert_called_once_with("/path/to/test.txt", "rb") + # The file name is derived from the path and the bytes are read off disk + assert_request( + self, + records_api.with_raw_response.upload_file_v_2, + { + "table_name": "test_table", + "column_name": "file_column", + "file": ("test.txt", b"test file content"), + "skyflow_id": "123", + "return_file_metadata": False + }, + expected_args=(VAULT_ID,) + ) self.assertEqual(result.skyflow_id, "123") self.assertIsNone(result.errors) @@ -651,6 +784,19 @@ def test_upload_file_with_base64_successful(self, mock_validate): # Call upload_file result = self.vault.upload_file(request) mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + # base64 is decoded to bytes and the caller-supplied file name is used + assert_request( + self, + records_api.with_raw_response.upload_file_v_2, + { + "table_name": "test_table", + "column_name": "file_column", + "file": ("test.txt", b"test file content"), + "skyflow_id": "123", + "return_file_metadata": False + }, + expected_args=(VAULT_ID,) + ) self.assertEqual(result.skyflow_id, "123") self.assertIsNone(result.errors) @@ -679,6 +825,19 @@ def test_upload_file_with_file_object_successful(self, mock_validate): # Call upload_file result = self.vault.upload_file(request) mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + # The file object is forwarded as-is, named from its own .name attribute + assert_request( + self, + records_api.with_raw_response.upload_file_v_2, + { + "table_name": "test_table", + "column_name": "file_column", + "file": ("test.txt", mock_file), + "skyflow_id": "123", + "return_file_metadata": False + }, + expected_args=(VAULT_ID,) + ) self.assertEqual(result.skyflow_id, "123") self.assertIsNone(result.errors) @@ -739,6 +898,19 @@ def test_upload_file_without_skyflow_id_successful(self, mock_validate): result = self.vault.upload_file(request) mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) self.assertIsNone(request.skyflow_id) + # skyflow_id is still sent, explicitly as None, so the vault generates one + assert_request( + self, + records_api.with_raw_response.upload_file_v_2, + { + "table_name": "test_table", + "column_name": "file_column", + "file": ("test.txt", b"test file content"), + "skyflow_id": None, + "return_file_metadata": False + }, + expected_args=(VAULT_ID,) + ) self.assertEqual(result.skyflow_id, "generated-id-123") self.assertIsNone(result.errors) From 4ec4cc8e53e4ecc1534ae11e9eb6f4c66cff0de1 Mon Sep 17 00:00:00 2001 From: yaswanth-pula-skyflow Date: Tue, 4 Aug 2026 13:09:37 +0530 Subject: [PATCH 7/7] SK-3039:Remove unnesscary patch on bearerToken tests. --- tests/client/test_skyflow.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/tests/client/test_skyflow.py b/tests/client/test_skyflow.py index d2c07392..f1d83ae1 100644 --- a/tests/client/test_skyflow.py +++ b/tests/client/test_skyflow.py @@ -442,9 +442,8 @@ def _build_and_generate(self, credentials, mock_gen): vault_client.initialize_client_configuration() return mock_gen.call_args[0][1] - @patch("skyflow.vault.client.client.Skyflow") @patch("skyflow.vault.client.client.generate_bearer_token_from_creds", return_value=("token", "bearer")) - def test_string_context_and_roles_reach_token_engine(self, mock_gen, _mock_api): + def test_string_context_and_roles_reach_token_engine(self, mock_gen): options = self._build_and_generate( { "credentials_string": self.CREDENTIALS_STRING, @@ -456,9 +455,8 @@ def test_string_context_and_roles_reach_token_engine(self, mock_gen, _mock_api): self.assertEqual(options["role_ids"], ["role_id_1", "role_id_2"]) self.assertEqual(options["ctx"], "user_12345") - @patch("skyflow.vault.client.client.Skyflow") @patch("skyflow.vault.client.client.generate_bearer_token_from_creds", return_value=("token", "bearer")) - def test_dict_context_reaches_token_engine(self, mock_gen, _mock_api): + def test_dict_context_reaches_token_engine(self, mock_gen): options = self._build_and_generate( { "credentials_string": self.CREDENTIALS_STRING,