diff --git a/src/azure-cli/azure/cli/command_modules/sql/_help.py b/src/azure-cli/azure/cli/command_modules/sql/_help.py index be720297a4a..69604702916 100644 --- a/src/azure-cli/azure/cli/command_modules/sql/_help.py +++ b/src/azure-cli/azure/cli/command_modules/sql/_help.py @@ -73,6 +73,11 @@ text: | az sql db audit-policy update -g mygroup -s myserver -n mydb --state Enabled \\ --lats Enabled --lawri myworkspaceresourceid + - name: Set the fields included in audit events sent to Azure Monitor. + text: | + az sql db audit-policy update -g mygroup -s myserver -n mydb --state Enabled \\ + --lats Enabled --lawri myworkspaceresourceid \\ + --required-fields event_time action_id statement - name: Disable a log analytics auditing policy. text: | az sql db audit-policy update -g mygroup -s myserver -n mydb @@ -1490,6 +1495,11 @@ text: | az sql server audit-policy update -g mygroup -n myserver --state Enabled \\ --lats Enabled --lawri myworkspaceresourceid + - name: Set the fields included in audit events sent to Azure Monitor. + text: | + az sql server audit-policy update -g mygroup -n myserver --state Enabled \\ + --lats Enabled --lawri myworkspaceresourceid \\ + --required-fields event_time action_id statement - name: Disable a log analytics auditing policy. text: | az sql server audit-policy update -g mygroup -n myserver diff --git a/src/azure-cli/azure/cli/command_modules/sql/_params.py b/src/azure-cli/azure/cli/command_modules/sql/_params.py index cc2445cd5a8..ffda2371122 100644 --- a/src/azure-cli/azure/cli/command_modules/sql/_params.py +++ b/src/azure-cli/azure/cli/command_modules/sql/_params.py @@ -1203,6 +1203,12 @@ def _configure_security_policy_storage_params(arg_ctx): 'Example: --actions FAILED_DATABASE_AUTHENTICATION_GROUP BATCH_COMPLETED_GROUP', nargs='+') + c.argument('required_fields', + arg_group=policy_arg_group, + help='List of fields to include in audit events. Can only be specified when the Azure Monitor ' + 'target is enabled.', + nargs='+') + c.argument('retention_days', arg_group=policy_arg_group, help='The number of days to retain audit logs.') @@ -2056,6 +2062,12 @@ def _configure_security_policy_storage_params(arg_ctx): 'Example: --actions FAILED_DATABASE_AUTHENTICATION_GROUP BATCH_COMPLETED_GROUP', nargs='+') + c.argument('required_fields', + arg_group=policy_arg_group, + help='List of fields to include in audit events. Can only be specified when the Azure Monitor ' + 'target is enabled.', + nargs='+') + c.argument('retention_days', arg_group=policy_arg_group, help='The number of days to retain audit logs.') diff --git a/src/azure-cli/azure/cli/command_modules/sql/commands.py b/src/azure-cli/azure/cli/command_modules/sql/commands.py index c92f56b1e8a..492121ab95d 100644 --- a/src/azure-cli/azure/cli/command_modules/sql/commands.py +++ b/src/azure-cli/azure/cli/command_modules/sql/commands.py @@ -277,12 +277,18 @@ def load_command_table(self, _): operations_tmpl='azure.mgmt.sql.operations#DatabaseBlobAuditingPoliciesOperations.{}', client_factory=get_sql_database_blob_auditing_policies_operations) + audit_policy_custom = CliCommandType( + operations_tmpl='azure.cli.command_modules.sql.custom#{}') + with self.command_group('sql db audit-policy', database_blob_auditing_policies_operations, client_factory=get_sql_database_blob_auditing_policies_operations) as g: g.custom_show_command('show', 'db_audit_policy_show') - g.generic_update_command('update', custom_func_name='db_audit_policy_update') + g.generic_update_command('update', + setter_name='db_audit_policy_set', + setter_type=audit_policy_custom, + custom_func_name='db_audit_policy_update') g.wait_command('wait') server_blob_auditing_policies_operations = CliCommandType( @@ -295,7 +301,8 @@ def load_command_table(self, _): g.custom_show_command('show', 'server_audit_policy_show') g.generic_update_command('update', - setter_name='begin_create_or_update', + setter_name='server_audit_policy_set', + setter_type=audit_policy_custom, custom_func_name='server_audit_policy_update', supports_no_wait=True) g.wait_command('wait') diff --git a/src/azure-cli/azure/cli/command_modules/sql/custom.py b/src/azure-cli/azure/cli/command_modules/sql/custom.py index 0bcdd69f980..9329616ddaa 100644 --- a/src/azure-cli/azure/cli/command_modules/sql/custom.py +++ b/src/azure-cli/azure/cli/command_modules/sql/custom.py @@ -2112,6 +2112,8 @@ def _get_database_keys_for_update(akvKeys, akvKeysToRemove): # sql db audit-policy & threat-policy ##### +_AUDITING_API_VERSION = '2026-08-01-preview' + def _find_storage_account_resource_group(cli_ctx, name): ''' @@ -2300,6 +2302,100 @@ def _get_diagnostic_settings( return list(azure_monitor_client.diagnostic_settings.list(diagnostic_settings_url)) +def _attach_required_fields(pipeline_response, deserialized, _headers): + ''' + The installed azure-mgmt-sql SDK's generated models do not define 'required_fields' (a + 2026-08-01-preview-only property), so the SDK's own deserializer silently drops it from the + response. Read it from the raw response body via the SDK's public 'cls' extensibility hook + (supported by every generated operation method) and attach it to the model instance the SDK + already built, instead of hand-parsing the response ourselves. + ''' + raw_body = pipeline_response.http_response.json() + deserialized.required_fields = raw_body.get('properties', {}).get('requiredFields') + return deserialized + + +def _get_audit_policy_preview(cmd, client, resource_group_name, server_name, database_name=None): # pylint: disable=unused-argument + ''' + Calls the real SDK client, overriding only the api_version (a publicly documented kwarg on + every generated operation method) so the 2026-08-01-preview surface -- the first to define + 'requiredFields' -- is used instead of the SDK's own pinned default. + ''' + kwargs = { + 'resource_group_name': resource_group_name, + 'server_name': server_name, + 'api_version': _AUDITING_API_VERSION, + 'cls': _attach_required_fields, + } + if database_name is not None: + kwargs['database_name'] = database_name + return client.get(**kwargs) + + +def _set_audit_policy_preview( + cmd, # pylint: disable=unused-argument + client, + parameters, + resource_group_name, + server_name, + database_name=None, + no_wait=False): + ''' + The installed SDK's models don't define 'required_fields', so parameters.serialize() won't + include it. Serialize what the SDK knows about, inject the extra field, then hand the SDK + client raw JSON bytes instead of the typed model: both 'create_or_update' and + 'begin_create_or_update' publicly accept 'parameters: Union[, IO[bytes]]'. This keeps + the request on the SDK's own client (auth, retries, and -- for the server-level LRO -- its + own ARMPolling-based poller), rather than issuing a raw HTTP call ourselves. + ''' + import json + + body = parameters.serialize() + body.setdefault('properties', {})['requiredFields'] = parameters.required_fields + raw_parameters = json.dumps(body).encode('utf-8') + + if database_name is not None: + return client.create_or_update( + resource_group_name=resource_group_name, + server_name=server_name, + database_name=database_name, + parameters=raw_parameters, + api_version=_AUDITING_API_VERSION, + cls=_attach_required_fields) + + return sdk_no_wait( + no_wait, + client.begin_create_or_update, + resource_group_name=resource_group_name, + server_name=server_name, + parameters=raw_parameters, + api_version=_AUDITING_API_VERSION, + cls=_attach_required_fields) + + +def db_audit_policy_set(cmd, client, resource_group_name, server_name, database_name, parameters): + if hasattr(parameters, 'required_fields'): + return _set_audit_policy_preview( + cmd, client, parameters, resource_group_name, server_name, database_name) + return client.create_or_update( + resource_group_name=resource_group_name, + server_name=server_name, + database_name=database_name, + parameters=parameters) + + +def server_audit_policy_set(cmd, client, resource_group_name, server_name, parameters, no_wait=False): + if hasattr(parameters, 'required_fields'): + return _set_audit_policy_preview( + cmd, client, parameters, resource_group_name, server_name, no_wait=no_wait) + return sdk_no_wait( + no_wait, + client.begin_create_or_update, + resource_group_name=resource_group_name, + server_name=server_name, + parameters=parameters) + + def _fetch_first_audit_diagnostic_setting(diagnostic_settings, category_name): return next((ds for ds in diagnostic_settings if hasattr(ds, 'logs') and next((log for log in ds.logs if log.enabled and @@ -2361,11 +2457,15 @@ def _audit_policy_show( resource_group_name=resource_group_name, server_name=server_name) else: - audit_policy = client.get( + audit_policy = _get_audit_policy_preview( + cmd=cmd, + client=client, resource_group_name=resource_group_name, server_name=server_name) else: - audit_policy = client.get( + audit_policy = _get_audit_policy_preview( + cmd=cmd, + client=client, resource_group_name=resource_group_name, server_name=server_name, database_name=database_name) @@ -2470,6 +2570,7 @@ def _audit_policy_validate_arguments( storage_endpoint=None, storage_account_access_key=None, retention_days=None, + required_fields=None, log_analytics_target_state=None, log_analytics_workspace_resource_id=None, event_hub_target_state=None, @@ -2492,7 +2593,7 @@ def _audit_policy_validate_arguments( event_hub_name is not None if not state and not blob_storage_arguments_provided and\ - not log_analytics_arguments_provided and not event_hub_arguments_provided: + not log_analytics_arguments_provided and not event_hub_arguments_provided and required_fields is None: raise CLIError('Either state or blob storage or log analytics or event hub arguments are missing') if _is_audit_policy_state_enabled(state) and\ @@ -2503,7 +2604,8 @@ def _audit_policy_validate_arguments( if _is_audit_policy_state_disabled(state) and\ (blob_storage_arguments_provided or log_analytics_arguments_provided or - event_hub_name): + event_hub_name or + required_fields is not None): raise CLIError('No additional arguments should be provided once state is disabled') if (_is_audit_policy_state_none_or_disabled(blob_storage_target_state)) and\ @@ -2883,6 +2985,7 @@ def _audit_policy_update_global_settings( storage_endpoint=None, storage_account_access_key=None, audit_actions_and_groups=None, + required_fields=None, retention_days=None, log_analytics_target_state=None, event_hub_target_state=None): @@ -2925,6 +3028,12 @@ def _audit_policy_update_global_settings( log_analytics_target_state=log_analytics_target_state, event_hub_target_state=event_hub_target_state) + if required_fields is not None: + if not _is_audit_policy_state_enabled(instance.state) or not instance.is_azure_monitor_target_enabled: + raise ValidationError( + 'required-fields can only be specified when auditing and the Azure Monitor target are enabled') + instance.required_fields = required_fields + def _audit_policy_update_rollback( cmd, @@ -2974,6 +3083,7 @@ def _audit_policy_update( storage_endpoint=None, storage_account_access_key=None, audit_actions_and_groups=None, + required_fields=None, retention_days=None, category_name=None, log_analytics_target_state=None, @@ -2990,6 +3100,7 @@ def _audit_policy_update( storage_endpoint=storage_endpoint, storage_account_access_key=storage_account_access_key, retention_days=retention_days, + required_fields=required_fields, log_analytics_target_state=log_analytics_target_state, log_analytics_workspace_resource_id=log_analytics_workspace_resource_id, event_hub_target_state=event_hub_target_state, @@ -3034,6 +3145,7 @@ def _audit_policy_update( storage_endpoint=storage_endpoint, storage_account_access_key=storage_account_access_key, audit_actions_and_groups=audit_actions_and_groups, + required_fields=required_fields, retention_days=retention_days, log_analytics_target_state=log_analytics_target_state, event_hub_target_state=event_hub_target_state) @@ -3065,6 +3177,7 @@ def server_audit_policy_update( storage_endpoint=None, storage_account_access_key=None, audit_actions_and_groups=None, + required_fields=None, retention_days=None, log_analytics_target_state=None, log_analytics_workspace_resource_id=None, @@ -3087,6 +3200,7 @@ def server_audit_policy_update( storage_endpoint=storage_endpoint, storage_account_access_key=storage_account_access_key, audit_actions_and_groups=audit_actions_and_groups, + required_fields=required_fields, retention_days=retention_days, category_name='SQLSecurityAuditEvents', log_analytics_target_state=log_analytics_target_state, @@ -3108,6 +3222,7 @@ def db_audit_policy_update( storage_endpoint=None, storage_account_access_key=None, audit_actions_and_groups=None, + required_fields=None, retention_days=None, log_analytics_target_state=None, log_analytics_workspace_resource_id=None, @@ -3130,6 +3245,7 @@ def db_audit_policy_update( storage_endpoint=storage_endpoint, storage_account_access_key=storage_account_access_key, audit_actions_and_groups=audit_actions_and_groups, + required_fields=required_fields, retention_days=retention_days, category_name='SQLSecurityAuditEvents', log_analytics_target_state=log_analytics_target_state, diff --git a/src/azure-cli/azure/cli/command_modules/sql/tests/latest/test_sql_audit_policy.py b/src/azure-cli/azure/cli/command_modules/sql/tests/latest/test_sql_audit_policy.py new file mode 100644 index 00000000000..45a5149391b --- /dev/null +++ b/src/azure-cli/azure/cli/command_modules/sql/tests/latest/test_sql_audit_policy.py @@ -0,0 +1,241 @@ +# -------------------------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. See License.txt in the project root for license information. +# -------------------------------------------------------------------------------------------- + +import json +import unittest +from types import SimpleNamespace +from unittest.mock import ANY, MagicMock, patch + +from azure.cli.core.azclierror import ValidationError +from azure.cli.command_modules.sql import custom + + +class SqlAuditPolicyTest(unittest.TestCase): + + def test_attach_required_fields_reads_raw_response_body(self): + ''' + The installed SDK's models don't define 'required_fields', so its generic deserializer + silently drops it. _attach_required_fields is the 'cls' callback that recovers it from + the raw HTTP response and attaches it to the SDK's own deserialized model instance. + ''' + deserialized = SimpleNamespace() + pipeline_response = MagicMock() + pipeline_response.http_response.json.return_value = { + 'properties': {'requiredFields': ['event_time', 'statement']} + } + + result = custom._attach_required_fields(pipeline_response, deserialized, {}) + + self.assertIs(deserialized, result) + self.assertEqual(['event_time', 'statement'], result.required_fields) + + def test_attach_required_fields_handles_missing_required_fields(self): + deserialized = SimpleNamespace() + pipeline_response = MagicMock() + pipeline_response.http_response.json.return_value = {'properties': {'state': 'Enabled'}} + + result = custom._attach_required_fields(pipeline_response, deserialized, {}) + + self.assertIsNone(result.required_fields) + + def test_get_audit_policy_preview_calls_sdk_client_for_database(self): + client = MagicMock() + client.get.return_value = SimpleNamespace(required_fields=['event_time']) + + result = custom._get_audit_policy_preview( + MagicMock(), client, 'resource-group', 'server', 'database') + + client.get.assert_called_once_with( + resource_group_name='resource-group', + server_name='server', + database_name='database', + api_version=custom._AUDITING_API_VERSION, + cls=custom._attach_required_fields) + self.assertIs(client.get.return_value, result) + + def test_get_audit_policy_preview_calls_sdk_client_for_server(self): + client = MagicMock() + client.get.return_value = SimpleNamespace(required_fields=None) + + custom._get_audit_policy_preview(MagicMock(), client, 'resource-group', 'server') + + call_kwargs = client.get.call_args.kwargs + self.assertNotIn('database_name', call_kwargs) + self.assertEqual(custom._AUDITING_API_VERSION, call_kwargs['api_version']) + self.assertIs(custom._attach_required_fields, call_kwargs['cls']) + + @patch('azure.cli.command_modules.sql.custom._get_audit_policy_preview') + def test_show_commands_use_preview_audit_policy(self, get_audit_policy_preview): + required_fields = ['event_time', 'action_id'] + get_audit_policy_preview.return_value = SimpleNamespace( + state='Disabled', + required_fields=required_fields) + + server_result = custom.server_audit_policy_show( + MagicMock(), MagicMock(), 'server', 'resource-group') + get_audit_policy_preview.assert_called_once_with( + cmd=ANY, + client=ANY, + resource_group_name='resource-group', + server_name='server') + self.assertEqual(required_fields, server_result.required_fields) + + get_audit_policy_preview.reset_mock() + database_result = custom.db_audit_policy_show( + MagicMock(), MagicMock(), 'server', 'resource-group', 'database') + get_audit_policy_preview.assert_called_once_with( + cmd=ANY, + client=ANY, + resource_group_name='resource-group', + server_name='server', + database_name='database') + self.assertEqual(required_fields, database_result.required_fields) + + def test_set_audit_policy_preview_database_calls_create_or_update(self): + ''' + Database-level auditing settings PUT is synchronous, so this goes through the SDK's + plain 'create_or_update', with 'parameters' as raw JSON bytes (a publicly documented + '@overload' of that method) instead of the typed model, so 'requiredFields' -- which the + typed model doesn't know about -- still reaches the wire. + ''' + client = MagicMock() + parameters = SimpleNamespace( + required_fields=['event_time', 'action_id', 'statement'], + serialize=MagicMock(return_value={'properties': {'state': 'Enabled'}})) + + custom._set_audit_policy_preview( + MagicMock(), client, parameters, 'resource-group', 'server', 'database') + + client.create_or_update.assert_called_once() + call_kwargs = client.create_or_update.call_args.kwargs + self.assertEqual('resource-group', call_kwargs['resource_group_name']) + self.assertEqual('server', call_kwargs['server_name']) + self.assertEqual('database', call_kwargs['database_name']) + self.assertEqual(custom._AUDITING_API_VERSION, call_kwargs['api_version']) + self.assertIs(custom._attach_required_fields, call_kwargs['cls']) + sent_body = json.loads(call_kwargs['parameters']) + self.assertEqual( + ['event_time', 'action_id', 'statement'], + sent_body['properties']['requiredFields']) + + def test_set_audit_policy_preview_server_calls_begin_create_or_update(self): + ''' + Server-level auditing settings PUT can be a long-running operation, so this must go + through 'begin_create_or_update' -- the SDK's own ARMPolling-based poller -- rather than + any hand-rolled polling loop. + ''' + client = MagicMock() + parameters = SimpleNamespace( + required_fields=['event_time'], + serialize=MagicMock(return_value={'properties': {'state': 'Enabled'}})) + + custom._set_audit_policy_preview( + MagicMock(), client, parameters, 'resource-group', 'server') + + client.begin_create_or_update.assert_called_once() + call_kwargs = client.begin_create_or_update.call_args.kwargs + self.assertEqual(custom._AUDITING_API_VERSION, call_kwargs['api_version']) + self.assertIs(custom._attach_required_fields, call_kwargs['cls']) + self.assertNotIn('polling', call_kwargs) + sent_body = json.loads(call_kwargs['parameters']) + self.assertEqual(['event_time'], sent_body['properties']['requiredFields']) + + def test_set_audit_policy_preview_server_no_wait_disables_polling(self): + client = MagicMock() + parameters = SimpleNamespace( + required_fields=['event_time'], + serialize=MagicMock(return_value={'properties': {'state': 'Enabled'}})) + + custom._set_audit_policy_preview( + MagicMock(), client, parameters, 'resource-group', 'server', no_wait=True) + + call_kwargs = client.begin_create_or_update.call_args.kwargs + self.assertEqual(False, call_kwargs['polling']) + + @patch('azure.cli.command_modules.sql.custom._set_audit_policy_preview') + def test_db_audit_policy_set_routes_to_preview_when_required_fields_present(self, set_preview): + parameters = SimpleNamespace(required_fields=['event_time']) + client = MagicMock() + cmd = MagicMock() + + custom.db_audit_policy_set(cmd, client, 'rg', 'server', 'database', parameters) + + set_preview.assert_called_once_with(cmd, client, parameters, 'rg', 'server', 'database') + client.create_or_update.assert_not_called() + + @patch('azure.cli.command_modules.sql.custom._set_audit_policy_preview') + def test_db_audit_policy_set_falls_back_to_sdk_without_required_fields(self, set_preview): + parameters = SimpleNamespace() # no 'required_fields' attribute + client = MagicMock() + + custom.db_audit_policy_set(MagicMock(), client, 'rg', 'server', 'database', parameters) + + set_preview.assert_not_called() + client.create_or_update.assert_called_once_with( + resource_group_name='rg', server_name='server', database_name='database', + parameters=parameters) + + @patch('azure.cli.command_modules.sql.custom._set_audit_policy_preview') + def test_server_audit_policy_set_routes_to_preview_when_required_fields_present(self, set_preview): + parameters = SimpleNamespace(required_fields=['event_time']) + client = MagicMock() + cmd = MagicMock() + + custom.server_audit_policy_set(cmd, client, 'rg', 'server', parameters, no_wait=True) + + set_preview.assert_called_once_with(cmd, client, parameters, 'rg', 'server', no_wait=True) + client.begin_create_or_update.assert_not_called() + + @patch('azure.cli.command_modules.sql.custom._set_audit_policy_preview') + def test_server_audit_policy_set_falls_back_to_sdk_without_required_fields(self, set_preview): + parameters = SimpleNamespace() # no 'required_fields' attribute + client = MagicMock() + + custom.server_audit_policy_set(MagicMock(), client, 'rg', 'server', parameters) + + set_preview.assert_not_called() + client.begin_create_or_update.assert_called_once() + + @patch('azure.cli.command_modules.sql.custom._audit_policy_update_apply_azure_monitor_target_enabled') + @patch('azure.cli.command_modules.sql.custom._audit_policy_update_apply_blob_storage_details') + def test_required_fields_are_applied_when_azure_monitor_is_enabled( + self, _, apply_azure_monitor_target_enabled): + instance = SimpleNamespace( + state='Enabled', + storage_endpoint=None, + storage_account_access_key=None, + is_azure_monitor_target_enabled=True, + audit_actions_and_groups=['BATCH_COMPLETED_GROUP']) + + custom._audit_policy_update_global_settings( + cmd=MagicMock(), + instance=instance, + required_fields=['event_time', 'statement']) + + self.assertEqual(['event_time', 'statement'], instance.required_fields) + apply_azure_monitor_target_enabled.assert_called_once() + + @patch('azure.cli.command_modules.sql.custom._audit_policy_update_apply_azure_monitor_target_enabled') + @patch('azure.cli.command_modules.sql.custom._audit_policy_update_apply_blob_storage_details') + def test_required_fields_require_azure_monitor( + self, _, apply_azure_monitor_target_enabled): + instance = SimpleNamespace( + state='Enabled', + storage_endpoint=None, + storage_account_access_key=None, + is_azure_monitor_target_enabled=False, + audit_actions_and_groups=['BATCH_COMPLETED_GROUP']) + + with self.assertRaisesRegex(ValidationError, 'Azure Monitor target'): + custom._audit_policy_update_global_settings( + cmd=MagicMock(), + instance=instance, + required_fields=['event_time']) + + apply_azure_monitor_target_enabled.assert_called_once() + + +if __name__ == '__main__': + unittest.main()