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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@

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

with open('README.md', 'r', encoding='utf-8') as f:
long_description = f.read()
Expand Down
3 changes: 3 additions & 0 deletions skyflow/utils/_skyflow_messages.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."
Expand Down Expand Up @@ -109,6 +111,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."
Expand Down
2 changes: 1 addition & 1 deletion skyflow/utils/_version.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
SDK_VERSION = '2.1.2'
SDK_VERSION = '2.1.2.dev0+5e4b2fb'
1 change: 1 addition & 0 deletions skyflow/utils/validations/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
96 changes: 72 additions & 24 deletions skyflow/utils/validations/_validations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -81,6 +82,54 @@ def validate_api_key(api_key: str, logger = None) -> bool:

return True

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,
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(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)
if config_id_type and config_id else SkyflowMessages.Error.EMPTY_CONTEXT.value
)
# Scalars are accepted because the token engine and the auth service both support
# them as ctx claims - bool is listed explicitly even though it is an int subclass.
validate_required_field(
logger, credentials, CredentialField.CONTEXT, (str, dict, bool, int, float),
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)

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)]

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

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 CredentialField.CONTEXT in credentials:
validate_required_field(
logger, credentials, CredentialField.CONTEXT, str,
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,
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
)
validate_token_options(logger, credentials, config_id_type, config_id)

if CredentialField.CREDENTIALS_STRING in credentials:
validate_required_field(
Expand Down Expand Up @@ -287,7 +320,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

Expand Down Expand Up @@ -496,6 +529,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 == "":
Expand All @@ -505,10 +543,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)
Expand Down Expand Up @@ -538,6 +572,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))
Expand Down Expand Up @@ -650,6 +688,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)

Expand Down Expand Up @@ -708,7 +748,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)

Expand Down Expand Up @@ -766,18 +806,26 @@ 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)

# Column Name
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)
Expand All @@ -796,7 +844,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)
Expand Down
14 changes: 10 additions & 4 deletions skyflow/vault/client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -74,10 +75,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)
}
validate_token_options(self.__logger, credentials)

options = {}
if CredentialField.ROLES in credentials:
options[OptionField.ROLE_IDS] = credentials.get(CredentialField.ROLES)

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)

Expand Down
47 changes: 47 additions & 0 deletions tests/client/test_skyflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -420,6 +420,53 @@ 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.generate_bearer_token_from_creds", return_value=("token", "bearer"))
def test_string_context_and_roles_reach_token_engine(self, mock_gen):
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.generate_bearer_token_from_creds", return_value=("token", "bearer"))
def test_dict_context_reaches_token_engine(self, mock_gen):
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()
Expand Down
Loading
Loading