Skip to content
Open
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
8 changes: 8 additions & 0 deletions lib/crewai-tools/src/crewai_tools/tools/nl2sql/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,14 @@ def researcher(self) -> Agent:
)
```

The tool infers the SQL dialect from `db_uri` and includes dialect-specific
query-generation guidance for the agent. You can override the detected dialect
when needed:

```python
nl2sql = NL2SQLTool(db_uri="sqlite:///example.db", dialect="sqlite")
```

## Example

The primary task goal was:
Expand Down
34 changes: 34 additions & 0 deletions lib/crewai-tools/src/crewai_tools/tools/nl2sql/nl2sql_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

try:
from sqlalchemy import create_engine, text
from sqlalchemy.engine import make_url
from sqlalchemy.orm import sessionmaker

SQLALCHEMY_AVAILABLE = True
Expand All @@ -24,6 +25,19 @@

logger = logging.getLogger(__name__)

_DIALECT_GUIDANCE = {
"postgresql": (
"Generate PostgreSQL-compatible SQL. PostgreSQL features such as ILIKE, "
"DATE_TRUNC, INTERVAL, and :: casts are available."
),
"sqlite": (
"Generate SQLite-compatible SQL. Use LOWER(column) LIKE LOWER(pattern) "
"for case-insensitive matching, SQLite date/time functions such as "
"strftime(), and || for string concatenation. Do not use ILIKE, "
"DATE_TRUNC, INTERVAL, or :: casts."
),
}

# Commands allowed in read-only mode
# NOTE: WITH is intentionally excluded — writable CTEs start with WITH, so the
# CTE body must be inspected separately (see _validate_statement).
Expand Down Expand Up @@ -237,6 +251,14 @@ class NL2SQLTool(BaseTool):
title="Database URI",
description="The URI of the database to connect to.",
)
dialect: str | None = Field(
default=None,
title="SQL Dialect",
description=(
"SQL dialect used when generating queries. Inferred from db_uri when "
"not provided."
),
)
allow_dml: bool = Field(
default=False,
title="Allow DML",
Expand Down Expand Up @@ -269,6 +291,18 @@ def model_post_init(self, __context: Any) -> None:
"`pip install crewai-tools[sqlalchemy]`"
)

self.dialect = (
self.dialect.strip().lower()
if self.dialect and self.dialect.strip()
else make_url(self.db_uri).get_backend_name()
Comment thread
coderabbitai[bot] marked this conversation as resolved.
)
guidance = _DIALECT_GUIDANCE.get(
self.dialect,
f"Generate SQL compatible with the {self.dialect} dialect.",
)
if guidance not in self.description:
self.description = f"{self.description.rstrip()} {guidance}"

if self.allow_dml:
logger.warning(
"NL2SQLTool: allow_dml=True — write operations (INSERT/UPDATE/"
Expand Down
61 changes: 61 additions & 0 deletions lib/crewai-tools/tests/tools/test_nl2sql_dialect_guidance.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
"""Tests for dialect-specific NL2SQLTool generation guidance."""

from unittest.mock import patch

import pytest


pytest.importorskip("sqlalchemy")

from crewai_tools.tools.nl2sql.nl2sql_tool import NL2SQLTool # noqa: E402


def _make_tool(db_uri: str, **kwargs: object) -> NL2SQLTool:
with (
patch.object(NL2SQLTool, "_fetch_available_tables", return_value=[]),
patch.object(NL2SQLTool, "_fetch_all_available_columns", return_value=[]),
):
return NL2SQLTool(db_uri=db_uri, **kwargs)


def test_sqlite_dialect_is_inferred_from_uri() -> None:
tool = _make_tool("sqlite:///example.db")

assert tool.dialect == "sqlite"
assert "Generate SQLite-compatible SQL" in tool.description
assert "Do not use ILIKE" in tool.description


def test_postgresql_dialect_is_inferred_from_uri() -> None:
tool = _make_tool("postgresql://user:password@localhost/database")

assert tool.dialect == "postgresql"
assert "Generate PostgreSQL-compatible SQL" in tool.description
assert "ILIKE" in tool.description


def test_explicit_dialect_overrides_uri() -> None:
tool = _make_tool(
"postgresql://user:password@localhost/database", dialect=" SQLite "
)

assert tool.dialect == "sqlite"
assert "Generate SQLite-compatible SQL" in tool.description


def test_custom_description_is_preserved() -> None:
tool = _make_tool(
"sqlite:///example.db",
description="Use the reporting database.",
)

assert tool.description.startswith("Use the reporting database.")
assert "Generate SQLite-compatible SQL" in tool.description


def test_dialect_guidance_is_not_duplicated_on_restore() -> None:
tool = _make_tool("sqlite:///example.db")

restored = _make_tool(**tool.model_dump())

assert restored.description.count("Generate SQLite-compatible SQL") == 1
Loading