Skip to content

Commit 2f1890f

Browse files
Merge pull request #273 from skyflowapi/release/26.8.1
SK-3039:Fix Skyflow clients bearertokens roles & context.
2 parents 20bc6cc + 55741fb commit 2f1890f

13 files changed

Lines changed: 888 additions & 55 deletions

File tree

setup.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
if sys.version_info < (3, 9):
99
raise RuntimeError("skyflow requires Python 3.9+")
10-
current_version = '2.1.2'
10+
current_version = '2.1.2.dev0+5e4b2fb'
1111

1212
with open('README.md', 'r', encoding='utf-8') as f:
1313
long_description = f.read()

skyflow/utils/_skyflow_messages.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -58,6 +58,8 @@ class Error(Enum):
5858
INVALID_ROLES_KEY_TYPE = f"{error_prefix} Validation error. Invalid roles. Specify roles as an array."
5959
EMPTY_ROLES_IN_CONFIG = f"{error_prefix} Validation error. Invalid roles for {{}} with id {{}}. Specify at least one role."
6060
EMPTY_ROLES = f"{error_prefix} Validation error. Invalid roles. Specify at least one role."
61+
INVALID_ROLE_ELEMENT_TYPE_IN_CONFIG = f"{error_prefix} Validation error. Invalid roles for {{}} with id {{}}. Each role must be a non-empty string."
62+
INVALID_ROLE_ELEMENT_TYPE = f"{error_prefix} Validation error. Invalid roles. Each role must be a non-empty string."
6163
EMPTY_CONTEXT_IN_CONFIG = f"{error_prefix} Initialization failed. Invalid context provided for {{}} with id {{}}. Specify context as type Context."
6264
EMPTY_CONTEXT = f"{error_prefix} Initialization failed. Invalid context provided. Specify context as type Context."
6365
INVALID_CONTEXT_IN_CONFIG = f"{error_prefix} Initialization failed. Invalid context for {{}} with id {{}}. Specify a valid context."
@@ -109,6 +111,7 @@ class Error(Enum):
109111
EMPTY_RECORD_IDS_IN_DELETE = f"{error_prefix} Validation error. 'record ids' array can't be empty. Specify one or more record ids."
110112
BULK_DELETE_FAILURE = f"{error_prefix} Delete operation failed."
111113
EMPTY_SKYFLOW_ID= f"{error_prefix} Validation error. skyflow_id can't be empty."
114+
INVALID_SKYFLOW_ID_TYPE = f"{error_prefix} Validation error. 'skyflow_id' has a value of type {{}}. Specify 'skyflow_id' as a string."
112115
INVALID_FILE_COLUMN_NAME= f"{error_prefix} Validation error. 'column_name' can't be empty."
113116

114117
INVALID_QUERY_TYPE = f"{error_prefix} Validation error. Query parameter is of type {{}}. Specify as a string."

skyflow/utils/_version.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
SDK_VERSION = '2.1.2'
1+
SDK_VERSION = '2.1.2.dev0+5e4b2fb'

skyflow/utils/validations/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
validate_update_vault_config,
66
validate_update_connection_config,
77
validate_credentials,
8+
validate_token_options,
89
validate_log_level,
910
validate_delete_request,
1011
validate_query_request,

skyflow/utils/validations/_validations.py

Lines changed: 72 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
import json
33
import os
44
from skyflow.service_account import is_expired
5+
from skyflow.service_account._utils import _validate_and_resolve_ctx
56
from skyflow.utils.enums import LogLevel, Env, RedactionType, TokenMode, DetectEntities, DetectOutputTranscriptions, \
67
MaskingMethod
78
from skyflow.error import SkyflowError
@@ -81,6 +82,54 @@ def validate_api_key(api_key: str, logger = None) -> bool:
8182

8283
return True
8384

