Skip to content
Closed
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
50 changes: 49 additions & 1 deletion backend/README.md
Original file line number Diff line number Diff line change
@@ -1 +1,49 @@
# FastAPI Project - Backend
# 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.
3 changes: 3 additions & 0 deletions backend/apps/datasource/crud/datasource.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ""
Expand Down
8 changes: 7 additions & 1 deletion backend/apps/template/generate_sql/generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]):
Expand Down
4 changes: 4 additions & 0 deletions backend/common/core/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand Down
241 changes: 241 additions & 0 deletions tests/test_table_sample_data.py
Original file line number Diff line number Diff line change
@@ -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": (
"<db-engine>{engine}</db-engine>\n<schema>{schema}</schema>\n"
"<sample-data>{sample_data}</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