diff --git a/common/tests/utils/validations/test__validations.py b/common/tests/utils/validations/test__validations.py index b10fc00..7684841 100644 --- a/common/tests/utils/validations/test__validations.py +++ b/common/tests/utils/validations/test__validations.py @@ -7,6 +7,7 @@ validate_update_vault_config, validate_credentials, validate_log_level, + validate_non_empty_string_list, ) VALID_VAULT_CONFIG = { @@ -165,5 +166,31 @@ def test_uses_injected_messages(self): self.assertIn("FAKE", ctx.exception.message) +class TestValidateNonEmptyStringList(unittest.TestCase): + def test_valid_list_passes(self): + validate_non_empty_string_list(None, ["a", "b"], "boom") # should not raise + + def test_non_list_raises_with_given_error(self): + with self.assertRaises(SkyflowError) as ctx: + validate_non_empty_string_list(None, "not-a-list", "boom") + self.assertEqual(ctx.exception.message, "boom") + + def test_empty_list_raises(self): + with self.assertRaises(SkyflowError): + validate_non_empty_string_list(None, [], "boom") + + def test_none_raises(self): + with self.assertRaises(SkyflowError): + validate_non_empty_string_list(None, None, "boom") + + def test_non_string_entry_raises(self): + with self.assertRaises(SkyflowError): + validate_non_empty_string_list(None, ["a", 1], "boom") + + def test_blank_string_entry_raises(self): + with self.assertRaises(SkyflowError): + validate_non_empty_string_list(None, ["a", " "], "boom") + + if __name__ == "__main__": unittest.main() diff --git a/common/utils/validations/__init__.py b/common/utils/validations/__init__.py index d49cc5d..f510089 100644 --- a/common/utils/validations/__init__.py +++ b/common/utils/validations/__init__.py @@ -4,6 +4,7 @@ validate_credentials, validate_log_level, validate_keys, + validate_non_empty_string_list, validate_vault_config, validate_update_vault_config, ) diff --git a/common/utils/validations/_validations.py b/common/utils/validations/_validations.py index f600b6a..2ce3bfd 100644 --- a/common/utils/validations/_validations.py +++ b/common/utils/validations/_validations.py @@ -173,6 +173,11 @@ def validate_keys(logger, config, config_keys, messages=None): raise SkyflowError(messages.Error.INVALID_KEY.value.format(key), invalid_input_error_code) +def validate_non_empty_string_list(logger, value, error): + if not isinstance(value, list) or not value or not all(isinstance(item, str) and item.strip() for item in value): + raise SkyflowError(error, invalid_input_error_code) + + def validate_vault_config(logger, config, messages=None): messages = messages or SkyflowMessages log_info(messages.Info.VALIDATING_VAULT_CONFIG.value, logger) diff --git a/flowvault/skyflow_flowvault/utils/_skyflow_messages.py b/flowvault/skyflow_flowvault/utils/_skyflow_messages.py index 3559629..ad6bb53 100644 --- a/flowvault/skyflow_flowvault/utils/_skyflow_messages.py +++ b/flowvault/skyflow_flowvault/utils/_skyflow_messages.py @@ -18,8 +18,8 @@ class SkyflowMessages: class Error(Enum): EMPTY_RECORDS_IN_INSERT = f"{error_prefix} Insert failed. Specify at least one record to insert." INVALID_RECORDS_TYPE_IN_INSERT = f"{error_prefix} Insert failed. 'records' must be a list of dicts." - INVALID_RECORD_DATA_IN_INSERT = f"{error_prefix} Insert failed. Each record's 'values' must be a non-empty dict." - INVALID_TABLE_NAME_IN_INSERT = f"{error_prefix} Insert failed. 'table' must be a non-empty string." + INVALID_RECORD_DATA_IN_INSERT = f"{error_prefix} Validation error. Each record's 'values' must be a non-empty dict." + INVALID_TABLE_NAME_IN_INSERT = f"{error_prefix} Validation error. 'table' must be a non-empty string." INVALID_UPSERT_TYPE_IN_INSERT = f"{error_prefix} Insert failed. 'upsert' must be a dict." INVALID_UPSERT_UNIQUE_COLUMNS_IN_INSERT = f"{error_prefix} Insert failed. Upsert's 'unique_columns' must be a non-empty list of strings." INVALID_UPSERT_UPDATE_TYPE_IN_INSERT = f"{error_prefix} Insert failed. Upsert's 'update_type' must be an UpsertType value." @@ -44,14 +44,75 @@ class Error(Enum): "provided per-record -- InsertRequest's request-level 'upsert' cannot be used while " "'table' is set on individual records." ) - EMPTY_KEY_IN_INSERT_DATA = f"{error_prefix} Insert failed. Each record's 'values' must not contain a null or empty key." + EMPTY_KEY_IN_INSERT_DATA = f"{error_prefix} Validation error. Each record's 'values' must not contain a null or empty key." EMPTY_VALUE_IN_INSERT_DATA = f"{error_prefix} Insert failed. Each record's 'values' must not contain a null or empty value." + MISSING_TABLE_NAME_IN_GET = f"{error_prefix} Get failed. Specify a table name." + MISSING_IDS_OR_UNIQUE_VALUES_IN_GET = f"{error_prefix} Get failed. Specify at least one of 'ids' or 'unique_values'." + INVALID_IDS_IN_GET = f"{error_prefix} Get failed. 'ids' must be a non-empty list of strings." + + EMPTY_RECORDS_IN_UPDATE = f"{error_prefix} Update failed. Specify at least one record to update." + INVALID_RECORDS_TYPE_IN_UPDATE = f"{error_prefix} Update failed. 'records' must be a list of dicts." + MISSING_SKYFLOW_ID_IN_UPDATE = f"{error_prefix} Update failed. Each record must specify a non-empty 'skyflow_id'." + INVALID_UPDATE_TYPE_IN_UPDATE = f"{error_prefix} Update failed. 'update_type' must be an UpsertType value." + TABLE_NAME_IN_BOTH_PLACES_IN_UPDATE = ( + f"{error_prefix} Update failed. 'table' cannot be set on UpdateRequest at the same " + "time as any record's 'table' -- specify a table name outside the records " + "(request-level, applying to all of them) or inside each record, but not both at once." + ) + TABLE_NAME_MISSING_IN_UPDATE = ( + f"{error_prefix} Update failed. 'table' is not set on UpdateRequest, so every record " + "must set its own 'table' -- either set 'table' once at the request level, or set it " + "individually on every record." + ) + + MISSING_TABLE_NAME_IN_DELETE = f"{error_prefix} Delete failed. Specify a table name." + MISSING_IDS_OR_UNIQUE_VALUES_IN_DELETE = f"{error_prefix} Delete failed. Specify at least one of 'ids' or 'unique_values'." + INVALID_IDS_IN_DELETE = f"{error_prefix} Delete failed. 'ids' must be a non-empty list of strings." + + EMPTY_TOKENS_IN_DETOKENIZE = f"{error_prefix} Detokenize failed. Specify at least one token to detokenize." + INVALID_TOKENS_TYPE_IN_DETOKENIZE = f"{error_prefix} Detokenize failed. 'tokens' must be a non-empty list of strings." + INVALID_TOKEN_GROUP_REDACTIONS_IN_DETOKENIZE = f"{error_prefix} Detokenize failed. 'token_group_redactions' must be a list of dicts with 'token_group_name' and 'redaction' keys." + + EMPTY_VALUES_IN_TOKENIZE = f"{error_prefix} Tokenize failed. Specify at least one value to tokenize." + INVALID_VALUES_TYPE_IN_TOKENIZE = f"{error_prefix} Tokenize failed. 'values' must be a list of dicts." + MISSING_TOKEN_GROUP_NAMES_IN_TOKENIZE = f"{error_prefix} Tokenize failed. Each value must specify a non-empty 'token_group_names' list of strings." + class Info(Enum): VALIDATE_INSERT_REQUEST = f"{INFO}: [{error_prefix}] Validating insert request." INSERT_TRIGGERED = f"{INFO}: [{error_prefix}] Insert method triggered." INSERT_REQUEST_RESOLVED = f"{INFO}: [{error_prefix}] Insert request resolved." INSERT_SUCCESS = f"{INFO}: [{error_prefix}] Data inserted." + VALIDATE_GET_REQUEST = f"{INFO}: [{error_prefix}] Validating get request." + GET_TRIGGERED = f"{INFO}: [{error_prefix}] Get method triggered." + GET_REQUEST_RESOLVED = f"{INFO}: [{error_prefix}] Get request resolved." + GET_SUCCESS = f"{INFO}: [{error_prefix}] Data fetched." + + VALIDATE_UPDATE_REQUEST = f"{INFO}: [{error_prefix}] Validating update request." + UPDATE_TRIGGERED = f"{INFO}: [{error_prefix}] Update method triggered." + UPDATE_REQUEST_RESOLVED = f"{INFO}: [{error_prefix}] Update request resolved." + UPDATE_SUCCESS = f"{INFO}: [{error_prefix}] Data updated." + + VALIDATE_DELETE_REQUEST = f"{INFO}: [{error_prefix}] Validating delete request." + DELETE_TRIGGERED = f"{INFO}: [{error_prefix}] Delete method triggered." + DELETE_REQUEST_RESOLVED = f"{INFO}: [{error_prefix}] Delete request resolved." + DELETE_SUCCESS = f"{INFO}: [{error_prefix}] Data deleted." + + VALIDATE_DETOKENIZE_REQUEST = f"{INFO}: [{error_prefix}] Validating detokenize request." + DETOKENIZE_TRIGGERED = f"{INFO}: [{error_prefix}] Detokenize method triggered." + DETOKENIZE_REQUEST_RESOLVED = f"{INFO}: [{error_prefix}] Detokenize request resolved." + DETOKENIZE_SUCCESS = f"{INFO}: [{error_prefix}] Tokens detokenized." + + VALIDATE_TOKENIZE_REQUEST = f"{INFO}: [{error_prefix}] Validating tokenize request." + TOKENIZE_TRIGGERED = f"{INFO}: [{error_prefix}] Tokenize method triggered." + TOKENIZE_REQUEST_RESOLVED = f"{INFO}: [{error_prefix}] Tokenize request resolved." + TOKENIZE_SUCCESS = f"{INFO}: [{error_prefix}] Values tokenized." + class ErrorLogs(Enum): INSERT_RECORDS_REJECTED = f"{ERROR}: [{error_prefix}] Insert call resulted in failure." + GET_RECORDS_REJECTED = f"{ERROR}: [{error_prefix}] Get call resulted in failure." + UPDATE_RECORDS_REJECTED = f"{ERROR}: [{error_prefix}] Update call resulted in failure." + DELETE_RECORDS_REJECTED = f"{ERROR}: [{error_prefix}] Delete call resulted in failure." + DETOKENIZE_RECORDS_REJECTED = f"{ERROR}: [{error_prefix}] Detokenize call resulted in failure." + TOKENIZE_RECORDS_REJECTED = f"{ERROR}: [{error_prefix}] Tokenize call resulted in failure." diff --git a/flowvault/skyflow_flowvault/utils/validations/__init__.py b/flowvault/skyflow_flowvault/utils/validations/__init__.py index 499bed6..9502717 100644 --- a/flowvault/skyflow_flowvault/utils/validations/__init__.py +++ b/flowvault/skyflow_flowvault/utils/validations/__init__.py @@ -1 +1,10 @@ -from ._validations import validate_vault_config, validate_update_vault_config, validate_insert_request +from ._validations import ( + validate_vault_config, + validate_update_vault_config, + validate_insert_request, + validate_get_request, + validate_update_request, + validate_delete_request, + validate_detokenize_request, + validate_tokenize_request, +) diff --git a/flowvault/skyflow_flowvault/utils/validations/_validations.py b/flowvault/skyflow_flowvault/utils/validations/_validations.py index a8ea74d..f8340e7 100644 --- a/flowvault/skyflow_flowvault/utils/validations/_validations.py +++ b/flowvault/skyflow_flowvault/utils/validations/_validations.py @@ -3,6 +3,7 @@ from common.utils.validations import ( validate_keys, validate_credentials, + validate_non_empty_string_list, validate_vault_config, validate_update_vault_config, ) @@ -11,6 +12,7 @@ VALID_INSERT_RECORD_KEYS = ["values", "table", "upsert"] VALID_UPSERT_KEYS = ["update_type", "unique_columns"] +VALID_UPDATE_RECORD_KEYS = ["skyflow_id", "values", "tokens", "table"] invalid_input_error_code = CommonMessages.ErrorCodes.INVALID_INPUT.value @@ -77,3 +79,88 @@ def validate_insert_request(logger, request): else: if request.upsert is not None: raise SkyflowError(SkyflowMessages.Error.REQUEST_LEVEL_UPSERT_NOT_ALLOWED_IN_INSERT.value, invalid_input_error_code) + + +def validate_get_request(logger, request): + if not request.table: + raise SkyflowError(SkyflowMessages.Error.MISSING_TABLE_NAME_IN_GET.value, invalid_input_error_code) + + if not request.ids and not request.unique_values: + raise SkyflowError(SkyflowMessages.Error.MISSING_IDS_OR_UNIQUE_VALUES_IN_GET.value, invalid_input_error_code) + + if request.ids is not None: + validate_non_empty_string_list(logger, request.ids, SkyflowMessages.Error.INVALID_IDS_IN_GET.value) + + +def validate_update_request(logger, request): + if not isinstance(request.records, list) or not all(isinstance(r, dict) for r in request.records): + raise SkyflowError(SkyflowMessages.Error.INVALID_RECORDS_TYPE_IN_UPDATE.value, invalid_input_error_code) + + if not request.records: + raise SkyflowError(SkyflowMessages.Error.EMPTY_RECORDS_IN_UPDATE.value, invalid_input_error_code) + + if request.update_type is not None and not isinstance(request.update_type, UpsertType): + raise SkyflowError(SkyflowMessages.Error.INVALID_UPDATE_TYPE_IN_UPDATE.value, invalid_input_error_code) + + for record in request.records: + validate_keys(logger, record, VALID_UPDATE_RECORD_KEYS) + skyflow_id = record.get("skyflow_id") + if not isinstance(skyflow_id, str) or not skyflow_id.strip(): + raise SkyflowError(SkyflowMessages.Error.MISSING_SKYFLOW_ID_IN_UPDATE.value, invalid_input_error_code) + + table_at_request_level = request.table is not None + + if table_at_request_level: + for record in request.records: + if record.get("table") is not None: + raise SkyflowError(SkyflowMessages.Error.TABLE_NAME_IN_BOTH_PLACES_IN_UPDATE.value, invalid_input_error_code) + else: + for record in request.records: + if record.get("table") is None: + raise SkyflowError(SkyflowMessages.Error.TABLE_NAME_MISSING_IN_UPDATE.value, invalid_input_error_code) + + +def validate_delete_request(logger, request): + if not request.table: + raise SkyflowError(SkyflowMessages.Error.MISSING_TABLE_NAME_IN_DELETE.value, invalid_input_error_code) + + if not request.ids and not request.unique_values: + raise SkyflowError(SkyflowMessages.Error.MISSING_IDS_OR_UNIQUE_VALUES_IN_DELETE.value, invalid_input_error_code) + + if request.ids is not None: + validate_non_empty_string_list(logger, request.ids, SkyflowMessages.Error.INVALID_IDS_IN_DELETE.value) + + +def validate_detokenize_request(logger, request): + if ( + not isinstance(request.tokens, list) or not all(isinstance(t, str) and t.strip() for t in request.tokens) + ): + raise SkyflowError(SkyflowMessages.Error.INVALID_TOKENS_TYPE_IN_DETOKENIZE.value, invalid_input_error_code) + + if not request.tokens: + raise SkyflowError(SkyflowMessages.Error.EMPTY_TOKENS_IN_DETOKENIZE.value, invalid_input_error_code) + + if request.token_group_redactions is not None: + valid = ( + isinstance(request.token_group_redactions, list) + and all( + isinstance(entry, dict) and isinstance(entry.get("token_group_name"), str) and entry.get("token_group_name").strip() + for entry in request.token_group_redactions + ) + ) + if not valid: + raise SkyflowError(SkyflowMessages.Error.INVALID_TOKEN_GROUP_REDACTIONS_IN_DETOKENIZE.value, invalid_input_error_code) + + +def validate_tokenize_request(logger, request): + if not isinstance(request.values, list) or not all(isinstance(v, dict) for v in request.values): + raise SkyflowError(SkyflowMessages.Error.INVALID_VALUES_TYPE_IN_TOKENIZE.value, invalid_input_error_code) + + if not request.values: + raise SkyflowError(SkyflowMessages.Error.EMPTY_VALUES_IN_TOKENIZE.value, invalid_input_error_code) + + for value in request.values: + token_group_names = value.get("token_group_names") + validate_non_empty_string_list( + logger, token_group_names, SkyflowMessages.Error.MISSING_TOKEN_GROUP_NAMES_IN_TOKENIZE.value + ) diff --git a/flowvault/skyflow_flowvault/vault/client/client.py b/flowvault/skyflow_flowvault/vault/client/client.py index 5dc4c47..f206fe8 100644 --- a/flowvault/skyflow_flowvault/vault/client/client.py +++ b/flowvault/skyflow_flowvault/vault/client/client.py @@ -10,5 +10,5 @@ def resolve_vault_url(self, cluster_id, env, vault_id, logger=None): def initialize_api_client(self, vault_url, bearer_token): self._api_client = SkyflowAuth(base_url=vault_url) - def get_insert_api(self): + def get_flowservice_api(self): return self._api_client.flowservice diff --git a/flowvault/skyflow_flowvault/vault/controller/_vault.py b/flowvault/skyflow_flowvault/vault/controller/_vault.py index a6abaa4..abb9563 100644 --- a/flowvault/skyflow_flowvault/vault/controller/_vault.py +++ b/flowvault/skyflow_flowvault/vault/controller/_vault.py @@ -4,11 +4,39 @@ from common.utils.constants import SKY_META_DATA_HEADER from common.utils.logger import log_info, log_error_log from common.vault.base_vault_controller import BaseVaultController -from skyflow_flowvault.generated.rest import V1InsertRecordData, V1Upsert +from skyflow_flowvault.generated.rest import ( + V1ColumnRedactions, + V1FlowTokenizeRequestObject, + V1InsertRecordData, + V1TokenGroupRedactions, + V1UniqueValue, + V1UpdateRecordData, + V1Upsert, +) from skyflow_flowvault.generated.rest.core import ApiError from skyflow_flowvault.utils import SkyflowMessages, get_metrics -from skyflow_flowvault.utils.validations import validate_insert_request -from skyflow_flowvault.vault.data import InsertRequest, InsertResponse +from skyflow_flowvault.utils.validations import ( + validate_insert_request, + validate_get_request, + validate_update_request, + validate_delete_request, + validate_detokenize_request, + validate_tokenize_request, +) +from skyflow_flowvault.vault.data import ( + InsertRequest, + InsertResponse, + GetRequest, + GetResponse, + UpdateRequest, + UpdateResponse, + DeleteRequest, + DeleteResponse, + DetokenizeRequest, + DetokenizeResponse, + TokenizeRequest, + TokenizeResponse, +) REQUEST_ID_HEADER = "x-request-id" @@ -29,7 +57,7 @@ def insert(self, request: InsertRequest) -> InsertResponse: log_info(SkyflowMessages.Info.INSERT_REQUEST_RESOLVED.value, self._vault_client.get_logger()) self._vault_client.initialize_client_configuration() - insert_api = self._vault_client.get_insert_api() + insert_api = self._vault_client.get_flowservice_api() needs_per_record_table = any(r.get("table") is not None for r in request.values) needs_per_record_upsert = any(r.get("upsert") is not None for r in request.values) @@ -61,20 +89,164 @@ def insert(self, request: InsertRequest) -> InsertResponse: log_info(SkyflowMessages.Info.INSERT_SUCCESS.value, self._vault_client.get_logger()) return InsertResponse(inserted_fields=inserted_fields, errors=errors if errors else None) - def get(self, request): - raise NotImplementedError("VaultController.get is not implemented yet") + def get(self, request: GetRequest) -> GetResponse: + log_info(SkyflowMessages.Info.VALIDATE_GET_REQUEST.value, self._vault_client.get_logger()) + validate_get_request(self._vault_client.get_logger(), request) + self._validate_table_name_if_present(request.table) + log_info(SkyflowMessages.Info.GET_REQUEST_RESOLVED.value, self._vault_client.get_logger()) + self._vault_client.initialize_client_configuration() + + flowservice_api = self._vault_client.get_flowservice_api() + items = request.ids or request.unique_values or [] + + try: + log_info(SkyflowMessages.Info.GET_TRIGGERED.value, self._vault_client.get_logger()) + raw_response = flowservice_api.with_raw_response.get( + vault_id=self._vault_client.get_vault_id(), + table_name=request.table, + skyflow_i_ds=request.ids, + unique_values=self.__to_v1_unique_values(request.unique_values), + columns=request.columns, + column_redactions=self.__to_v1_column_redactions(request.column_redactions), + limit=request.limit, + offset=request.offset, + request_options={'additional_headers': self.__build_headers()}, + ) + request_id = self.__extract_request_id(raw_response.headers) + records, errors = self.__split_success_and_errors( + raw_response.data.records or [], 0, request_id, include_data=True, + ) + except Exception as e: + log_error_log(SkyflowMessages.ErrorLogs.GET_RECORDS_REJECTED.value, self._vault_client.get_logger()) + records, errors = [], self.__errors_from_exception(e, items, 0) + + log_info(SkyflowMessages.Info.GET_SUCCESS.value, self._vault_client.get_logger()) + return GetResponse(records=records, errors=errors if errors else None) + + def update(self, request: UpdateRequest) -> UpdateResponse: + log_info(SkyflowMessages.Info.VALIDATE_UPDATE_REQUEST.value, self._vault_client.get_logger()) + validate_update_request(self._vault_client.get_logger(), request) + self._validate_table_name_if_present(request.table) + for record in request.records: + self._validate_table_name_if_present(record.get("table")) + if record.get("values") is not None: + self._validate_field_values(record.get("values")) + log_info(SkyflowMessages.Info.UPDATE_REQUEST_RESOLVED.value, self._vault_client.get_logger()) + self._vault_client.initialize_client_configuration() + + flowservice_api = self._vault_client.get_flowservice_api() + + needs_per_record_table = any(r.get("table") is not None for r in request.records) + + wire_records = [ + self.__build_update_wire_record(record, request, needs_per_record_table) + for record in request.records + ] + + try: + log_info(SkyflowMessages.Info.UPDATE_TRIGGERED.value, self._vault_client.get_logger()) + top_level_kwargs = self.__omit_none( + table_name=None if needs_per_record_table else request.table, + update_type=request.update_type.value if request.update_type else None, + ) + raw_response = flowservice_api.with_raw_response.update( + vault_id=self._vault_client.get_vault_id(), + records=wire_records, + request_options={'additional_headers': self.__build_headers()}, + **top_level_kwargs, + ) + request_id = self.__extract_request_id(raw_response.headers) + records, errors = self.__split_success_and_errors( + raw_response.data.records or [], 0, request_id, include_data=True, + ) + except Exception as e: + log_error_log(SkyflowMessages.ErrorLogs.UPDATE_RECORDS_REJECTED.value, self._vault_client.get_logger()) + records, errors = [], self.__errors_from_exception(e, request.records, 0) - def update(self, request): - raise NotImplementedError("VaultController.update is not implemented yet") + log_info(SkyflowMessages.Info.UPDATE_SUCCESS.value, self._vault_client.get_logger()) + return UpdateResponse(records=records, errors=errors if errors else None) + + def delete(self, request: DeleteRequest) -> DeleteResponse: + log_info(SkyflowMessages.Info.VALIDATE_DELETE_REQUEST.value, self._vault_client.get_logger()) + validate_delete_request(self._vault_client.get_logger(), request) + self._validate_table_name_if_present(request.table) + log_info(SkyflowMessages.Info.DELETE_REQUEST_RESOLVED.value, self._vault_client.get_logger()) + self._vault_client.initialize_client_configuration() + + flowservice_api = self._vault_client.get_flowservice_api() + items = request.ids or request.unique_values or [] + + try: + log_info(SkyflowMessages.Info.DELETE_TRIGGERED.value, self._vault_client.get_logger()) + raw_response = flowservice_api.with_raw_response.delete( + vault_id=self._vault_client.get_vault_id(), + table_name=request.table, + skyflow_i_ds=request.ids, + unique_values=self.__to_v1_unique_values(request.unique_values), + request_options={'additional_headers': self.__build_headers()}, + ) + request_id = self.__extract_request_id(raw_response.headers) + records, errors = self.__split_success_and_errors(raw_response.data.records or [], 0, request_id) + except Exception as e: + log_error_log(SkyflowMessages.ErrorLogs.DELETE_RECORDS_REJECTED.value, self._vault_client.get_logger()) + records, errors = [], self.__errors_from_exception(e, items, 0) - def delete(self, request): - raise NotImplementedError("VaultController.delete is not implemented yet") + log_info(SkyflowMessages.Info.DELETE_SUCCESS.value, self._vault_client.get_logger()) + return DeleteResponse(records=records, errors=errors if errors else None) def query(self, request): raise NotImplementedError("VaultController.query is not implemented yet") - def detokenize(self, request): - raise NotImplementedError("VaultController.detokenize is not implemented yet") + def detokenize(self, request: DetokenizeRequest) -> DetokenizeResponse: + log_info(SkyflowMessages.Info.VALIDATE_DETOKENIZE_REQUEST.value, self._vault_client.get_logger()) + validate_detokenize_request(self._vault_client.get_logger(), request) + log_info(SkyflowMessages.Info.DETOKENIZE_REQUEST_RESOLVED.value, self._vault_client.get_logger()) + self._vault_client.initialize_client_configuration() + + flowservice_api = self._vault_client.get_flowservice_api() + + try: + log_info(SkyflowMessages.Info.DETOKENIZE_TRIGGERED.value, self._vault_client.get_logger()) + raw_response = flowservice_api.with_raw_response.detokenize( + vault_id=self._vault_client.get_vault_id(), + tokens=request.tokens, + token_group_redactions=self.__to_v1_token_group_redactions(request.token_group_redactions), + request_options={'additional_headers': self.__build_headers()}, + ) + request_id = self.__extract_request_id(raw_response.headers) + records, errors = self.__split_detokenize_success_and_errors(raw_response.data.response or [], 0, request_id) + except Exception as e: + log_error_log(SkyflowMessages.ErrorLogs.DETOKENIZE_RECORDS_REJECTED.value, self._vault_client.get_logger()) + records, errors = [], self.__errors_from_exception(e, request.tokens, 0) + + log_info(SkyflowMessages.Info.DETOKENIZE_SUCCESS.value, self._vault_client.get_logger()) + return DetokenizeResponse(records=records, errors=errors if errors else None) + + def tokenize(self, request: TokenizeRequest) -> TokenizeResponse: + log_info(SkyflowMessages.Info.VALIDATE_TOKENIZE_REQUEST.value, self._vault_client.get_logger()) + validate_tokenize_request(self._vault_client.get_logger(), request) + log_info(SkyflowMessages.Info.TOKENIZE_REQUEST_RESOLVED.value, self._vault_client.get_logger()) + self._vault_client.initialize_client_configuration() + + flowservice_api = self._vault_client.get_flowservice_api() + + wire_values = [self.__build_tokenize_wire_value(value) for value in request.values] + + try: + log_info(SkyflowMessages.Info.TOKENIZE_TRIGGERED.value, self._vault_client.get_logger()) + raw_response = flowservice_api.with_raw_response.tokenize( + vault_id=self._vault_client.get_vault_id(), + data=wire_values, + request_options={'additional_headers': self.__build_headers()}, + ) + request_id = self.__extract_request_id(raw_response.headers) + records, errors = self.__split_tokenize_success_and_errors(raw_response.data.response or [], 0, request_id) + except Exception as e: + log_error_log(SkyflowMessages.ErrorLogs.TOKENIZE_RECORDS_REJECTED.value, self._vault_client.get_logger()) + records, errors = [], self.__errors_from_exception(e, request.values, 0) + + log_info(SkyflowMessages.Info.TOKENIZE_SUCCESS.value, self._vault_client.get_logger()) + return TokenizeResponse(records=records, errors=errors if errors else None) def __build_wire_record(self, record, request, needs_per_record_table, needs_per_record_upsert): return V1InsertRecordData(data=record["values"], **self.__omit_none( @@ -82,6 +254,23 @@ def __build_wire_record(self, record, request, needs_per_record_table, needs_per upsert=self.__to_v1_upsert(record.get("upsert") or request.upsert) if needs_per_record_upsert else None, )) + def __build_update_wire_record(self, record, request, needs_per_record_table): + return V1UpdateRecordData( + skyflow_id=record.get("skyflow_id"), + data=record.get("values"), + **self.__omit_none( + tokens=record.get("tokens"), + table_name=(record.get("table") or request.table) if needs_per_record_table else None, + ), + ) + + def __build_tokenize_wire_value(self, value): + return V1FlowTokenizeRequestObject( + value=value.get("value"), + token_group_names=value.get("token_group_names"), + **self.__omit_none(token=value.get("token")), + ) + def __omit_none(self, **kwargs): return {k: v for k, v in kwargs.items() if v is not None} @@ -101,11 +290,31 @@ def __to_v1_upsert(self, upsert): unique_columns=upsert.get("unique_columns"), ) + def __to_v1_unique_values(self, unique_values): + if unique_values is None: + return None + return [V1UniqueValue(data=value) for value in unique_values] + + def __to_v1_column_redactions(self, column_redactions): + if column_redactions is None: + return None + return [ + V1ColumnRedactions(column_name=entry.get("column_name"), redaction=entry.get("redaction")) + for entry in column_redactions + ] + + def __to_v1_token_group_redactions(self, token_group_redactions): + if token_group_redactions is None: + return None + return [ + V1TokenGroupRedactions(token_group_name=entry.get("token_group_name"), redaction=entry.get("redaction")) + for entry in token_group_redactions + ] + def __extract_request_id(self, headers): return headers.get(REQUEST_ID_HEADER) if headers else None - def __split_success_and_errors(self, records, start_index, request_id): - + def __split_success_and_errors(self, records, start_index, request_id, include_data=False): successes, errors = [], [] for offset, record in enumerate(records): request_index = start_index + offset @@ -116,10 +325,54 @@ def __split_success_and_errors(self, records, start_index, request_id): 'request_index': request_index, 'skyflow_id': record.skyflow_id, } - success.update(self.__flatten_tokens(record.tokens)) + success.update(self.__flatten_tokens(getattr(record, 'tokens', None))) + if include_data: + data = getattr(record, 'data', None) + if data: + success['data'] = data + hashed_data = getattr(record, 'hashed_data', None) + if hashed_data: + success['hashed_data'] = hashed_data successes.append(success) return successes, errors + def __split_detokenize_success_and_errors(self, responses, start_index, request_id): + successes, errors = [], [] + for offset, resp in enumerate(responses): + request_index = start_index + offset + if resp.error is not None: + errors.append({ + 'request_index': request_index, 'token': resp.token, 'error': resp.error, + 'code': resp.http_code, 'request_id': request_id, + }) + else: + successes.append({ + 'request_index': request_index, + 'token': resp.token, + 'value': resp.value, + 'token_group_name': resp.token_group_name, + }) + return successes, errors + + def __split_tokenize_success_and_errors(self, responses, start_index, request_id): + successes, errors = [], [] + for offset, resp in enumerate(responses): + request_index = start_index + offset + for token in (resp.tokens or []): + if token.error is not None: + errors.append({ + 'request_index': request_index, 'token_group_name': token.token_group_name, + 'error': token.error, 'code': token.http_code, 'request_id': request_id, + }) + else: + successes.append({ + 'request_index': request_index, + 'value': resp.value, + 'token_group_name': token.token_group_name, + 'token': token.token, + }) + return successes, errors + def __flatten_tokens(self, tokens): if not tokens: return {} diff --git a/flowvault/skyflow_flowvault/vault/data/__init__.py b/flowvault/skyflow_flowvault/vault/data/__init__.py index 62ae85c..849d169 100644 --- a/flowvault/skyflow_flowvault/vault/data/__init__.py +++ b/flowvault/skyflow_flowvault/vault/data/__init__.py @@ -1,3 +1,13 @@ from ._insert_request import InsertRequest from ._insert_response import InsertResponse from ._upsert import Upsert +from ._get_request import GetRequest +from ._get_response import GetResponse +from ._update_request import UpdateRequest +from ._update_response import UpdateResponse +from ._delete_request import DeleteRequest +from ._delete_response import DeleteResponse +from ._detokenize_request import DetokenizeRequest +from ._detokenize_response import DetokenizeResponse +from ._tokenize_request import TokenizeRequest +from ._tokenize_response import TokenizeResponse diff --git a/flowvault/skyflow_flowvault/vault/data/_delete_request.py b/flowvault/skyflow_flowvault/vault/data/_delete_request.py new file mode 100644 index 0000000..b985fef --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_delete_request.py @@ -0,0 +1,5 @@ +class DeleteRequest: + def __init__(self, table: str, ids: list = None, unique_values: list = None): + self.table = table + self.ids = ids + self.unique_values = unique_values diff --git a/flowvault/skyflow_flowvault/vault/data/_delete_response.py b/flowvault/skyflow_flowvault/vault/data/_delete_response.py new file mode 100644 index 0000000..a33a76e --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_delete_response.py @@ -0,0 +1,10 @@ +class DeleteResponse: + def __init__(self, records=None, errors=None): + self.records = records + self.errors = errors + + def __repr__(self): + return f"DeleteResponse(records={self.records}, errors={self.errors})" + + def __str__(self): + return self.__repr__() diff --git a/flowvault/skyflow_flowvault/vault/data/_detokenize_request.py b/flowvault/skyflow_flowvault/vault/data/_detokenize_request.py new file mode 100644 index 0000000..fae7271 --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_detokenize_request.py @@ -0,0 +1,4 @@ +class DetokenizeRequest: + def __init__(self, tokens: list, token_group_redactions: list = None): + self.tokens = tokens + self.token_group_redactions = token_group_redactions diff --git a/flowvault/skyflow_flowvault/vault/data/_detokenize_response.py b/flowvault/skyflow_flowvault/vault/data/_detokenize_response.py new file mode 100644 index 0000000..82997f6 --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_detokenize_response.py @@ -0,0 +1,10 @@ +class DetokenizeResponse: + def __init__(self, records=None, errors=None): + self.records = records + self.errors = errors + + def __repr__(self): + return f"DetokenizeResponse(records={self.records}, errors={self.errors})" + + def __str__(self): + return self.__repr__() diff --git a/flowvault/skyflow_flowvault/vault/data/_get_request.py b/flowvault/skyflow_flowvault/vault/data/_get_request.py new file mode 100644 index 0000000..ea3f33b --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_get_request.py @@ -0,0 +1,10 @@ +class GetRequest: + def __init__(self, table: str, ids: list = None, unique_values: list = None, columns: list = None, + column_redactions: list = None, limit: int = None, offset: int = None): + self.table = table + self.ids = ids + self.unique_values = unique_values + self.columns = columns + self.column_redactions = column_redactions + self.limit = limit + self.offset = offset diff --git a/flowvault/skyflow_flowvault/vault/data/_get_response.py b/flowvault/skyflow_flowvault/vault/data/_get_response.py new file mode 100644 index 0000000..5e431e8 --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_get_response.py @@ -0,0 +1,10 @@ +class GetResponse: + def __init__(self, records=None, errors=None): + self.records = records + self.errors = errors + + def __repr__(self): + return f"GetResponse(records={self.records}, errors={self.errors})" + + def __str__(self): + return self.__repr__() diff --git a/flowvault/skyflow_flowvault/vault/data/_tokenize_request.py b/flowvault/skyflow_flowvault/vault/data/_tokenize_request.py new file mode 100644 index 0000000..ace7861 --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_tokenize_request.py @@ -0,0 +1,3 @@ +class TokenizeRequest: + def __init__(self, values: list): + self.values = values diff --git a/flowvault/skyflow_flowvault/vault/data/_tokenize_response.py b/flowvault/skyflow_flowvault/vault/data/_tokenize_response.py new file mode 100644 index 0000000..77689ec --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_tokenize_response.py @@ -0,0 +1,10 @@ +class TokenizeResponse: + def __init__(self, records=None, errors=None): + self.records = records + self.errors = errors + + def __repr__(self): + return f"TokenizeResponse(records={self.records}, errors={self.errors})" + + def __str__(self): + return self.__repr__() diff --git a/flowvault/skyflow_flowvault/vault/data/_update_request.py b/flowvault/skyflow_flowvault/vault/data/_update_request.py new file mode 100644 index 0000000..09c9802 --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_update_request.py @@ -0,0 +1,5 @@ +class UpdateRequest: + def __init__(self, records: list, table: str = None, update_type=None): + self.records = records + self.table = table + self.update_type = update_type diff --git a/flowvault/skyflow_flowvault/vault/data/_update_response.py b/flowvault/skyflow_flowvault/vault/data/_update_response.py new file mode 100644 index 0000000..2b07fe3 --- /dev/null +++ b/flowvault/skyflow_flowvault/vault/data/_update_response.py @@ -0,0 +1,10 @@ +class UpdateResponse: + def __init__(self, records=None, errors=None): + self.records = records + self.errors = errors + + def __repr__(self): + return f"UpdateResponse(records={self.records}, errors={self.errors})" + + def __str__(self): + return self.__repr__() diff --git a/flowvault/tests/utils/validations/test__validations.py b/flowvault/tests/utils/validations/test__validations.py index 10b36b6..e7038e0 100644 --- a/flowvault/tests/utils/validations/test__validations.py +++ b/flowvault/tests/utils/validations/test__validations.py @@ -3,8 +3,23 @@ from common.errors import SkyflowError from common.utils.enums import Env from skyflow_flowvault.utils.enums import UpsertType -from skyflow_flowvault.utils.validations import validate_insert_request, validate_vault_config -from skyflow_flowvault.vault.data import InsertRequest +from skyflow_flowvault.utils.validations import ( + validate_insert_request, + validate_get_request, + validate_update_request, + validate_delete_request, + validate_detokenize_request, + validate_tokenize_request, + validate_vault_config, +) +from skyflow_flowvault.vault.data import ( + InsertRequest, + GetRequest, + UpdateRequest, + DeleteRequest, + DetokenizeRequest, + TokenizeRequest, +) class TestValidateInsertRequest(unittest.TestCase): @@ -156,6 +171,215 @@ def test_per_record_upsert_is_also_validated(self): validate_insert_request(None, request) +class TestValidateGetRequest(unittest.TestCase): + def test_valid_request_with_ids(self): + request = GetRequest(table="t1", ids=["id1"]) + validate_get_request(None, request) # should not raise + + def test_valid_request_with_unique_values(self): + request = GetRequest(table="t1", unique_values=[{"email": "a@b.com"}]) + validate_get_request(None, request) # should not raise + + def test_missing_table_raises(self): + request = GetRequest(table=None, ids=["id1"]) + with self.assertRaises(SkyflowError): + validate_get_request(None, request) + + def test_empty_table_raises(self): + request = GetRequest(table="", ids=["id1"]) + with self.assertRaises(SkyflowError): + validate_get_request(None, request) + + def test_missing_ids_and_unique_values_raises(self): + request = GetRequest(table="t1") + with self.assertRaises(SkyflowError): + validate_get_request(None, request) + + def test_ids_must_be_a_list(self): + request = GetRequest(table="t1", ids="not-a-list") + with self.assertRaises(SkyflowError): + validate_get_request(None, request) + + def test_ids_must_be_non_empty(self): + request = GetRequest(table="t1", ids=[]) + with self.assertRaises(SkyflowError): + validate_get_request(None, request) + + def test_ids_must_be_strings(self): + request = GetRequest(table="t1", ids=[123]) + with self.assertRaises(SkyflowError): + validate_get_request(None, request) + + +class TestValidateUpdateRequest(unittest.TestCase): + def test_valid_request_with_request_level_table(self): + request = UpdateRequest(records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1") + validate_update_request(None, request) # should not raise + + def test_valid_request_with_per_record_table(self): + request = UpdateRequest(records=[{"skyflow_id": "id1", "values": {"a": 1}, "table": "t1"}]) + validate_update_request(None, request) # should not raise + + def test_records_must_be_a_list(self): + request = UpdateRequest(records="not-a-list", table="t1") + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_records_must_not_be_empty(self): + request = UpdateRequest(records=[], table="t1") + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_missing_skyflow_id_raises(self): + request = UpdateRequest(records=[{"values": {"a": 1}}], table="t1") + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_empty_skyflow_id_raises(self): + request = UpdateRequest(records=[{"skyflow_id": " ", "values": {"a": 1}}], table="t1") + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_record_with_unknown_key_raises(self): + request = UpdateRequest(records=[{"skyflow_id": "id1", "unexpected": 1}], table="t1") + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_table_in_both_places_raises(self): + request = UpdateRequest(records=[{"skyflow_id": "id1", "values": {"a": 1}, "table": "t2"}], table="t1") + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_table_missing_from_one_record_raises(self): + request = UpdateRequest(records=[ + {"skyflow_id": "id1", "values": {"a": 1}, "table": "t1"}, + {"skyflow_id": "id2", "values": {"a": 2}}, + ]) + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_invalid_update_type_raises(self): + request = UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1", update_type="REPLACE", + ) + with self.assertRaises(SkyflowError): + validate_update_request(None, request) + + def test_valid_update_type_enum_is_valid(self): + request = UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1", update_type=UpsertType.REPLACE, + ) + validate_update_request(None, request) # should not raise + + +class TestValidateDeleteRequest(unittest.TestCase): + def test_valid_request_with_ids(self): + request = DeleteRequest(table="t1", ids=["id1"]) + validate_delete_request(None, request) # should not raise + + def test_valid_request_with_unique_values(self): + request = DeleteRequest(table="t1", unique_values=[{"email": "a@b.com"}]) + validate_delete_request(None, request) # should not raise + + def test_missing_table_raises(self): + request = DeleteRequest(table=None, ids=["id1"]) + with self.assertRaises(SkyflowError): + validate_delete_request(None, request) + + def test_missing_ids_and_unique_values_raises(self): + request = DeleteRequest(table="t1") + with self.assertRaises(SkyflowError): + validate_delete_request(None, request) + + def test_ids_must_be_non_empty(self): + request = DeleteRequest(table="t1", ids=[]) + with self.assertRaises(SkyflowError): + validate_delete_request(None, request) + + def test_ids_must_be_strings(self): + request = DeleteRequest(table="t1", ids=[123]) + with self.assertRaises(SkyflowError): + validate_delete_request(None, request) + + +class TestValidateDetokenizeRequest(unittest.TestCase): + def test_valid_request(self): + request = DetokenizeRequest(tokens=["tok1", "tok2"]) + validate_detokenize_request(None, request) # should not raise + + def test_valid_request_with_token_group_redactions(self): + request = DetokenizeRequest( + tokens=["tok1"], token_group_redactions=[{"token_group_name": "g1", "redaction": "mask1"}], + ) + validate_detokenize_request(None, request) # should not raise + + def test_tokens_must_be_a_list(self): + request = DetokenizeRequest(tokens="not-a-list") + with self.assertRaises(SkyflowError): + validate_detokenize_request(None, request) + + def test_tokens_must_not_be_empty(self): + request = DetokenizeRequest(tokens=[]) + with self.assertRaises(SkyflowError): + validate_detokenize_request(None, request) + + def test_tokens_must_be_strings(self): + request = DetokenizeRequest(tokens=[123]) + with self.assertRaises(SkyflowError): + validate_detokenize_request(None, request) + + def test_empty_string_token_raises(self): + request = DetokenizeRequest(tokens=[" "]) + with self.assertRaises(SkyflowError): + validate_detokenize_request(None, request) + + def test_invalid_token_group_redactions_raises(self): + request = DetokenizeRequest(tokens=["tok1"], token_group_redactions=["not-a-dict"]) + with self.assertRaises(SkyflowError): + validate_detokenize_request(None, request) + + def test_token_group_redactions_missing_name_raises(self): + request = DetokenizeRequest(tokens=["tok1"], token_group_redactions=[{"redaction": "mask1"}]) + with self.assertRaises(SkyflowError): + validate_detokenize_request(None, request) + + +class TestValidateTokenizeRequest(unittest.TestCase): + def test_valid_request(self): + request = TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1"]}]) + validate_tokenize_request(None, request) # should not raise + + def test_values_must_be_a_list(self): + request = TokenizeRequest(values="not-a-list") + with self.assertRaises(SkyflowError): + validate_tokenize_request(None, request) + + def test_values_must_not_be_empty(self): + request = TokenizeRequest(values=[]) + with self.assertRaises(SkyflowError): + validate_tokenize_request(None, request) + + def test_values_must_be_dicts(self): + request = TokenizeRequest(values=["not-a-dict"]) + with self.assertRaises(SkyflowError): + validate_tokenize_request(None, request) + + def test_missing_token_group_names_raises(self): + request = TokenizeRequest(values=[{"value": "a@b.com"}]) + with self.assertRaises(SkyflowError): + validate_tokenize_request(None, request) + + def test_empty_token_group_names_raises(self): + request = TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": []}]) + with self.assertRaises(SkyflowError): + validate_tokenize_request(None, request) + + def test_non_string_token_group_names_raises(self): + request = TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": [123]}]) + with self.assertRaises(SkyflowError): + validate_tokenize_request(None, request) + + class TestValidateVaultConfig(unittest.TestCase): def test_valid_config(self): config = { diff --git a/flowvault/tests/vault/client/test__client.py b/flowvault/tests/vault/client/test__client.py index 2accba2..6d9edac 100644 --- a/flowvault/tests/vault/client/test__client.py +++ b/flowvault/tests/vault/client/test__client.py @@ -45,9 +45,9 @@ def test_initialize_api_client_does_not_pass_token(self, mock_skyflow_auth): self.assertEqual(kwargs.get("base_url"), "https://test-vault-url.com") self.assertNotIn("token", kwargs) - def test_get_insert_api_returns_flowservice(self): + def test_get_flowservice_api_returns_flowservice(self): self.vault_client._api_client = MagicMock() - result = self.vault_client.get_insert_api() + result = self.vault_client.get_flowservice_api() self.assertEqual(result, self.vault_client._api_client.flowservice) diff --git a/flowvault/tests/vault/controller/test__vault.py b/flowvault/tests/vault/controller/test__vault.py index 449ede9..9cbf4e1 100644 --- a/flowvault/tests/vault/controller/test__vault.py +++ b/flowvault/tests/vault/controller/test__vault.py @@ -1,23 +1,63 @@ import unittest +from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch from common.errors import SkyflowError from skyflow_flowvault.generated.rest.core import ApiError from skyflow_flowvault.vault.controller import VaultController -from skyflow_flowvault.vault.data import InsertRequest +from skyflow_flowvault.vault.data import ( + InsertRequest, + GetRequest, + UpdateRequest, + DeleteRequest, + DetokenizeRequest, + TokenizeRequest, +) from skyflow_flowvault.utils.enums import UpsertType class FakeRecordResponseObject: - def __init__(self, skyflow_id=None, tokens=None, data=None, error=None, http_code=None, table_name=None): + def __init__(self, skyflow_id=None, tokens=None, data=None, hashed_data=None, error=None, http_code=None, table_name=None): self.skyflow_id = skyflow_id self.tokens = tokens self.data = data + self.hashed_data = hashed_data self.error = error self.http_code = http_code self.table_name = table_name +class FakeDeleteResponseObject: + def __init__(self, skyflow_id=None, error=None, http_code=None): + self.skyflow_id = skyflow_id + self.error = error + self.http_code = http_code + + +class FakeDetokenizeResponseObject: + def __init__(self, token=None, value=None, token_group_name=None, error=None, http_code=None, metadata=None): + self.token = token + self.value = value + self.token_group_name = token_group_name + self.error = error + self.http_code = http_code + self.metadata = metadata + + +class FakeTokenizeResponseObjectToken: + def __init__(self, token_group_name=None, token=None, error=None, http_code=None): + self.token_group_name = token_group_name + self.token = token + self.error = error + self.http_code = http_code + + +class FakeTokenizeResponseObject: + def __init__(self, value=None, tokens=None): + self.value = value + self.tokens = tokens + + class FakeV1InsertResponse: def __init__(self, records): self.records = records @@ -40,7 +80,7 @@ def setUp(self): self.vault_client.get_logger.return_value = Mock() self.vault_client.get_current_bearer_token.return_value = None self.insert_api = MagicMock() - self.vault_client.get_insert_api.return_value = self.insert_api + self.vault_client.get_flowservice_api.return_value = self.insert_api self.vault = VaultController(self.vault_client) # ------------------------------------------------------------------ # @@ -371,5 +411,780 @@ def test_no_authorization_header_when_no_token_available(self): self.assertNotIn("Authorization", headers) +def fake_get_raw_response(records, headers=None): + return SimpleNamespace(data=SimpleNamespace(records=records), headers=headers or {}) + + +class TestVaultGet(unittest.TestCase): + def setUp(self): + self.vault_client = Mock() + self.vault_client.get_vault_id.return_value = "vault123" + self.vault_client.get_logger.return_value = Mock() + self.vault_client.get_current_bearer_token.return_value = None + self.get_api = MagicMock() + self.vault_client.get_flowservice_api.return_value = self.get_api + self.vault = VaultController(self.vault_client) + + # ------------------------------------------------------------------ # + # validation / initialization sequencing + # ------------------------------------------------------------------ # + + @patch("skyflow_flowvault.vault.controller._vault.validate_get_request") + def test_get_validates_before_initializing_client(self, mock_validate): + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([]) + request = GetRequest(table="t1", ids=["id1"]) + + self.vault.get(request) + + mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + self.vault_client.initialize_client_configuration.assert_called_once() + + def test_get_raises_for_invalid_request(self): + with self.assertRaises(SkyflowError): + self.vault.get(GetRequest(table="t1")) + self.vault_client.initialize_client_configuration.assert_not_called() + + def test_get_raises_on_invalid_table_name(self): + with self.assertRaises(SkyflowError): + self.vault.get(GetRequest(table=" ", ids=["id1"])) + self.get_api.with_raw_response.get.assert_not_called() + + # ------------------------------------------------------------------ # + # request -> wire field mapping + # ------------------------------------------------------------------ # + + def test_maps_table_and_ids(self): + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([]) + + self.vault.get(GetRequest(table="t1", ids=["id1", "id2"])) + + _, kwargs = self.get_api.with_raw_response.get.call_args + self.assertEqual(kwargs["vault_id"], "vault123") + self.assertEqual(kwargs["table_name"], "t1") + self.assertEqual(kwargs["skyflow_i_ds"], ["id1", "id2"]) + + def test_maps_unique_values(self): + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([]) + + self.vault.get(GetRequest(table="t1", unique_values=[{"email": "a@b.com"}])) + + _, kwargs = self.get_api.with_raw_response.get.call_args + self.assertEqual(len(kwargs["unique_values"]), 1) + self.assertEqual(kwargs["unique_values"][0].data, {"email": "a@b.com"}) + + def test_maps_column_redactions(self): + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([]) + + self.vault.get(GetRequest( + table="t1", ids=["id1"], column_redactions=[{"column_name": "ssn", "redaction": "mask1"}], + )) + + _, kwargs = self.get_api.with_raw_response.get.call_args + self.assertEqual(len(kwargs["column_redactions"]), 1) + self.assertEqual(kwargs["column_redactions"][0].column_name, "ssn") + self.assertEqual(kwargs["column_redactions"][0].redaction, "mask1") + + def test_maps_limit_offset_columns(self): + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([]) + + self.vault.get(GetRequest(table="t1", ids=["id1"], columns=["a", "b"], limit=10, offset=5)) + + _, kwargs = self.get_api.with_raw_response.get.call_args + self.assertEqual(kwargs["columns"], ["a", "b"]) + self.assertEqual(kwargs["limit"], 10) + self.assertEqual(kwargs["offset"], 5) + + # ------------------------------------------------------------------ # + # response shape -- includes data, unlike insert + # ------------------------------------------------------------------ # + + def test_successful_records_include_data_and_tokens(self): + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([ + FakeRecordResponseObject( + skyflow_id="id1", + tokens={"name": [{"token": "tok1", "tokenGroupName": "deterministic_string"}]}, + data={"name": "john doe"}, + hashed_data={"name": "a1b2c3"}, + table_name="t1", + ), + ], headers={"x-request-id": "req-1"}) + + response = self.vault.get(GetRequest(table="t1", ids=["id1"])) + + self.assertEqual(len(response.records), 1) + record = response.records[0] + self.assertEqual(record["hashed_data"], {"name": "a1b2c3"}) + self.assertEqual(record["request_index"], 0) + self.assertEqual(record["skyflow_id"], "id1") + self.assertEqual(record["name"], "tok1") + self.assertEqual(record["data"], {"name": "john doe"}) + self.assertIsNone(response.errors) + + def test_mixed_success_and_error_records_are_split(self): + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([ + FakeRecordResponseObject(skyflow_id="id1", data={"a": 1}), + FakeRecordResponseObject(error="not found", http_code=404), + ], headers={"x-request-id": "req-2"}) + + response = self.vault.get(GetRequest(table="t1", ids=["id1", "id2"])) + + self.assertEqual(len(response.records), 1) + self.assertEqual(response.records[0]["data"], {"a": 1}) + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "not found") + self.assertEqual(response.errors[0]["code"], 404) + self.assertEqual(response.errors[0]["request_id"], "req-2") + + # ------------------------------------------------------------------ # + # transport failure + # ------------------------------------------------------------------ # + + def test_transport_exception_marks_every_id_as_an_error(self): + self.get_api.with_raw_response.get.side_effect = Exception("network blip") + + response = self.vault.get(GetRequest(table="t1", ids=["id1", "id2"])) + + self.assertEqual(len(response.records), 0) + self.assertEqual(len(response.errors), 2) + self.assertTrue(all("network blip" in e["error"] for e in response.errors)) + + def test_api_error_with_structured_body_splits_into_one_error_per_row(self): + api_error = ApiError( + status_code=404, + headers={"x-request-id": "req-3"}, + body={"records": [{"error": "not found", "httpCode": 404}]}, + ) + self.get_api.with_raw_response.get.side_effect = api_error + + response = self.vault.get(GetRequest(table="t1", ids=["id1"])) + + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "not found") + self.assertEqual(response.errors[0]["code"], 404) + self.assertEqual(response.errors[0]["request_id"], "req-3") + + # ------------------------------------------------------------------ # + # per-call Authorization header injection + # ------------------------------------------------------------------ # + + def test_injects_authorization_header_from_current_bearer_token(self): + self.vault_client.get_current_bearer_token.return_value = "the-current-token" + self.get_api.with_raw_response.get.return_value = fake_get_raw_response([]) + + self.vault.get(GetRequest(table="t1", ids=["id1"])) + + _, kwargs = self.get_api.with_raw_response.get.call_args + headers = kwargs["request_options"]["additional_headers"] + self.assertEqual(headers.get("Authorization"), "Bearer the-current-token") + + +def fake_update_raw_response(records, headers=None): + return SimpleNamespace(data=SimpleNamespace(records=records), headers=headers or {}) + + +class TestVaultUpdate(unittest.TestCase): + def setUp(self): + self.vault_client = Mock() + self.vault_client.get_vault_id.return_value = "vault123" + self.vault_client.get_logger.return_value = Mock() + self.vault_client.get_current_bearer_token.return_value = None + self.update_api = MagicMock() + self.vault_client.get_flowservice_api.return_value = self.update_api + self.vault = VaultController(self.vault_client) + + # ------------------------------------------------------------------ # + # validation / initialization sequencing + # ------------------------------------------------------------------ # + + @patch("skyflow_flowvault.vault.controller._vault.validate_update_request") + def test_update_validates_before_initializing_client(self, mock_validate): + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([]) + request = UpdateRequest(records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1") + + self.vault.update(request) + + mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + self.vault_client.initialize_client_configuration.assert_called_once() + + def test_update_raises_for_invalid_request(self): + with self.assertRaises(SkyflowError): + self.vault.update(UpdateRequest(records=[], table="t1")) + self.vault_client.initialize_client_configuration.assert_not_called() + + def test_update_raises_on_empty_key(self): + with self.assertRaises(SkyflowError): + self.vault.update(UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"": "value"}}], table="t1", + )) + self.update_api.with_raw_response.update.assert_not_called() + + def test_update_raises_on_invalid_table_name(self): + with self.assertRaises(SkyflowError): + self.vault.update(UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"a": 1}}], table=" ", + )) + + # ------------------------------------------------------------------ # + # request -> wire field mapping + # ------------------------------------------------------------------ # + + def test_maps_request_level_table_and_update_type(self): + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([]) + request = UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1", update_type=UpsertType.REPLACE, + ) + + self.vault.update(request) + + _, kwargs = self.update_api.with_raw_response.update.call_args + self.assertEqual(kwargs["vault_id"], "vault123") + self.assertEqual(kwargs["table_name"], "t1") + self.assertEqual(kwargs["update_type"], "REPLACE") + self.assertEqual(len(kwargs["records"]), 1) + self.assertEqual(kwargs["records"][0].skyflow_id, "id1") + self.assertEqual(kwargs["records"][0].data, {"a": 1}) + self.assertIsNone(kwargs["records"][0].table_name) + + def test_maps_per_record_table_when_request_level_unset(self): + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([]) + request = UpdateRequest(records=[ + {"skyflow_id": "id1", "values": {"a": 1}, "table": "t2"}, + ]) + + self.vault.update(request) + + _, kwargs = self.update_api.with_raw_response.update.call_args + self.assertNotIn("table_name", kwargs) + self.assertEqual(kwargs["records"][0].table_name, "t2") + + def test_maps_per_record_tokens(self): + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([]) + request = UpdateRequest(records=[ + {"skyflow_id": "id1", "values": {"a": 1}, "tokens": {"a": "tok1"}, "table": "t1"}, + ]) + + self.vault.update(request) + + _, kwargs = self.update_api.with_raw_response.update.call_args + self.assertEqual(kwargs["records"][0].tokens, {"a": "tok1"}) + + def test_no_update_type_is_omitted_not_sent_as_none(self): + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([]) + request = UpdateRequest(records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1") + + self.vault.update(request) + + _, kwargs = self.update_api.with_raw_response.update.call_args + self.assertNotIn("update_type", kwargs) + + # ------------------------------------------------------------------ # + # response shape -- includes data, like get + # ------------------------------------------------------------------ # + + def test_successful_records_include_data_and_tokens(self): + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([ + FakeRecordResponseObject( + skyflow_id="id1", + tokens={"name": [{"token": "tok1", "tokenGroupName": "deterministic_string"}]}, + data={"name": "john doe"}, + ), + ], headers={"x-request-id": "req-1"}) + + response = self.vault.update(UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"name": "john doe"}}], table="t1", + )) + + self.assertEqual(len(response.records), 1) + record = response.records[0] + self.assertEqual(record["skyflow_id"], "id1") + self.assertEqual(record["name"], "tok1") + self.assertEqual(record["data"], {"name": "john doe"}) + self.assertIsNone(response.errors) + + def test_mixed_success_and_error_records_are_split(self): + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([ + FakeRecordResponseObject(skyflow_id="id1", data={"a": 1}), + FakeRecordResponseObject(error="not found", http_code=404), + ], headers={"x-request-id": "req-2"}) + + response = self.vault.update(UpdateRequest(records=[ + {"skyflow_id": "id1", "values": {"a": 1}}, + {"skyflow_id": "id2", "values": {"a": 2}}, + ], table="t1")) + + self.assertEqual(len(response.records), 1) + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "not found") + self.assertEqual(response.errors[0]["code"], 404) + + # ------------------------------------------------------------------ # + # transport failure + # ------------------------------------------------------------------ # + + def test_transport_exception_marks_every_record_as_an_error(self): + self.update_api.with_raw_response.update.side_effect = Exception("network blip") + records = [{"skyflow_id": "id1", "values": {"a": 1}}, {"skyflow_id": "id2", "values": {"a": 2}}] + + response = self.vault.update(UpdateRequest(records=records, table="t1")) + + self.assertEqual(len(response.records), 0) + self.assertEqual(len(response.errors), 2) + self.assertTrue(all("network blip" in e["error"] for e in response.errors)) + + def test_api_error_with_structured_body_splits_into_one_error_per_row(self): + api_error = ApiError( + status_code=404, + headers={"x-request-id": "req-3"}, + body={"records": [{"error": "not found", "httpCode": 404}]}, + ) + self.update_api.with_raw_response.update.side_effect = api_error + + response = self.vault.update(UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1", + )) + + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "not found") + self.assertEqual(response.errors[0]["code"], 404) + self.assertEqual(response.errors[0]["request_id"], "req-3") + + # ------------------------------------------------------------------ # + # per-call Authorization header injection + # ------------------------------------------------------------------ # + + def test_injects_authorization_header_from_current_bearer_token(self): + self.vault_client.get_current_bearer_token.return_value = "the-current-token" + self.update_api.with_raw_response.update.return_value = fake_update_raw_response([]) + + self.vault.update(UpdateRequest(records=[{"skyflow_id": "id1", "values": {"a": 1}}], table="t1")) + + _, kwargs = self.update_api.with_raw_response.update.call_args + headers = kwargs["request_options"]["additional_headers"] + self.assertEqual(headers.get("Authorization"), "Bearer the-current-token") + + +def fake_delete_raw_response(records, headers=None): + return SimpleNamespace(data=SimpleNamespace(records=records), headers=headers or {}) + + +class TestVaultDelete(unittest.TestCase): + def setUp(self): + self.vault_client = Mock() + self.vault_client.get_vault_id.return_value = "vault123" + self.vault_client.get_logger.return_value = Mock() + self.vault_client.get_current_bearer_token.return_value = None + self.delete_api = MagicMock() + self.vault_client.get_flowservice_api.return_value = self.delete_api + self.vault = VaultController(self.vault_client) + + # ------------------------------------------------------------------ # + # validation / initialization sequencing + # ------------------------------------------------------------------ # + + @patch("skyflow_flowvault.vault.controller._vault.validate_delete_request") + def test_delete_validates_before_initializing_client(self, mock_validate): + self.delete_api.with_raw_response.delete.return_value = fake_delete_raw_response([]) + request = DeleteRequest(table="t1", ids=["id1"]) + + self.vault.delete(request) + + mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + self.vault_client.initialize_client_configuration.assert_called_once() + + def test_delete_raises_for_invalid_request(self): + with self.assertRaises(SkyflowError): + self.vault.delete(DeleteRequest(table="t1")) + self.vault_client.initialize_client_configuration.assert_not_called() + + def test_delete_raises_on_invalid_table_name(self): + with self.assertRaises(SkyflowError): + self.vault.delete(DeleteRequest(table=" ", ids=["id1"])) + self.delete_api.with_raw_response.delete.assert_not_called() + + # ------------------------------------------------------------------ # + # request -> wire field mapping + # ------------------------------------------------------------------ # + + def test_maps_table_and_ids(self): + self.delete_api.with_raw_response.delete.return_value = fake_delete_raw_response([]) + + self.vault.delete(DeleteRequest(table="t1", ids=["id1", "id2"])) + + _, kwargs = self.delete_api.with_raw_response.delete.call_args + self.assertEqual(kwargs["vault_id"], "vault123") + self.assertEqual(kwargs["table_name"], "t1") + self.assertEqual(kwargs["skyflow_i_ds"], ["id1", "id2"]) + + def test_maps_unique_values(self): + self.delete_api.with_raw_response.delete.return_value = fake_delete_raw_response([]) + + self.vault.delete(DeleteRequest(table="t1", unique_values=[{"email": "a@b.com"}])) + + _, kwargs = self.delete_api.with_raw_response.delete.call_args + self.assertEqual(len(kwargs["unique_values"]), 1) + self.assertEqual(kwargs["unique_values"][0].data, {"email": "a@b.com"}) + + # ------------------------------------------------------------------ # + # response shape -- no tokens/data field at all on V1DeleteResponseObject; + # this is the regression test pinning the getattr(..., 'tokens', None) fix + # ------------------------------------------------------------------ # + + def test_successful_records_have_no_data_or_tokens_keys(self): + self.delete_api.with_raw_response.delete.return_value = fake_delete_raw_response([ + FakeDeleteResponseObject(skyflow_id="id1"), + ], headers={"x-request-id": "req-1"}) + + response = self.vault.delete(DeleteRequest(table="t1", ids=["id1"])) + + self.assertEqual(len(response.records), 1) + record = response.records[0] + self.assertEqual(record["request_index"], 0) + self.assertEqual(record["skyflow_id"], "id1") + self.assertNotIn("data", record) + self.assertNotIn("hashed_data", record) + self.assertNotIn("tokens", record) + self.assertIsNone(response.errors) + + def test_mixed_success_and_error_records_are_split(self): + self.delete_api.with_raw_response.delete.return_value = fake_delete_raw_response([ + FakeDeleteResponseObject(skyflow_id="id1"), + FakeDeleteResponseObject(error="not found", http_code=404), + ], headers={"x-request-id": "req-2"}) + + response = self.vault.delete(DeleteRequest(table="t1", ids=["id1", "id2"])) + + self.assertEqual(len(response.records), 1) + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "not found") + self.assertEqual(response.errors[0]["code"], 404) + + # ------------------------------------------------------------------ # + # transport failure + # ------------------------------------------------------------------ # + + def test_transport_exception_marks_every_id_as_an_error(self): + self.delete_api.with_raw_response.delete.side_effect = Exception("network blip") + + response = self.vault.delete(DeleteRequest(table="t1", ids=["id1", "id2"])) + + self.assertEqual(len(response.records), 0) + self.assertEqual(len(response.errors), 2) + self.assertTrue(all("network blip" in e["error"] for e in response.errors)) + + def test_api_error_with_structured_body_splits_into_one_error_per_row(self): + api_error = ApiError( + status_code=404, + headers={"x-request-id": "req-3"}, + body={"records": [{"error": "not found", "httpCode": 404}]}, + ) + self.delete_api.with_raw_response.delete.side_effect = api_error + + response = self.vault.delete(DeleteRequest(table="t1", ids=["id1"])) + + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "not found") + self.assertEqual(response.errors[0]["code"], 404) + self.assertEqual(response.errors[0]["request_id"], "req-3") + + # ------------------------------------------------------------------ # + # per-call Authorization header injection + # ------------------------------------------------------------------ # + + def test_injects_authorization_header_from_current_bearer_token(self): + self.vault_client.get_current_bearer_token.return_value = "the-current-token" + self.delete_api.with_raw_response.delete.return_value = fake_delete_raw_response([]) + + self.vault.delete(DeleteRequest(table="t1", ids=["id1"])) + + _, kwargs = self.delete_api.with_raw_response.delete.call_args + headers = kwargs["request_options"]["additional_headers"] + self.assertEqual(headers.get("Authorization"), "Bearer the-current-token") + + +def fake_detokenize_raw_response(response, headers=None): + return SimpleNamespace(data=SimpleNamespace(response=response), headers=headers or {}) + + +class TestVaultDetokenize(unittest.TestCase): + def setUp(self): + self.vault_client = Mock() + self.vault_client.get_vault_id.return_value = "vault123" + self.vault_client.get_logger.return_value = Mock() + self.vault_client.get_current_bearer_token.return_value = None + self.detokenize_api = MagicMock() + self.vault_client.get_flowservice_api.return_value = self.detokenize_api + self.vault = VaultController(self.vault_client) + + # ------------------------------------------------------------------ # + # validation / initialization sequencing + # ------------------------------------------------------------------ # + + @patch("skyflow_flowvault.vault.controller._vault.validate_detokenize_request") + def test_detokenize_validates_before_initializing_client(self, mock_validate): + self.detokenize_api.with_raw_response.detokenize.return_value = fake_detokenize_raw_response([]) + request = DetokenizeRequest(tokens=["tok1"]) + + self.vault.detokenize(request) + + mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + self.vault_client.initialize_client_configuration.assert_called_once() + + def test_detokenize_raises_for_invalid_request(self): + with self.assertRaises(SkyflowError): + self.vault.detokenize(DetokenizeRequest(tokens=[])) + self.vault_client.initialize_client_configuration.assert_not_called() + + # ------------------------------------------------------------------ # + # request -> wire field mapping + # ------------------------------------------------------------------ # + + def test_maps_tokens(self): + self.detokenize_api.with_raw_response.detokenize.return_value = fake_detokenize_raw_response([]) + + self.vault.detokenize(DetokenizeRequest(tokens=["tok1", "tok2"])) + + _, kwargs = self.detokenize_api.with_raw_response.detokenize.call_args + self.assertEqual(kwargs["vault_id"], "vault123") + self.assertEqual(kwargs["tokens"], ["tok1", "tok2"]) + + def test_maps_token_group_redactions(self): + self.detokenize_api.with_raw_response.detokenize.return_value = fake_detokenize_raw_response([]) + + self.vault.detokenize(DetokenizeRequest( + tokens=["tok1"], token_group_redactions=[{"token_group_name": "g1", "redaction": "mask1"}], + )) + + _, kwargs = self.detokenize_api.with_raw_response.detokenize.call_args + self.assertEqual(len(kwargs["token_group_redactions"]), 1) + self.assertEqual(kwargs["token_group_redactions"][0].token_group_name, "g1") + self.assertEqual(kwargs["token_group_redactions"][0].redaction, "mask1") + + # ------------------------------------------------------------------ # + # response shape -- keyed by token, not skyflow_id + # ------------------------------------------------------------------ # + + def test_successful_records_include_value_and_token_group_name(self): + self.detokenize_api.with_raw_response.detokenize.return_value = fake_detokenize_raw_response([ + FakeDetokenizeResponseObject(token="tok1", value="john doe", token_group_name="deterministic_string"), + ], headers={"x-request-id": "req-1"}) + + response = self.vault.detokenize(DetokenizeRequest(tokens=["tok1"])) + + self.assertEqual(len(response.records), 1) + record = response.records[0] + self.assertEqual(record["request_index"], 0) + self.assertEqual(record["token"], "tok1") + self.assertEqual(record["value"], "john doe") + self.assertEqual(record["token_group_name"], "deterministic_string") + self.assertIsNone(response.errors) + + def test_mixed_success_and_error_records_are_split(self): + self.detokenize_api.with_raw_response.detokenize.return_value = fake_detokenize_raw_response([ + FakeDetokenizeResponseObject(token="tok1", value="john doe"), + FakeDetokenizeResponseObject(token="tok2", error="invalid token", http_code=404), + ], headers={"x-request-id": "req-2"}) + + response = self.vault.detokenize(DetokenizeRequest(tokens=["tok1", "tok2"])) + + self.assertEqual(len(response.records), 1) + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["token"], "tok2") + self.assertEqual(response.errors[0]["error"], "invalid token") + self.assertEqual(response.errors[0]["code"], 404) + + # ------------------------------------------------------------------ # + # transport failure + # ------------------------------------------------------------------ # + + def test_transport_exception_marks_every_token_as_an_error(self): + self.detokenize_api.with_raw_response.detokenize.side_effect = Exception("network blip") + + response = self.vault.detokenize(DetokenizeRequest(tokens=["tok1", "tok2"])) + + self.assertEqual(len(response.records), 0) + self.assertEqual(len(response.errors), 2) + self.assertTrue(all("network blip" in e["error"] for e in response.errors)) + + def test_api_error_with_structured_body_splits_into_one_error_per_row(self): + api_error = ApiError( + status_code=404, + headers={"x-request-id": "req-3"}, + body={"records": [{"error": "invalid token", "httpCode": 404}]}, + ) + self.detokenize_api.with_raw_response.detokenize.side_effect = api_error + + response = self.vault.detokenize(DetokenizeRequest(tokens=["tok1"])) + + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "invalid token") + self.assertEqual(response.errors[0]["code"], 404) + self.assertEqual(response.errors[0]["request_id"], "req-3") + + # ------------------------------------------------------------------ # + # per-call Authorization header injection + # ------------------------------------------------------------------ # + + def test_injects_authorization_header_from_current_bearer_token(self): + self.vault_client.get_current_bearer_token.return_value = "the-current-token" + self.detokenize_api.with_raw_response.detokenize.return_value = fake_detokenize_raw_response([]) + + self.vault.detokenize(DetokenizeRequest(tokens=["tok1"])) + + _, kwargs = self.detokenize_api.with_raw_response.detokenize.call_args + headers = kwargs["request_options"]["additional_headers"] + self.assertEqual(headers.get("Authorization"), "Bearer the-current-token") + + +def fake_tokenize_raw_response(response, headers=None): + return SimpleNamespace(data=SimpleNamespace(response=response), headers=headers or {}) + + +class TestVaultTokenize(unittest.TestCase): + def setUp(self): + self.vault_client = Mock() + self.vault_client.get_vault_id.return_value = "vault123" + self.vault_client.get_logger.return_value = Mock() + self.vault_client.get_current_bearer_token.return_value = None + self.tokenize_api = MagicMock() + self.vault_client.get_flowservice_api.return_value = self.tokenize_api + self.vault = VaultController(self.vault_client) + + # ------------------------------------------------------------------ # + # validation / initialization sequencing + # ------------------------------------------------------------------ # + + @patch("skyflow_flowvault.vault.controller._vault.validate_tokenize_request") + def test_tokenize_validates_before_initializing_client(self, mock_validate): + self.tokenize_api.with_raw_response.tokenize.return_value = fake_tokenize_raw_response([]) + request = TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1"]}]) + + self.vault.tokenize(request) + + mock_validate.assert_called_once_with(self.vault_client.get_logger(), request) + self.vault_client.initialize_client_configuration.assert_called_once() + + def test_tokenize_raises_for_invalid_request(self): + with self.assertRaises(SkyflowError): + self.vault.tokenize(TokenizeRequest(values=[])) + self.vault_client.initialize_client_configuration.assert_not_called() + + # ------------------------------------------------------------------ # + # request -> wire field mapping + # ------------------------------------------------------------------ # + + def test_maps_values_and_token_group_names(self): + self.tokenize_api.with_raw_response.tokenize.return_value = fake_tokenize_raw_response([]) + + self.vault.tokenize(TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1", "g2"]}])) + + _, kwargs = self.tokenize_api.with_raw_response.tokenize.call_args + self.assertEqual(kwargs["vault_id"], "vault123") + self.assertEqual(len(kwargs["data"]), 1) + self.assertEqual(kwargs["data"][0].value, "a@b.com") + self.assertEqual(kwargs["data"][0].token_group_names, ["g1", "g2"]) + + def test_no_byot_token_is_omitted_not_sent_as_none(self): + self.tokenize_api.with_raw_response.tokenize.return_value = fake_tokenize_raw_response([]) + + self.vault.tokenize(TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1"]}])) + + _, kwargs = self.tokenize_api.with_raw_response.tokenize.call_args + self.assertIsNone(kwargs["data"][0].token) + + def test_maps_byot_token_when_present(self): + self.tokenize_api.with_raw_response.tokenize.return_value = fake_tokenize_raw_response([]) + + self.vault.tokenize(TokenizeRequest( + values=[{"value": "a@b.com", "token_group_names": ["g1"], "token": "custom-tok"}], + )) + + _, kwargs = self.tokenize_api.with_raw_response.tokenize.call_args + self.assertEqual(kwargs["data"][0].token, "custom-tok") + + # ------------------------------------------------------------------ # + # response shape -- one value fans out to a list of per-token-group results + # ------------------------------------------------------------------ # + + def test_successful_value_fans_out_to_one_record_per_token_group(self): + self.tokenize_api.with_raw_response.tokenize.return_value = fake_tokenize_raw_response([ + FakeTokenizeResponseObject(value="a@b.com", tokens=[ + FakeTokenizeResponseObjectToken(token_group_name="g1", token="tok-g1"), + FakeTokenizeResponseObjectToken(token_group_name="g2", token="tok-g2"), + ]), + ], headers={"x-request-id": "req-1"}) + + response = self.vault.tokenize(TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1", "g2"]}])) + + self.assertEqual(len(response.records), 2) + self.assertEqual(response.records[0]["value"], "a@b.com") + self.assertEqual(response.records[0]["token_group_name"], "g1") + self.assertEqual(response.records[0]["token"], "tok-g1") + self.assertEqual(response.records[1]["token_group_name"], "g2") + self.assertEqual(response.records[1]["token"], "tok-g2") + self.assertIsNone(response.errors) + + def test_mixed_success_and_error_token_groups_are_split(self): + self.tokenize_api.with_raw_response.tokenize.return_value = fake_tokenize_raw_response([ + FakeTokenizeResponseObject(value="a@b.com", tokens=[ + FakeTokenizeResponseObjectToken(token_group_name="g1", token="tok-g1"), + FakeTokenizeResponseObjectToken(token_group_name="g2", error="group not found", http_code=404), + ]), + ], headers={"x-request-id": "req-2"}) + + response = self.vault.tokenize(TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1", "g2"]}])) + + self.assertEqual(len(response.records), 1) + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["token_group_name"], "g2") + self.assertEqual(response.errors[0]["error"], "group not found") + self.assertEqual(response.errors[0]["code"], 404) + + # ------------------------------------------------------------------ # + # transport failure + # ------------------------------------------------------------------ # + + def test_transport_exception_marks_every_value_as_an_error(self): + self.tokenize_api.with_raw_response.tokenize.side_effect = Exception("network blip") + values = [ + {"value": "a@b.com", "token_group_names": ["g1"]}, + {"value": "b@c.com", "token_group_names": ["g1"]}, + ] + + response = self.vault.tokenize(TokenizeRequest(values=values)) + + self.assertEqual(len(response.records), 0) + self.assertEqual(len(response.errors), 2) + self.assertTrue(all("network blip" in e["error"] for e in response.errors)) + + def test_api_error_with_structured_body_splits_into_one_error_per_row(self): + api_error = ApiError( + status_code=404, + headers={"x-request-id": "req-3"}, + body={"records": [{"error": "group not found", "httpCode": 404}]}, + ) + self.tokenize_api.with_raw_response.tokenize.side_effect = api_error + + response = self.vault.tokenize(TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1"]}])) + + self.assertEqual(len(response.errors), 1) + self.assertEqual(response.errors[0]["error"], "group not found") + self.assertEqual(response.errors[0]["code"], 404) + self.assertEqual(response.errors[0]["request_id"], "req-3") + + # ------------------------------------------------------------------ # + # per-call Authorization header injection + # ------------------------------------------------------------------ # + + def test_injects_authorization_header_from_current_bearer_token(self): + self.vault_client.get_current_bearer_token.return_value = "the-current-token" + self.tokenize_api.with_raw_response.tokenize.return_value = fake_tokenize_raw_response([]) + + self.vault.tokenize(TokenizeRequest(values=[{"value": "a@b.com", "token_group_names": ["g1"]}])) + + _, kwargs = self.tokenize_api.with_raw_response.tokenize.call_args + headers = kwargs["request_options"]["additional_headers"] + self.assertEqual(headers.get("Authorization"), "Bearer the-current-token") + + if __name__ == "__main__": unittest.main() diff --git a/flowvault/tests/vault/data/test_data_classes.py b/flowvault/tests/vault/data/test_data_classes.py index 33b274c..c2ab2b1 100644 --- a/flowvault/tests/vault/data/test_data_classes.py +++ b/flowvault/tests/vault/data/test_data_classes.py @@ -2,7 +2,20 @@ from common.vault.data import BaseInsertRequest, BaseInsertResponse from skyflow_flowvault.utils.enums import UpsertType -from skyflow_flowvault.vault.data import InsertRequest, InsertResponse +from skyflow_flowvault.vault.data import ( + InsertRequest, + InsertResponse, + GetRequest, + GetResponse, + UpdateRequest, + UpdateResponse, + DeleteRequest, + DeleteResponse, + DetokenizeRequest, + DetokenizeResponse, + TokenizeRequest, + TokenizeResponse, +) class TestInsertRequest(unittest.TestCase): @@ -53,5 +66,167 @@ def test_repr_does_not_raise(self): self.assertIn("InsertResponse", repr(response)) +class TestGetRequest(unittest.TestCase): + def test_required_and_optional_defaults(self): + request = GetRequest(table="t1", ids=["id1"]) + self.assertEqual(request.table, "t1") + self.assertEqual(request.ids, ["id1"]) + self.assertIsNone(request.unique_values) + self.assertIsNone(request.columns) + self.assertIsNone(request.column_redactions) + self.assertIsNone(request.limit) + self.assertIsNone(request.offset) + + def test_all_fields_stored(self): + request = GetRequest( + table="t1", ids=["id1"], unique_values=[{"email": "a@b.com"}], columns=["a", "b"], + column_redactions=[{"column_name": "a", "redaction": "mask1"}], limit=10, offset=5, + ) + self.assertEqual(request.unique_values, [{"email": "a@b.com"}]) + self.assertEqual(request.columns, ["a", "b"]) + self.assertEqual(request.column_redactions, [{"column_name": "a", "redaction": "mask1"}]) + self.assertEqual(request.limit, 10) + self.assertEqual(request.offset, 5) + + +class TestGetResponse(unittest.TestCase): + def test_shape(self): + records = [{"request_index": 0, "skyflow_id": "id1", "data": {"a": 1}}] + response = GetResponse(records=records, errors=[]) + self.assertIs(response.records, records) + self.assertEqual(response.errors, []) + + def test_defaults(self): + response = GetResponse() + self.assertIsNone(response.records) + self.assertIsNone(response.errors) + + def test_repr_and_str_do_not_raise(self): + response = GetResponse(records=[], errors=[{"request_index": 0, "error": "boom"}]) + self.assertIn("GetResponse", repr(response)) + self.assertIn("GetResponse", str(response)) + + +class TestUpdateRequest(unittest.TestCase): + def test_required_and_optional_defaults(self): + request = UpdateRequest(records=[{"skyflow_id": "id1", "values": {"a": 1}}]) + self.assertEqual(request.records, [{"skyflow_id": "id1", "values": {"a": 1}}]) + self.assertIsNone(request.table) + self.assertIsNone(request.update_type) + + def test_all_fields_stored(self): + request = UpdateRequest( + records=[{"skyflow_id": "id1", "values": {"a": 1}, "tokens": {"a": "tok"}, "table": "t2"}], + table="t1", update_type=UpsertType.REPLACE, + ) + self.assertEqual(request.table, "t1") + self.assertEqual(request.update_type, UpsertType.REPLACE) + self.assertEqual(request.records[0]["tokens"], {"a": "tok"}) + + +class TestUpdateResponse(unittest.TestCase): + def test_shape(self): + records = [{"request_index": 0, "skyflow_id": "id1"}] + response = UpdateResponse(records=records, errors=[]) + self.assertIs(response.records, records) + self.assertEqual(response.errors, []) + + def test_defaults(self): + response = UpdateResponse() + self.assertIsNone(response.records) + self.assertIsNone(response.errors) + + def test_repr_and_str_do_not_raise(self): + response = UpdateResponse(records=[], errors=[{"request_index": 0, "error": "boom"}]) + self.assertIn("UpdateResponse", repr(response)) + self.assertIn("UpdateResponse", str(response)) + + +class TestDeleteRequest(unittest.TestCase): + def test_required_and_optional_defaults(self): + request = DeleteRequest(table="t1", ids=["id1"]) + self.assertEqual(request.table, "t1") + self.assertEqual(request.ids, ["id1"]) + self.assertIsNone(request.unique_values) + + def test_unique_values_stored(self): + request = DeleteRequest(table="t1", unique_values=[{"email": "a@b.com"}]) + self.assertEqual(request.unique_values, [{"email": "a@b.com"}]) + + +class TestDeleteResponse(unittest.TestCase): + def test_shape(self): + records = [{"request_index": 0, "skyflow_id": "id1"}] + response = DeleteResponse(records=records, errors=[]) + self.assertIs(response.records, records) + self.assertEqual(response.errors, []) + + def test_defaults(self): + response = DeleteResponse() + self.assertIsNone(response.records) + self.assertIsNone(response.errors) + + def test_repr_and_str_do_not_raise(self): + response = DeleteResponse(records=[], errors=[{"request_index": 0, "error": "boom"}]) + self.assertIn("DeleteResponse", repr(response)) + self.assertIn("DeleteResponse", str(response)) + + +class TestDetokenizeRequest(unittest.TestCase): + def test_required_and_optional_defaults(self): + request = DetokenizeRequest(tokens=["tok1", "tok2"]) + self.assertEqual(request.tokens, ["tok1", "tok2"]) + self.assertIsNone(request.token_group_redactions) + + def test_token_group_redactions_stored(self): + request = DetokenizeRequest( + tokens=["tok1"], token_group_redactions=[{"token_group_name": "g1", "redaction": "mask1"}], + ) + self.assertEqual(request.token_group_redactions, [{"token_group_name": "g1", "redaction": "mask1"}]) + + +class TestDetokenizeResponse(unittest.TestCase): + def test_shape(self): + records = [{"request_index": 0, "token": "tok1", "value": "john"}] + response = DetokenizeResponse(records=records, errors=[]) + self.assertIs(response.records, records) + self.assertEqual(response.errors, []) + + def test_defaults(self): + response = DetokenizeResponse() + self.assertIsNone(response.records) + self.assertIsNone(response.errors) + + def test_repr_and_str_do_not_raise(self): + response = DetokenizeResponse(records=[], errors=[{"request_index": 0, "error": "boom"}]) + self.assertIn("DetokenizeResponse", repr(response)) + self.assertIn("DetokenizeResponse", str(response)) + + +class TestTokenizeRequest(unittest.TestCase): + def test_values_stored(self): + values = [{"value": "a@b.com", "token_group_names": ["g1"]}] + request = TokenizeRequest(values=values) + self.assertEqual(request.values, values) + + +class TestTokenizeResponse(unittest.TestCase): + def test_shape(self): + records = [{"request_index": 0, "token": "tok1", "token_group_name": "g1"}] + response = TokenizeResponse(records=records, errors=[]) + self.assertIs(response.records, records) + self.assertEqual(response.errors, []) + + def test_defaults(self): + response = TokenizeResponse() + self.assertIsNone(response.records) + self.assertIsNone(response.errors) + + def test_repr_and_str_do_not_raise(self): + response = TokenizeResponse(records=[], errors=[{"request_index": 0, "error": "boom"}]) + self.assertIn("TokenizeResponse", repr(response)) + self.assertIn("TokenizeResponse", str(response)) + + if __name__ == "__main__": unittest.main() diff --git a/tests/contract/adapters/v3_adapter.py b/tests/contract/adapters/v3_adapter.py index d5a7ceb..396ad5c 100644 --- a/tests/contract/adapters/v3_adapter.py +++ b/tests/contract/adapters/v3_adapter.py @@ -18,7 +18,7 @@ def build_vault(): vault_client = VaultClient(config) vault_client.initialize_client_configuration = MagicMock() # skip real credential/URL resolution insert_api = MagicMock() - vault_client.get_insert_api = MagicMock(return_value=insert_api) + vault_client.get_flowservice_api = MagicMock(return_value=insert_api) vault = VaultController(vault_client) return vault, insert_api