85+
def validate_token_options(logger, credentials, config_id_type=None, config_id=None):
86+
"""Validate the roles/context token options nested under credentials.
87+
88+
Shared by config validation and VaultClient, so a directly constructed client
89+
cannot silently generate an unscoped or context-less token.
90+
"""
91+
if CredentialField.ROLES in credentials:
92+
empty_roles_error = (
93+
SkyflowMessages.Error.EMPTY_ROLES_IN_CONFIG.value.format(config_id_type, config_id)
94+
if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value
95+
)
96+
validate_required_field(
97+
logger, credentials, CredentialField.ROLES, list,
98+
empty_roles_error,
99+
SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE_IN_CONFIG.value.format(config_id_type, config_id)
100+
if config_id_type and config_id else SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE.value
101+
)
102+
if not credentials.get(CredentialField.ROLES):
103+
raise SkyflowError(empty_roles_error, invalid_input_error_code)
104+
105+
invalid_role_element_error = (
106+
SkyflowMessages.Error.INVALID_ROLE_ELEMENT_TYPE_IN_CONFIG.value.format(config_id_type, config_id)
107+
if config_id_type and config_id else SkyflowMessages.Error.INVALID_ROLE_ELEMENT_TYPE.value
108+
)
109+
for role in credentials.get(CredentialField.ROLES):
110+
if not isinstance(role, str) or not role.strip():
111+
raise SkyflowError(invalid_role_element_error, invalid_input_error_code)
112+
113+
if CredentialField.CONTEXT in credentials:
114+
empty_context_error = (
115+
SkyflowMessages.Error.EMPTY_CONTEXT_IN_CONFIG.value.format(config_id_type, config_id)
116+
if config_id_type and config_id else SkyflowMessages.Error.EMPTY_CONTEXT.value
117+
)
118+
# Scalars are accepted because the token engine and the auth service both support
119+
# them as ctx claims - bool is listed explicitly even though it is an int subclass.
120+
validate_required_field(
121+
logger, credentials, CredentialField.CONTEXT, (str, dict, bool, int, float),
122+
empty_context_error,
123+
SkyflowMessages.Error.INVALID_CONTEXT_IN_CONFIG.value.format(config_id_type, config_id)
124+
if config_id_type and config_id else SkyflowMessages.Error.INVALID_CONTEXT.value
125+
)
126+
context = credentials.get(CredentialField.CONTEXT)
127+
if isinstance(context, dict):
128+
if not context:
129+
raise SkyflowError(empty_context_error, invalid_input_error_code)
130+
# Surface invalid ctx keys at config time instead of on the first API call.
131+
_validate_and_resolve_ctx(context)
132+
84133
def validate_credentials(logger, credentials, config_id_type=None, config_id=None):
85134
key_present = [k for k in [CredentialField.PATH, CredentialField.TOKEN, CredentialField.CREDENTIALS_STRING, CredentialField.API_KEY] if credentials.get(k)]
86135

