diff --git a/secator/cli.py b/secator/cli.py index ad3093dee..267604670 100644 --- a/secator/cli.py +++ b/secator/cli.py @@ -1127,11 +1127,6 @@ def list_aliases(silent): @click.pass_context def query(ctx, arg, output, output_folder, time_delta, fmt, workspace, report_filter, driver, dedupe, limit, save): """Query""" - # Empty query: return all results (subject to the enforced base query), - # optionally scoped by --report-filter / --workspace. - if not arg: - run_report_show(report_filter, output, time_delta, None, fmt, workspace, driver, dedupe, limit, output_folder) - return # 0. Save the expression under a name, then exit (reuse later with `secator q `). if save: @@ -1142,6 +1137,12 @@ def query(ctx, arg, output, output_folder, time_delta, fmt, workspace, report_fi else: console.print(Error(message='Invalid config, not saving it.')) return + + # Empty query: return all results (subject to the enforced base query), + # optionally scoped by --report-filter / --workspace. + if not arg: + run_report_show(report_filter, output, time_delta, None, fmt, workspace, driver, dedupe, limit, output_folder) + return # 1. Saved query name if arg in CONFIG.queries: diff --git a/tests/unit/test_query_cli.py b/tests/unit/test_query_cli.py index 22133de27..6ea74d7d5 100644 --- a/tests/unit/test_query_cli.py +++ b/tests/unit/test_query_cli.py @@ -127,35 +127,26 @@ def test_raw_expression(self): vulns = captured.get('results', {}).get('vulnerability', []) self.assertEqual(sorted(v['name'] for v in vulns), ['XSS']) - def test_save_query_round_trips(self): - # `secator q "" --save ` persists the expression under , - # and it becomes loadable/runnable as a named query. - from secator.cli import cli - from secator.config import CONFIG, Config - expr = "vulnerability.severity == 'critical'" - name = 'saved_crit' - CONFIG.queries.pop(name, None) - try: - # Patch Config.save (class level) so persistence doesn't clobber the real config file. - with mock.patch.object(Config, 'save', return_value=True) as mock_save: - result = self.cli_runner.invoke(cli, ['query', expr, '--save', name]) - self.assertIsNone(result.exception, str(result.exception)) - self.assertEqual(result.exit_code, 0) - mock_save.assert_called_once() - # Persisted under the name and resolvable by the named-query lookup path. - self.assertEqual(CONFIG.queries.get(name), expr) - result, captured = self._invoke(['query', name, '-ws', WS, '--driver', 'local']) - self.assertIsNone(result.exception, str(result.exception)) - vulns = captured.get('results', {}).get('vulnerability', []) - self.assertEqual(sorted(v['name'] for v in vulns), ['SQLi']) - finally: - CONFIG.queries.pop(name, None) - - def test_empty_query_returns_all(self): - # `secator q -rf tasks/1` (no query expression) must return all findings - # scoped by the report filter, not raise "Missing argument ARG". + def test_empty_arg_with_report_filter(self): + """secator q -rf tasks/1 should return all results without requiring an ARG.""" result, captured = self._invoke(['query', '-rf', 'tasks/1', '-ws', WS, '--driver', 'local']) self.assertIsNone(result.exception, str(result.exception)) self.assertEqual(result.exit_code, 0) vulns = captured.get('results', {}).get('vulnerability', []) self.assertEqual(sorted(v['name'] for v in vulns), ['SQLi', 'XSS']) + + def test_empty_arg_with_workspace(self): + """secator q -ws myws should return all results for the workspace without requiring an ARG.""" + result, captured = self._invoke(['query', '-ws', WS, '--driver', 'local']) + self.assertIsNone(result.exception, str(result.exception)) + self.assertEqual(result.exit_code, 0) + vulns = captured.get('results', {}).get('vulnerability', []) + self.assertEqual(sorted(v['name'] for v in vulns), ['SQLi', 'XSS']) + + def test_empty_arg_no_filter_shows_all(self): + """bare secator q (no args, no filters) should succeed and return all results.""" + result, captured = self._invoke(['query', '--driver', 'local']) + self.assertIsNone(result.exception, str(result.exception)) + self.assertEqual(result.exit_code, 0) + vulns = captured.get('results', {}).get('vulnerability', []) + self.assertEqual(sorted(v['name'] for v in vulns), ['SQLi', 'XSS'])