From a84352e1eec1bfd024085525bcdea2aee27ae26c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=89=A7=E7=AC=94?= Date: Mon, 7 Sep 2026 22:27:58 -0700 Subject: [PATCH] feat: make table sample data in SQL prompts configurable Add a default-on TABLE_SAMPLE_DATA_ENABLED setting. When disabled, skip automatic sample queries and omit sample values from SQL prompt templates. Preserve schema context and explicit query/preview behavior. Add isolated regression tests and document configuration and limitations. Refs dataease/SQLBot#1291 --- backend/README.md | 50 +++- backend/apps/datasource/crud/datasource.py | 3 + .../apps/template/generate_sql/generator.py | 8 +- backend/common/core/config.py | 4 + tests/test_table_sample_data.py | 241 ++++++++++++++++++ 5 files changed, 304 insertions(+), 2 deletions(-) create mode 100644 tests/test_table_sample_data.py diff --git a/backend/README.md b/backend/README.md index b6dbbb6d3..c854cfb8a 100644 --- a/backend/README.md +++ b/backend/README.md @@ -1 +1,49 @@ -# FastAPI Project - Backend \ No newline at end of file +# FastAPI Project - Backend + +## Table sample data in SQL generation + +`TABLE_SAMPLE_DATA_ENABLED` controls whether SQLBot fetches table sample rows +and includes them in SQL-generation prompts. It defaults to `true` to preserve +existing behavior, including the current three-row sample limit. + +To disable automatic sampling, set the following environment variable (or add it +to the project-root `.env` used when running the backend from `backend/`): + +```dotenv +TABLE_SAMPLE_DATA_ENABLED=false +``` + +For Docker Compose, pass it explicitly to the SQLBot service: + +```yaml +environment: + TABLE_SAMPLE_DATA_ENABLED: "false" +``` + +A Compose `.env` file alone does not automatically pass all its variables into a +container. For `docker run`, use `-e TABLE_SAMPLE_DATA_ENABLED=false`. +Settings are loaded at process startup: restart a source deployment, or recreate +the container with the new environment configuration. + +When disabled, the automatic sample-data helper returns before reading table +metadata or sample rows. The SQL prompt template also ignores pre-existing +`sample_data` values. Schema context, explicit SQL queries, manual data previews +and their existing permission checks remain unchanged. Less sample context may +affect SQL-generation quality. + +This is not a global data-loss-prevention switch. User messages, field comments, +SQL examples, query results and analysis prompts can still contain business +data. Existing logs are not deleted or retroactively redacted. + +### Focused regression tests + +From the repository root, with the backend development dependencies available: + +```bash +python -m pytest -q tests/test_table_sample_data.py +``` + +The tests execute actual source function definitions with dependency doubles to +avoid initializing database drivers, embedding models or X-Pack. They cover +configuration parsing, sampling, prompt rendering and unaffected query/preview +behavior; they are not full application or model-provider integration tests. diff --git a/backend/apps/datasource/crud/datasource.py b/backend/apps/datasource/crud/datasource.py index 11720a783..3d394ce92 100644 --- a/backend/apps/datasource/crud/datasource.py +++ b/backend/apps/datasource/crud/datasource.py @@ -500,6 +500,9 @@ def get_table_sample_data(ds: CoreDatasource, table_name: str, fields: list) -> def get_tables_sample_data(session: SessionDep, current_user: CurrentUser, ds: CoreDatasource, table_list: list[str] = None) -> str: """Get sample data (3 rows) for all tables to help AI understand the data""" + if not settings.TABLE_SAMPLE_DATA_ENABLED: + return "" + table_objs = get_table_obj_by_ds(session=session, current_user=current_user, ds=ds) if len(table_objs) == 0: return "" diff --git a/backend/apps/template/generate_sql/generator.py b/backend/apps/template/generate_sql/generator.py index 07eb61972..a13f3bb64 100644 --- a/backend/apps/template/generate_sql/generator.py +++ b/backend/apps/template/generate_sql/generator.py @@ -2,11 +2,17 @@ from apps.db.constant import DB from apps.template.template import get_base_template, get_sql_template as get_base_sql_template +from common.core.config import settings def get_sql_template(): template = get_base_template() - return template['template']['sql'] + sql_template = template['template']['sql'] + if not settings.TABLE_SAMPLE_DATA_ENABLED: + # Do not mutate shared templates or render pre-existing sample values. + sql_template = sql_template.copy() + sql_template['generate_basic_info'] = sql_template['generate_basic_info'].replace('{sample_data}', '') + return sql_template def get_sql_example_template(db_type: Union[str, DB]): diff --git a/backend/common/core/config.py b/backend/common/core/config.py index 4b9baeaec..a350e5b77 100644 --- a/backend/common/core/config.py +++ b/backend/common/core/config.py @@ -126,6 +126,9 @@ def SQLALCHEMY_DATABASE_URI(self) -> PostgresDsn | str: PG_POOL_RECYCLE: int = 3600 PG_POOL_PRE_PING: bool = True + # Include table sample rows in SQL generation context (disable for sensitive data). + TABLE_SAMPLE_DATA_ENABLED: bool = True + TABLE_EMBEDDING_ENABLED: bool = True TABLE_EMBEDDING_COUNT: int = 10 DS_EMBEDDING_COUNT: int = 10 @@ -138,6 +141,7 @@ def SQLALCHEMY_DATABASE_URI(self) -> PostgresDsn | str: 'PARSE_REASONING_BLOCK_ENABLED', 'PG_POOL_PRE_PING', 'TABLE_EMBEDDING_ENABLED', + 'TABLE_SAMPLE_DATA_ENABLED', mode='before') @classmethod def lowercase_bool(cls, v: Any) -> Any: diff --git a/tests/test_table_sample_data.py b/tests/test_table_sample_data.py new file mode 100644 index 000000000..bc2512a16 --- /dev/null +++ b/tests/test_table_sample_data.py @@ -0,0 +1,241 @@ +"""Isolated regression tests for the table-sample context switch (#1291). + +Load function definitions from the actual sources, not copied implementations. +This avoids importing database drivers, embedding models and X-Pack just to test +sampling and prompt construction. Settings uses the real Pydantic loader. These +are unit tests, not full application or model-provider integration tests. +""" + +import __future__ +import ast +import importlib.util +import json +from copy import deepcopy +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock, Mock + +import pytest +from pydantic import ValidationError + +ROOT = Path(__file__).resolve().parents[1] +BACKEND = ROOT / "backend" +SENTINEL = "SQLBOT_SAMPLE_SENTINEL_1291" + + +def load_functions(relative_path, names, namespace): + """Execute unchanged source definitions with explicit dependency doubles.""" + path = BACKEND / relative_path + source = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + definitions = [ + node for node in source.body + if isinstance(node, ast.FunctionDef) and node.name in names + ] + assert {node.name for node in definitions} == set(names) + module = ast.Module(body=definitions, type_ignores=[]) + code = compile( + module, str(path), "exec", flags=__future__.annotations.compiler_flag, + dont_inherit=True, + ) + exec(code, namespace) + return namespace + + +@pytest.fixture +def settings_class(monkeypatch, tmp_path): + # Do not load a developer's .env or carry a previous test's configuration. + monkeypatch.chdir(tmp_path) + monkeypatch.delenv("TABLE_SAMPLE_DATA_ENABLED", raising=False) + spec = importlib.util.spec_from_file_location( + "sqlbot_test_settings", BACKEND / "common/core/config.py" + ) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module.Settings + + +def test_sample_context_is_enabled_by_default(settings_class): + assert settings_class(_env_file=None).TABLE_SAMPLE_DATA_ENABLED is True + + +@pytest.mark.parametrize("value, expected", [ + ("true", True), ("false", False), ("TRUE", True), ("FALSE", False), + (" True ", True), (" False ", False), ("1", True), ("0", False), +]) +def test_environment_boolean_parsing(settings_class, monkeypatch, value, expected): + monkeypatch.setenv("TABLE_SAMPLE_DATA_ENABLED", value) + assert settings_class(_env_file=None).TABLE_SAMPLE_DATA_ENABLED is expected + + +def test_invalid_boolean_is_rejected(settings_class, monkeypatch): + monkeypatch.setenv("TABLE_SAMPLE_DATA_ENABLED", "not-a-boolean") + with pytest.raises(ValidationError): + settings_class(_env_file=None) + + +def test_dotenv_can_disable_sample_context(settings_class, tmp_path): + env_file = tmp_path / "sample.env" + env_file.write_text("TABLE_SAMPLE_DATA_ENABLED=false\n", encoding="utf-8") + assert settings_class(_env_file=env_file).TABLE_SAMPLE_DATA_ENABLED is False + + +@pytest.fixture +def sampling(): + fields = [SimpleNamespace( + field_name="sku", field_type="varchar", custom_comment="Product code", + )] + table = SimpleNamespace( + id=1, table_name="orders", custom_comment="Order records", embedding=None, + ) + table_obj = SimpleNamespace(schema="shop", table=table, fields=fields) + namespace = { + "settings": SimpleNamespace( + TABLE_SAMPLE_DATA_ENABLED=True, TABLE_EMBEDDING_ENABLED=False, + ), + "get_table_obj_by_ds": Mock(return_value=[table_obj]), + "exec_sql": Mock(return_value={"data": [{"sku": SENTINEL}]}), + "DB": SimpleNamespace(get_db=lambda _: SimpleNamespace(prefix='"', suffix='"')), + "equals_ignore_case": lambda first, second: first.lower() == second.lower(), + } + return load_functions( + "apps/datasource/crud/datasource.py", + {"get_tables_sample_data", "get_table_sample_data", "get_table_schema", "execSql", "preview"}, + namespace, + ) + + +def test_disabled_does_not_read_metadata_or_rows(sampling): + sampling["settings"].TABLE_SAMPLE_DATA_ENABLED = False + result = sampling["get_tables_sample_data"]( + object(), object(), SimpleNamespace(type="pg"), + ) + assert result == "" + sampling["get_table_obj_by_ds"].assert_not_called() + sampling["exec_sql"].assert_not_called() + + +def test_enabled_preserves_sample_query_and_output(sampling): + session, user = object(), object() + datasource = SimpleNamespace(type="pg") + result = sampling["get_tables_sample_data"](session, user, datasource) + sampling["get_table_obj_by_ds"].assert_called_once_with( + session=session, current_user=user, ds=datasource, + ) + sampling["exec_sql"].assert_called_once_with( + ds=datasource, sql='SELECT "sku" FROM "orders" LIMIT 3', origin_column=True, + ) + assert result == '# Table: orders\n[\n {\n "sku": "' + SENTINEL + '"\n }\n]' + + +@pytest.mark.parametrize("table_list", [[], ["another_table"]]) +def test_selected_table_filter_is_preserved(sampling, table_list): + assert sampling["get_tables_sample_data"]( + object(), object(), SimpleNamespace(type="pg"), table_list, + ) == "" + sampling["exec_sql"].assert_not_called() + + +def test_no_authorized_fields_does_not_sample(sampling): + sampling["get_table_obj_by_ds"].return_value[0].fields = [] + assert sampling["get_tables_sample_data"]( + object(), object(), SimpleNamespace(type="pg"), + ) == "" + sampling["exec_sql"].assert_not_called() + + +def test_sample_limit_remains_three_rows(sampling): + sampling["exec_sql"].return_value = {"data": [{"sku": n} for n in range(5)]} + result = sampling["get_tables_sample_data"]( + object(), object(), SimpleNamespace(type="pg"), + ) + assert json.loads(result.split("\n", 1)[1]) == [{"sku": 0}, {"sku": 1}, {"sku": 2}] + + +def test_disabled_does_not_remove_schema(sampling): + sampling["settings"].TABLE_SAMPLE_DATA_ENABLED = False + schema, tables = sampling["get_table_schema"]( + object(), object(), SimpleNamespace(type="pg", table_relation=None), + "Show orders", embedding=False, + ) + assert tables == ["orders"] + assert "shop.orders" in schema + assert "sku:varchar, Product code" in schema + sampling["exec_sql"].assert_not_called() + + +def test_disabled_does_not_block_explicit_sql_execution(sampling): + sampling["settings"].TABLE_SAMPLE_DATA_ENABLED = False + sampling["CoreDatasource"] = SimpleNamespace(id=1) + sampling["select"] = Mock() + session, datasource = Mock(), object() + session.exec.return_value.first.return_value = datasource + result = sampling["execSql"](session, 1, "SELECT 1") + sampling["exec_sql"].assert_called_once_with(datasource, "SELECT 1", True) + assert result is sampling["exec_sql"].return_value + + +@pytest.fixture +def templates(): + template = {"template": {"sql": { + "generate_basic_info": ( + "{engine}\n{schema}\n" + "{sample_data}" + ), + "other_rule": "Preserve other rules", + }}} + namespace = { + "settings": SimpleNamespace(TABLE_SAMPLE_DATA_ENABLED=True), + "get_base_template": Mock(return_value=template), + } + return load_functions( + "apps/template/generate_sql/generator.py", {"get_sql_template"}, namespace, + ) + + +@pytest.mark.parametrize("enabled", [True, False]) +def test_template_controls_sample_insertion_without_removing_schema(templates, enabled): + templates["settings"].TABLE_SAMPLE_DATA_ENABLED = enabled + result = templates["get_sql_template"]() + rendered = result["generate_basic_info"].format( + engine="PostgreSQL", schema="orders(sku varchar)", sample_data=SENTINEL, + ) + assert (SENTINEL in rendered) is enabled + assert "PostgreSQL" in rendered + assert "orders(sku varchar)" in rendered + assert result["other_rule"] == "Preserve other rules" + + +def test_disabling_does_not_mutate_shared_templates(templates): + original = deepcopy(templates["get_base_template"].return_value) + templates["settings"].TABLE_SAMPLE_DATA_ENABLED = False + disabled = templates["get_sql_template"]() + assert "{sample_data}" not in disabled["generate_basic_info"] + assert templates["get_base_template"].return_value == original + templates["settings"].TABLE_SAMPLE_DATA_ENABLED = True + assert "{sample_data}" in templates["get_sql_template"]()["generate_basic_info"] + + +def test_disabled_does_not_block_manual_preview(sampling): + sampling["settings"].TABLE_SAMPLE_DATA_ENABLED = False + sampling.update({ + "CoreDatasource": MagicMock(), "CoreField": MagicMock(), "CoreTable": MagicMock(), + "is_normal_user": Mock(return_value=False), + "get_engine_config": Mock(return_value=SimpleNamespace(dbSchema="public")), + }) + datasource = SimpleNamespace(type="excel") + field = SimpleNamespace(field_name="sku", checked=True) + table = SimpleNamespace(id=1, table_name="orders") + ds_query, field_query, table_query = Mock(), Mock(), Mock() + ds_query.filter.return_value.first.return_value = datasource + field_query.filter.return_value.order_by.return_value.all.return_value = [field] + table_query.filter.return_value.first.return_value = table + session = Mock() + session.query.side_effect = [ds_query, field_query, table_query] + result = sampling["preview"](session, object(), 1, SimpleNamespace(table=table)) + sampling["exec_sql"].assert_called_once() + args = sampling["exec_sql"].call_args.args + assert args[0] is datasource + assert '"public"."orders"' in args[1] + assert "LIMIT 100" in args[1] + assert args[2] is True + assert result is sampling["exec_sql"].return_value