@@ -101,23 +150,7 @@ def validate_credentials(logger, credentials, config_id_type=None, config_id=Non
101150
log_error_log(error_message, logger)
102151
raise SkyflowError(error_message, invalid_input_error_code)
103152

104-
if CredentialField.ROLES in credentials:
105-
validate_required_field(
106-
logger, credentials, CredentialField.ROLES, list,
107-
SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE_IN_CONFIG.value.format(config_id_type, config_id)
108-
if config_id_type and config_id else SkyflowMessages.Error.INVALID_ROLES_KEY_TYPE.value,
109-
SkyflowMessages.Error.EMPTY_ROLES_IN_CONFIG.value.format(config_id_type, config_id)
110-
if config_id_type and config_id else SkyflowMessages.Error.EMPTY_ROLES.value
111-
)
112-
113-
if CredentialField.CONTEXT in credentials:
114-
validate_required_field(
115-
logger, credentials, CredentialField.CONTEXT, str,
116-
SkyflowMessages.Error.EMPTY_CONTEXT_IN_CONFIG.value.format(config_id_type, config_id)
117-
if config_id_type and config_id else SkyflowMessages.Error.EMPTY_CONTEXT.value,
118-
SkyflowMessages.Error.INVALID_CONTEXT_IN_CONFIG.value.format(config_id_type, config_id)
119-
if config_id_type and config_id else SkyflowMessages.Error.INVALID_CONTEXT.value
120-
)
153+
validate_token_options(logger, credentials, config_id_type, config_id)
121154

122155
if CredentialField.CREDENTIALS_STRING in credentials:
123156
validate_required_field(
@@ -287,7 +320,7 @@ def validate_update_connection_config(logger, config):
287320

288321
if ConfigField.CREDENTIALS not in config:
289322
raise SkyflowError(SkyflowMessages.Error.EMPTY_CREDENTIALS.value.format(ConfigType.CONNECTION, connection_id), invalid_input_error_code)
290-
validate_credentials(logger, config.get(ConfigField.CREDENTIALS))
323+
validate_credentials(logger, config.get(ConfigField.CREDENTIALS), ConfigType.CONNECTION, connection_id)
291324

292325
return True
293326

@@ -496,6 +529,11 @@ def validate_insert_request(logger, request):
496529
raise SkyflowError(SkyflowMessages.Error.INVALID_CONTINUE_ON_ERROR_TYPE.value, invalid_input_error_code)
497530

498531
if request.tokens:
532+
if not isinstance(request.tokens, list) or not request.tokens or not all(
533+
isinstance(t, dict) for t in request.tokens):
534+
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_TOKENS.value.format(RequestOperation.INSERT), logger=logger)
535+
raise SkyflowError(SkyflowMessages.Error.INVALID_TYPE_OF_DATA_IN_INSERT.value, invalid_input_error_code)
536+
499537
for i, item in enumerate(request.tokens, start=1):
500538
for key, value in item.items():
501539
if key is None or key == "":
@@ -505,10 +543,6 @@ def validate_insert_request(logger, request):
505543
if value is None or value == "":
506544
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_OR_NULL_KEY_IN_TOKENS.value.format(RequestOperation.INSERT, key),
507545
logger=logger)
508-
if not isinstance(request.tokens, list) or not request.tokens or not all(
509-
isinstance(t, dict) for t in request.tokens):
510-
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_TOKENS.value.format(RequestOperation.INSERT), logger=logger)
511-
raise SkyflowError(SkyflowMessages.Error.INVALID_TYPE_OF_DATA_IN_INSERT.value, invalid_input_error_code)
512546

513547
if request.token_mode == TokenMode.ENABLE and not request.tokens:
514548
raise SkyflowError(SkyflowMessages.Error.NO_TOKENS_IN_INSERT.value.format(request.token_mode), invalid_input_error_code)
@@ -538,6 +572,10 @@ def validate_delete_request(logger, request):
538572
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_IDS.value.format(RequestOperation.DELETE), logger=logger)
539573
raise SkyflowError(SkyflowMessages.Error.EMPTY_RECORD_IDS_IN_DELETE.value, invalid_input_error_code)
540574

575+
if not isinstance(request.ids, list):
576+
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_IDS.value.format(RequestOperation.DELETE), logger=logger)
577+
raise SkyflowError(SkyflowMessages.Error.INVALID_IDS_TYPE.value.format(type(request.ids)), invalid_input_error_code)
578+
541579
def validate_query_request(logger, request):
542580
if not isinstance(request.query, str):
543581
query_type = str(type(request.query))
@@ -650,6 +688,8 @@ def validate_update_request(logger, request):
650688
skyflow_id = request.data.get(ResponseField.SKYFLOW_ID)
651689
if skyflow_id is None:
652690
log_error_log(SkyflowMessages.ErrorLogs.SKYFLOW_ID_IS_REQUIRED.value.format(RequestOperation.UPDATE), logger=logger)
691+
elif not isinstance(skyflow_id, str):
692+
raise SkyflowError(SkyflowMessages.Error.INVALID_SKYFLOW_ID_TYPE.value.format(type(skyflow_id)), invalid_input_error_code)
653693
elif not skyflow_id.strip():
654694
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_SKYFLOW_ID.value.format(RequestOperation.UPDATE), logger=logger)
655695

