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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 23 additions & 9 deletions explorer/assistant/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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()

Expand Down Expand Up @@ -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)

Expand Down
8 changes: 6 additions & 2 deletions explorer/assistant/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
53 changes: 47 additions & 6 deletions explorer/tests/test_assistant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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"
Expand Down Expand Up @@ -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):
Expand All @@ -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",)]
Expand All @@ -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")]
Expand All @@ -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")]
Expand All @@ -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()
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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):
Expand Down
Loading