Skip to content

Commit e3f7d1a

Browse files
committed
fix mysql search table name validation
1 parent 5d4851e commit e3f7d1a

2 files changed

Lines changed: 133 additions & 1 deletion

File tree

lib/crewai-tools/src/crewai_tools/tools/mysql_search_tool/mysql_search_tool.py

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import re
12
from typing import Any
23

34
from pydantic import BaseModel, Field
@@ -6,6 +7,26 @@
67
from crewai_tools.tools.rag.rag_tool import RagTool
78

89

10+
_MYSQL_IDENTIFIER_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_$]*$")
11+
12+
13+
def _quote_mysql_table_name(table_name: str) -> str:
14+
identifier_parts = table_name.split(".")
15+
if (
16+
not identifier_parts
17+
or len(identifier_parts) > 2
18+
or any(
19+
not _MYSQL_IDENTIFIER_PATTERN.fullmatch(part) for part in identifier_parts
20+
)
21+
):
22+
raise ValueError(
23+
"MySQL table_name must be a valid table identifier or schema.table "
24+
"identifier"
25+
)
26+
27+
return ".".join(f"`{part}`" for part in identifier_parts)
28+
29+
930
class MySQLSearchToolSchema(BaseModel):
1031
"""Input for MySQLSearchTool."""
1132

@@ -32,7 +53,8 @@ def add( # type: ignore[override]
3253
table_name: str,
3354
**kwargs: Any,
3455
) -> None:
35-
super().add(f"SELECT * FROM {table_name};", **kwargs) # noqa: S608
56+
quoted_table_name = _quote_mysql_table_name(table_name)
57+
super().add(f"SELECT * FROM {quoted_table_name};", **kwargs) # noqa: S608
3658

3759
def _run( # type: ignore[override]
3860
self,
Lines changed: 110 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,110 @@
1+
from unittest.mock import MagicMock, patch
2+
3+
import pytest
4+
5+
from crewai_tools.rag.data_types import DataType
6+
from crewai_tools.tools.mysql_search_tool.mysql_search_tool import MySQLSearchTool
7+
from crewai_tools.tools.rag.rag_tool import RagTool
8+
9+
10+
@pytest.fixture
11+
def mock_rag_client() -> MagicMock:
12+
mock_client = MagicMock()
13+
mock_client.get_or_create_collection = MagicMock(return_value=None)
14+
mock_client.add_documents = MagicMock(return_value=None)
15+
mock_client.search = MagicMock(return_value=[])
16+
return mock_client
17+
18+
19+
def create_mysql_search_tool(
20+
mock_rag_client: MagicMock, table_name: str
21+
) -> MySQLSearchTool:
22+
with (
23+
patch(
24+
"crewai_tools.adapters.crewai_rag_adapter.get_rag_client",
25+
return_value=mock_rag_client,
26+
),
27+
patch(
28+
"crewai_tools.adapters.crewai_rag_adapter.create_client",
29+
return_value=mock_rag_client,
30+
),
31+
):
32+
return MySQLSearchTool(
33+
db_uri="mysql://user:password@localhost:3306/test_database",
34+
table_name=table_name,
35+
)
36+
37+
38+
@pytest.mark.parametrize(
39+
("table_name", "expected_query"),
40+
[
41+
("users", "SELECT * FROM `users`;"),
42+
("user_profiles_2026", "SELECT * FROM `user_profiles_2026`;"),
43+
("schema_name.users", "SELECT * FROM `schema_name`.`users`;"),
44+
("information_schema.tables", "SELECT * FROM `information_schema`.`tables`;"),
45+
],
46+
)
47+
def test_mysql_search_tool_quotes_valid_table_identifiers(
48+
mock_rag_client: MagicMock, table_name: str, expected_query: str
49+
) -> None:
50+
with patch.object(RagTool, "add", return_value=None) as mock_add:
51+
create_mysql_search_tool(mock_rag_client, table_name)
52+
53+
mock_add.assert_called_once_with(
54+
expected_query,
55+
data_type=DataType.MYSQL,
56+
metadata={"db_uri": "mysql://user:password@localhost:3306/test_database"},
57+
)
58+
59+
60+
@pytest.mark.parametrize(
61+
"table_name",
62+
[
63+
"users where 1=1",
64+
"users; drop table users;--",
65+
"users -- comment",
66+
"users/*comment*/",
67+
"`users`",
68+
"schema.users.extra",
69+
"schema.",
70+
".users",
71+
"123users",
72+
],
73+
)
74+
def test_mysql_search_tool_rejects_invalid_table_identifiers(
75+
mock_rag_client: MagicMock, table_name: str
76+
) -> None:
77+
with (
78+
patch.object(RagTool, "add", return_value=None) as mock_add,
79+
pytest.raises(ValueError, match="MySQL table_name must be a valid"),
80+
):
81+
create_mysql_search_tool(mock_rag_client, table_name)
82+
83+
mock_add.assert_not_called()
84+
85+
86+
def test_mysql_search_tool_still_runs_search_queries(
87+
mock_rag_client: MagicMock,
88+
) -> None:
89+
with patch.object(RagTool, "add", return_value=None):
90+
tool = create_mysql_search_tool(mock_rag_client, "users")
91+
92+
with patch.object(RagTool, "_run", return_value="Alice") as mock_run:
93+
result = tool._run("alice")
94+
95+
assert "Alice" in result
96+
mock_run.assert_called_once_with(
97+
query="alice", similarity_threshold=None, limit=None
98+
)
99+
100+
101+
def test_mysql_search_tool_uses_mysql_data_type_metadata(
102+
mock_rag_client: MagicMock,
103+
) -> None:
104+
with patch.object(RagTool, "add", return_value=None) as mock_add:
105+
create_mysql_search_tool(mock_rag_client, "users")
106+
107+
assert mock_add.call_args.kwargs == {
108+
"data_type": DataType.MYSQL,
109+
"metadata": {"db_uri": "mysql://user:password@localhost:3306/test_database"},
110+
}

0 commit comments

Comments
 (0)