@@ -708,7 +748,7 @@ def validate_detokenize_request(logger, request):
708748
raise SkyflowError(SkyflowMessages.Error.EMPTY_TOKENS_LIST_VALUE.value, invalid_input_error_code)
709749

710750
for item in request.data:
711-
if ResponseField.TOKEN not in item:
751+
if not isinstance(item, dict) or ResponseField.TOKEN not in item:
712752
raise SkyflowError(SkyflowMessages.Error.INVALID_TOKENS_LIST_VALUE.value.format(type(request.data)),
713753
invalid_input_error_code)
714754

@@ -766,18 +806,26 @@ def validate_file_upload_request(logger, request):
766806
table = getattr(request, FileUploadField.TABLE, None)
767807
if table is None:
768808
raise SkyflowError(SkyflowMessages.Error.INVALID_TABLE_VALUE.value, invalid_input_error_code)
809+
elif not isinstance(table, str):
810+
log_error_log(SkyflowMessages.ErrorLogs.TABLE_IS_REQUIRED.value.format(RequestOperation.FILE_UPLOAD), logger=logger)
811+
raise SkyflowError(SkyflowMessages.Error.INVALID_TABLE_VALUE.value, invalid_input_error_code)
769812
elif table.strip() == "":
770813
raise SkyflowError(SkyflowMessages.Error.EMPTY_TABLE_VALUE.value, invalid_input_error_code)
771814

772815
# Skyflow ID
773816
skyflow_id = getattr(request, FileUploadField.SKYFLOW_ID, None)
817+
if skyflow_id is not None and not isinstance(skyflow_id, str):
818+
raise SkyflowError(SkyflowMessages.Error.INVALID_SKYFLOW_ID_TYPE.value.format(type(skyflow_id)), invalid_input_error_code)
774819
if skyflow_id is not None and skyflow_id.strip() == "":
775820
raise SkyflowError(SkyflowMessages.Error.EMPTY_SKYFLOW_ID.value.format(RequestOperation.FILE_UPLOAD), invalid_input_error_code)
776821

777822
# Column Name
778823
column_name = getattr(request, FileUploadField.COLUMN_NAME, None)
779824
if column_name is None:
780825
raise SkyflowError(SkyflowMessages.Error.INVALID_FILE_COLUMN_NAME.value.format(type(column_name)), invalid_input_error_code)
826+
elif not isinstance(column_name, str):
827+
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_FILE_COLUMN_NAME.value, logger)
828+
raise SkyflowError(SkyflowMessages.Error.INVALID_COLUMN_NAME.value.format(type(column_name)), invalid_input_error_code)
781829
elif column_name.strip() == "":
782830
log_error_log(SkyflowMessages.ErrorLogs.EMPTY_FILE_COLUMN_NAME.value, logger)
783831
raise SkyflowError(SkyflowMessages.Error.INVALID_FILE_COLUMN_NAME.value.format(type(column_name)), invalid_input_error_code)
@@ -796,7 +844,7 @@ def validate_file_upload_request(logger, request):
796844

797845
# Check base64 if present
798846
if not is_none_or_empty(base64_str):
799-
if is_none_or_empty(file_name):
847+
if not isinstance(file_name, str) or is_none_or_empty(file_name):
800848
raise SkyflowError(SkyflowMessages.Error.INVALID_FILE_NAME.value, invalid_input_error_code)
801849
try:
802850
base64.b64decode(base64_str)

