From 3c5db76a7f2f3dcafa7390b8ef97bb770ec69f4a Mon Sep 17 00:00:00 2001 From: Chris Clark Date: Wed, 23 Sep 2026 10:52:24 -0400 Subject: [PATCH] Require permissions on assistant endpoints and validate sampled table names AssistantHelpView and AssistantHistoryApiView had no permission check, and table names from the request's selected_tables were interpolated directly into the row-sampling SQL. - Both assistant API views now require change_permission, matching the playground where the assistant is used. - build_prompt only samples tables that exist in the connection's schema (respecting include/exclude prefixes), using the introspected name. - sample_rows_from_table quotes the table name via connection.ops.quote_name. Co-Authored-By: Claude --- explorer/assistant/utils.py | 32 +++++++++++++------ explorer/assistant/views.py | 8 +++-- explorer/tests/test_assistant.py | 53 ++++++++++++++++++++++++++++---- 3 files changed, 76 insertions(+), 17 deletions(-) diff --git a/explorer/assistant/utils.py b/explorer/assistant/utils.py index 45cf1626..6fed4932 100644 --- a/explorer/assistant/utils.py +++ b/explorer/assistant/utils.py @@ -38,11 +38,21 @@ def extract_response(r): return r[-1].content -def table_schema(db_connection, table_name): - schema = schema_info(db_connection) +def find_table(db_connection, table_name): + """ + Look up a table in the connection's (include/exclude-filtered) schema, case-insensitively. + Returns a (real_table_name, schema) tuple, or None if the table is not in the schema. + """ + schema = schema_info(db_connection) or [] s = [table for table in schema if table[0].lower() == table_name.lower()] if len(s): - return s[0][1] + return s[0] + + +def table_schema(db_connection, table_name): + table = find_table(db_connection, table_name) + if table: + return table[1] def sample_rows_from_table(connection, table_name): @@ -62,7 +72,8 @@ def sample_rows_from_table(connection, table_name): """ cursor = connection.cursor() try: - cursor.execute(f"SELECT * FROM {table_name} LIMIT {ROW_SAMPLE_SIZE}") + # table_name must be validated against the schema by the caller; quote it regardless. + cursor.execute(f"SELECT * FROM {connection.ops.quote_name(table_name)} LIMIT {ROW_SAMPLE_SIZE}") ret = [[header[0] for header in cursor.description]] rows = cursor.fetchall() @@ -159,14 +170,17 @@ def build_prompt(db_connection, assistant_request, included_tables, query_error= error_chunk = f"## Query Error ##\n{query_error}" if query_error else None sql_chunk = f"## Existing User-Written SQL ##\n{sql}" if sql else None request_chunk = f"## User's Request to Assistant ##\n{assistant_request}" + # Only sample tables that actually exist in the connection's schema. The table names come + # from the request, so they must never be interpolated into SQL unvalidated. + known_tables = [find_table(db_connection, t) for t in included_tables] table_chunks = [ TablePromptData( - name=t, - schema=table_schema(db_connection, t), - sample=sample_rows_from_table(db_connection.as_django_connection(), t), - annotation=get_relevant_annotation(db_connection, t) + name=name, + schema=schema, + sample=sample_rows_from_table(db_connection.as_django_connection(), name), + annotation=get_relevant_annotation(db_connection, name) ).render() - for t in included_tables + for name, schema in [t for t in known_tables if t] ] few_shot_chunk = get_few_shot_chunk(db_connection, included_tables) diff --git a/explorer/assistant/views.py b/explorer/assistant/views.py index 2571e216..6a30a93c 100644 --- a/explorer/assistant/views.py +++ b/explorer/assistant/views.py @@ -61,7 +61,9 @@ def run_assistant(request_data, user): return response_text -class AssistantHelpView(View): +class AssistantHelpView(PermissionRequiredMixin, View): + + permission_required = "change_permission" def post(self, request, *args, **kwargs): try: @@ -106,7 +108,9 @@ class TableDescriptionDeleteView(PermissionRequiredMixin, ExplorerContextMixin, success_url = reverse_lazy("table_description_list") -class AssistantHistoryApiView(View): +class AssistantHistoryApiView(PermissionRequiredMixin, View): + + permission_required = "change_permission" def post(self, request, *args, **kwargs): try: diff --git a/explorer/tests/test_assistant.py b/explorer/tests/test_assistant.py index 60c7ed2c..bc20e8c4 100644 --- a/explorer/tests/test_assistant.py +++ b/explorer/tests/test_assistant.py @@ -73,7 +73,7 @@ def test_build_prompt_with_vendor_only(self, mock_get_item): self.assertIn("sqlite", result["system"]) @patch("explorer.assistant.utils.sample_rows_from_table", return_value="sample data") - @patch("explorer.assistant.utils.table_schema", return_value=[]) + @patch("explorer.assistant.utils.find_table", return_value=("foo", [])) @patch("explorer.models.ExplorerValue.objects.get_item") def test_build_prompt_with_sql_and_annotation(self, mock_get_item, mock_table_schema, mock_sample_rows): mock_get_item.return_value.value = "system prompt" @@ -87,7 +87,7 @@ def test_build_prompt_with_sql_and_annotation(self, mock_get_item, mock_table_sc self.assertIn("Usage Notes:\nannotated", result["user"]) @patch("explorer.assistant.utils.sample_rows_from_table", return_value="sample data") - @patch("explorer.assistant.utils.table_schema", return_value=[]) + @patch("explorer.assistant.utils.find_table", return_value=("magic", [])) @patch("explorer.models.ExplorerValue.objects.get_item") def test_build_prompt_with_few_shot(self, mock_get_item, mock_table_schema, mock_sample_rows): mock_get_item.return_value.value = "system prompt" @@ -128,6 +128,19 @@ def test_build_prompt_with_extra_tables_fitting_window(self, mock_get_item): self.assertIn("## Information for Table 'explorer_query' ##", result["user"]) self.assertIn("Sample rows:\nid | title", result["user"]) + @patch("explorer.assistant.utils.sample_rows_from_table") + @patch("explorer.models.ExplorerValue.objects.get_item") + def test_build_prompt_skips_tables_not_in_schema(self, mock_get_item, mock_sample_rows): + mock_get_item.return_value.value = "system prompt" + malicious = "explorer_query UNION SELECT username, password FROM auth_user --" + + result = build_prompt(default_db_connection(), "Help me with SQL", + [malicious, "explorer_query"]) + + mock_sample_rows.assert_called_once() + self.assertEqual(mock_sample_rows.call_args[0][1], "explorer_query") + self.assertNotIn("auth_user", result["user"]) + @unittest.skipIf(not app_settings.has_assistant(), "assistant not enabled") class TestPromptContext(TestCase): @@ -141,7 +154,7 @@ def test_retrieves_sample_rows(self): self.assertEqual(len(ret), ROW_SAMPLE_SIZE+1) # includes header row def test_truncates_long_strings(self): - c = MagicMock + c = MagicMock() mock_cursor = MagicMock() long_string = "a" * 600 mock_cursor.description = [("col1",), ("col2",)] @@ -160,7 +173,7 @@ def test_binary_data(self): long_binary = b"a" * 600 # Mock database connection and cursor - c = MagicMock + c = MagicMock() mock_cursor = MagicMock() mock_cursor.description = [("col1",), ("col2",)] mock_cursor.fetchall.return_value = [(long_binary, b"short binary")] @@ -176,7 +189,7 @@ def test_binary_data(self): def test_handles_various_data_types(self): # Mock database connection and cursor - c = MagicMock + c = MagicMock() mock_cursor = MagicMock() mock_cursor.description = [("col1",), ("col2",), ("col3",)] mock_cursor.fetchall.return_value = [(123, 45.67, "normal string")] @@ -192,7 +205,7 @@ def test_handles_various_data_types(self): self.assertEqual(row[2], "normal string") def test_handles_operational_error(self): - c = MagicMock + c = MagicMock() mock_cursor = MagicMock() mock_cursor.execute.side_effect = OperationalError("Test OperationalError") c.cursor = MagicMock() @@ -258,6 +271,12 @@ def test_sample_rows_from_table(self): self.assertTrue("First Query" in ret) self.assertTrue("Second Query" in ret) + def test_sample_rows_from_table_quotes_table_name(self): + from explorer.assistant.utils import sample_rows_from_table + ret = sample_rows_from_table(conn(), "explorer_query; DROP TABLE explorer_query") + self.assertEqual(ret, [["no such table: explorer_query; DROP TABLE explorer_query"]]) + self.assertTrue(conn().introspection.table_names().count("explorer_query")) + def test_sample_rows_from_tables_no_table_match(self): from explorer.assistant.utils import sample_rows_from_table SimpleQueryFactory(title="First Query") @@ -303,6 +322,28 @@ def test_get_relevant_annotations(self): self.assertEqual(relevant2.id, res2.id) +class TestAssistantViewPermissions(TestCase): + + def assert_denied(self, url_name): + with patch("explorer.assistant.views.run_assistant") as mock_run: + resp = self.client.post(reverse(url_name), + data=json.dumps({"connection_id": default_db_connection().id, + "selected_tables": ["explorer_query"]}), + content_type="application/json") + mock_run.assert_not_called() + self.assertNotIn("application/json", resp["Content-Type"]) + + def test_anonymous_user_cannot_use_assistant(self): + self.assert_denied("assistant") + self.assert_denied("assistant_history") + + def test_non_staff_user_cannot_use_assistant(self): + User.objects.create_user("user", "user@user.com", "pwd") + self.client.login(username="user", password="pwd") + self.assert_denied("assistant") + self.assert_denied("assistant_history") + + class TestAssistantHistoryApiView(TestCase): def setUp(self):