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
39 changes: 36 additions & 3 deletions llama-index-core/llama_index/core/utilities/sql_wrapper.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""SQL wrapper around SQLDatabase in langchain."""

from typing import Any, Dict, Iterable, List, Optional, Tuple
import re
from typing import Any, Dict, Iterable, List, Optional, Set, Tuple

from sqlalchemy import MetaData, create_engine, insert, inspect, text
from sqlalchemy.engine import Engine
Expand Down Expand Up @@ -212,6 +213,39 @@ def truncate_word(self, content: Any, *, length: int, suffix: str = "...") -> st

return content[: length - len(suffix)].rsplit(" ", 1)[0] + suffix

def _add_schema_prefix(self, command: str) -> str:
"""
Add schema prefix to table names in FROM/JOIN clauses.

Preserves CTE (Common Table Expression) names and already
schema-qualified identifiers so they are not double-prefixed.
"""
# Collect CTE names defined in WITH clauses
cte_names: Set[str] = set()
# First CTE: WITH [RECURSIVE] name AS (
for m in re.finditer(
r"\bWITH\s+(?:RECURSIVE\s+)?(\w+)\s+AS\s*\(", command, re.IGNORECASE
):
cte_names.add(m.group(1).lower())
# Subsequent CTEs: ), name AS (
for m in re.finditer(r"\)\s*,\s*(\w+)\s+AS\s*\(", command, re.IGNORECASE):
cte_names.add(m.group(1).lower())

def _replace(match: re.Match) -> str:
keyword = match.group(1)
table_ref = match.group(2)
# Skip CTE references and already schema-qualified names
if table_ref.lower() in cte_names or "." in table_ref:
return match.group(0)
return f"{keyword}{self._schema}.{table_ref}"

return re.sub(
r"\b((?:FROM|JOIN)\s+)(\w+(?:\.\w+)?)",
_replace,
command,
flags=re.IGNORECASE,
)

def run_sql(self, command: str) -> Tuple[str, Dict]:
"""
Execute a SQL statement and return a string representing the results.
Expand All @@ -222,8 +256,7 @@ def run_sql(self, command: str) -> Tuple[str, Dict]:
with self._engine.begin() as connection:
try:
if self._schema:
command = command.replace("FROM ", f"FROM {self._schema}.")
command = command.replace("JOIN ", f"JOIN {self._schema}.")
command = self._add_schema_prefix(command)
cursor = connection.execute(text(command))
except (ProgrammingError, OperationalError) as exc:
raise NotImplementedError(
Expand Down
105 changes: 105 additions & 0 deletions llama-index-core/tests/utilities/test_sql_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,3 +96,108 @@ def test_long_string_no_truncation(sql_database: SQLDatabase) -> None:
result_str, _ = sql_database.run_sql("SELECT * FROM test_table;")

assert result_str == f"[(1, '{long_string}')]"


# --- Tests for _add_schema_prefix (CTE support) ---


def test_schema_prefix_simple_from(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
result = sql_database._add_schema_prefix("SELECT * FROM users")
assert result == "SELECT * FROM myschema.users"


def test_schema_prefix_simple_join(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
result = sql_database._add_schema_prefix(
"SELECT * FROM users JOIN orders ON users.id = orders.user_id"
)
assert "FROM myschema.users" in result
assert "JOIN myschema.orders" in result


def test_schema_prefix_preserves_cte_names(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
cmd = "WITH my_cte AS (SELECT * FROM users) SELECT * FROM my_cte"
result = sql_database._add_schema_prefix(cmd)
assert "FROM myschema.users" in result
assert "FROM my_cte" in result
assert "myschema.my_cte" not in result


def test_schema_prefix_multiple_ctes(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
cmd = (
"WITH cte1 AS (SELECT * FROM users), "
"cte2 AS (SELECT * FROM orders) "
"SELECT * FROM cte1 JOIN cte2 ON cte1.id = cte2.user_id"
)
result = sql_database._add_schema_prefix(cmd)
assert "FROM myschema.users" in result
assert "FROM myschema.orders" in result
assert "FROM cte1" in result
assert "JOIN cte2" in result
assert "myschema.cte1" not in result
assert "myschema.cte2" not in result


def test_schema_prefix_recursive_cte(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
cmd = (
"WITH RECURSIVE subordinates AS ("
"SELECT id, name FROM employees WHERE manager_id IS NULL "
"UNION ALL "
"SELECT e.id, e.name FROM employees e "
"JOIN subordinates s ON e.manager_id = s.id"
") SELECT * FROM subordinates"
)
result = sql_database._add_schema_prefix(cmd)
assert "FROM myschema.employees" in result
assert "FROM subordinates" in result
assert "myschema.subordinates" not in result


def test_schema_prefix_skips_already_qualified(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
result = sql_database._add_schema_prefix("SELECT * FROM other_schema.users")
assert "FROM other_schema.users" in result
assert "myschema.other_schema" not in result


def test_schema_prefix_left_join(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
result = sql_database._add_schema_prefix(
"SELECT * FROM users LEFT JOIN orders ON users.id = orders.user_id"
)
assert "FROM myschema.users" in result
assert "JOIN myschema.orders" in result


def test_schema_prefix_subquery_not_prefixed(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
result = sql_database._add_schema_prefix(
"SELECT * FROM (SELECT id FROM users) AS sub"
)
# The subquery opener '(' should not be treated as a table name
assert "FROM myschema.users" in result
assert "myschema.(SELECT" not in result


def test_schema_prefix_cte_join_inside_body(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
cmd = (
"WITH active AS (SELECT * FROM users WHERE active = 1) "
"SELECT * FROM active JOIN orders ON active.id = orders.user_id"
)
result = sql_database._add_schema_prefix(cmd)
assert "FROM myschema.users" in result
assert "JOIN myschema.orders" in result
assert "FROM active " in result
assert "myschema.active" not in result


def test_schema_prefix_case_insensitive(sql_database: SQLDatabase) -> None:
sql_database._schema = "myschema"
result = sql_database._add_schema_prefix("select * from users join orders on 1=1")
assert "from myschema.users" in result
assert "join myschema.orders" in result
Loading