skyflow/vault/client/client.py

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from skyflow.utils import get_vault_url, get_credentials, SkyflowMessages
55
from skyflow.utils.logger import log_info
66
from skyflow.utils.constants import OptionField, CredentialField, ConfigField
7+
from skyflow.utils.validations import validate_token_options
78

89

910
class VaultClient:
@@ -74,10 +75,15 @@ def get_bearer_token(self, credentials):
7475
elif CredentialField.TOKEN in credentials:
7576
return credentials.get(CredentialField.TOKEN)
7677

77-
options = {
78-
OptionField.ROLE_IDS: self.__config.get(OptionField.ROLES),
79-
OptionField.CTX: self.__config.get(OptionField.CTX)
80-
}
78+
validate_token_options(self.__logger, credentials)
79+
80+
options = {}
81+
if CredentialField.ROLES in credentials:
82+
options[OptionField.ROLE_IDS] = credentials.get(CredentialField.ROLES)
83+
84+
if CredentialField.CONTEXT in credentials:
85+
options[OptionField.CTX] = credentials.get(CredentialField.CONTEXT)
86+
8187
if CredentialField.TOKEN_URI_OPTION in credentials and credentials.get(CredentialField.TOKEN_URI_OPTION):
8288
options[CredentialField.TOKEN_URI_OPTION] = credentials.get(CredentialField.TOKEN_URI_OPTION)
8389

tests/client/test_skyflow.py

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -420,6 +420,53 @@ def test_get_bearer_token_passes_token_uri_option(self, _mock_expired, mock_gen)
420420
self.assertEqual(options_passed["token_uri"], "https://custom-token-uri.com/token")
421421

422422

423+
class TestBearerTokenContextAndRolesEndToEnd(unittest.TestCase):
424+
"""roles/context configured via add_vault_config() must reach the token engine."""
425+
426+
CREDENTIALS_STRING = '{"clientID":"id","privateKey":"pk","keyID":"kid","tokenURI":"https://token.uri"}'
427+
428+
def _build_and_generate(self, credentials, mock_gen):
429+
config = {
430+
"vault_id": "VAULT_ID",
431+
"cluster_id": "CLUSTER_ID",
432+
"env": Env.DEV,
433+
"credentials": credentials,
434+
}
435+
client = (
436+
Skyflow.builder()
437+
.add_vault_config(config)
438+
.set_log_level(LogLevel.OFF)
439+
.build()
440+
)
441+
vault_client = client._Skyflow__builder.get_vault_config("VAULT_ID").get("vault_client")
442+
vault_client.initialize_client_configuration()
443+
return mock_gen.call_args[0][1]
444+
445+
@patch("skyflow.vault.client.client.generate_bearer_token_from_creds", return_value=("token", "bearer"))
446+
def test_string_context_and_roles_reach_token_engine(self, mock_gen):
447+
options = self._build_and_generate(
448+
{
449+
"credentials_string": self.CREDENTIALS_STRING,
450+
"roles": ["role_id_1", "role_id_2"],
451+
"context": "user_12345",
452+
},
453+
mock_gen,
454+
)
455+
self.assertEqual(options["role_ids"], ["role_id_1", "role_id_2"])
456+
self.assertEqual(options["ctx"], "user_12345")
457+
458+
@patch("skyflow.vault.client.client.generate_bearer_token_from_creds", return_value=("token", "bearer"))
459+
def test_dict_context_reaches_token_engine(self, mock_gen):
460+
options = self._build_and_generate(
461+
{
462+
"credentials_string": self.CREDENTIALS_STRING,
463+
"context": {"role": "admin", "department": "finance"},
464+
},
465+
mock_gen,
466+
)
467+
self.assertEqual(options["ctx"], {"role": "admin", "department": "finance"})
468+
469+
423470
class TestUpdateLogLevelDeprecation(unittest.TestCase):
424471
def _build_client(self):
425472
return Skyflow.builder().add_vault_config(VALID_VAULT_CONFIG).build()

0 commit comments

Comments
 (0)