FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
jigsawstack-python/tests/test_sql.py at refs/heads/main · InterfazeAI/jigsawstack-python · GitHub
Uh oh!
There was an error while loading.
Please reload this page
.
InterfazeAI
/
jigsawstack-python
Public
Notifications
You must be signed in to change notification settings
Fork
4
Star
20
Code
Issues
0
Pull requests
3
Actions
Projects
Security and quality
0
Insights
Additional navigation options
Code
Issues
Pull requests
Actions
Projects
Security and quality
Insights
Expand file tree
Breadcrumbs
jigsawstack-python
/
tests
/
test_sql.py
Copy path
More file actions
More file actions
Latest commit
History
History
History
282 lines (256 loc) · 7.85 KB
Breadcrumbs
jigsawstack-python
/
tests
/
test_sql.py
Copy path
File metadata and controls
282 lines (256 loc) · 7.85 KB
Raw
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
import
logging
import
os
import
pytest
from
dotenv
import
load_dotenv
import
jigsawstack
from
jigsawstack
.
exceptions
import
JigsawStackError
load_dotenv
()
logging
.
basicConfig
(
level
=
logging
.
INFO
)
logger
=
logging
.
getLogger
(
__name__
)
jigsaw
=
jigsawstack
.
JigsawStack
(
api_key
=
os
.
getenv
(
"JIGSAWSTACK_API_KEY"
),
base_url
=
os
.
getenv
(
"JIGSAWSTACK_BASE_URL"
)
+
"/api"
if
os
.
getenv
(
"JIGSAWSTACK_BASE_URL"
)
else
"https://api.jigsawstack.com"
,
headers
=
{
"x-jigsaw-skip-cache"
:
"true"
},
)
async_jigsaw
=
jigsawstack
.
AsyncJigsawStack
(
api_key
=
os
.
getenv
(
"JIGSAWSTACK_API_KEY"
),
base_url
=
os
.
getenv
(
"JIGSAWSTACK_BASE_URL"
)
+
"/api"
if
os
.
getenv
(
"JIGSAWSTACK_BASE_URL"
)
else
"https://api.jigsawstack.com"
,
headers
=
{
"x-jigsaw-skip-cache"
:
"true"
},
)
# Sample schemas for different databases
MYSQL_SCHEMA
=
"""
CREATE TABLE users (
id INT PRIMARY KEY AUTO_INCREMENT,
username VARCHAR(255) NOT NULL,
email VARCHAR(255) UNIQUE NOT NULL,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE orders (
id INT PRIMARY KEY AUTO_INCREMENT,
user_id INT,
product_name VARCHAR(255),
quantity INT,
price DECIMAL(10, 2),
order_date DATE,
FOREIGN KEY (user_id) REFERENCES users(id)
);
"""
POSTGRESQL_SCHEMA
=
"""
CREATE TABLE employees (
id SERIAL PRIMARY KEY,
name VARCHAR(100) NOT NULL,
department VARCHAR(50),
salary NUMERIC(10, 2),
hire_date DATE,
is_active BOOLEAN DEFAULT true
);
CREATE TABLE departments (
id SERIAL PRIMARY KEY,
name VARCHAR(50) UNIQUE NOT NULL,
budget NUMERIC(12, 2),
manager_id INTEGER REFERENCES employees(id)
);
"""
SQLITE_SCHEMA
=
"""
CREATE TABLE products (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL,
category TEXT,
price REAL,
stock_quantity INTEGER DEFAULT 0
);
CREATE TABLE sales (
id INTEGER PRIMARY KEY AUTOINCREMENT,
product_id INTEGER,
quantity INTEGER,
sale_date TEXT,
total_amount REAL,
FOREIGN KEY (product_id) REFERENCES products(id)
);
"""
TEST_CASES
=
[
{
"name"
:
"mysql_simple_select"
,
"params"
: {
"prompt"
:
"Get all users from the users table"
,
"sql_schema"
:
MYSQL_SCHEMA
,
"database"
:
"mysql"
,
},
},
{
"name"
:
"mysql_join_query"
,
"params"
: {
"prompt"
:
"Get all orders with user information for orders placed in the last 30 days"
,
"sql_schema"
:
MYSQL_SCHEMA
,
"database"
:
"mysql"
,
},
},
{
"name"
:
"mysql_aggregate_query"
,
"params"
: {
"prompt"
:
"Calculate the total revenue per user"
,
"sql_schema"
:
MYSQL_SCHEMA
,
"database"
:
"mysql"
,
},
},
{
"name"
:
"postgresql_simple_select"
,
"params"
: {
"prompt"
:
"Find all active employees"
,
"sql_schema"
:
POSTGRESQL_SCHEMA
,
"database"
:
"postgresql"
,
},
},
{
"name"
:
"postgresql_complex_join"
,
"params"
: {
"prompt"
:
"Get all departments with their manager names and department budgets greater than 100000"
,
"sql_schema"
:
POSTGRESQL_SCHEMA
,
"database"
:
"postgresql"
,
},
},
{
"name"
:
"postgresql_window_function"
,
"params"
: {
"prompt"
:
"Rank employees by salary within each department"
,
"sql_schema"
:
POSTGRESQL_SCHEMA
,
"database"
:
"postgresql"
,
},
},
{
"name"
:
"sqlite_simple_query"
,
"params"
: {
"prompt"
:
"List all products in the electronics category"
,
"sql_schema"
:
SQLITE_SCHEMA
,
"database"
:
"sqlite"
,
},
},
{
"name"
:
"sqlite_aggregate_with_group"
,
"params"
: {
"prompt"
:
"Calculate total sales amount for each product"
,
"sql_schema"
:
SQLITE_SCHEMA
,
"database"
:
"sqlite"
,
},
},
{
"name"
:
"default_database_type"
,
"params"
: {
"prompt"
:
"Select all records from users table where email contains 'example.com'"
,
"sql_schema"
:
MYSQL_SCHEMA
,
# No database specified, should use default
},
},
{
"name"
:
"complex_multi_table_query"
,
"params"
: {
"prompt"
:
"Find users who have placed more than 5 orders with total value exceeding 1000"
,
"sql_schema"
:
MYSQL_SCHEMA
,
"database"
:
"mysql"
,
},
},
{
"name"
:
"insert_query"
,
"params"
: {
"prompt"
:
"Insert a new user with username 'john_doe' and email 'john@example.com'"
,
"sql_schema"
:
MYSQL_SCHEMA
,
"database"
:
"mysql"
,
},
},
{
"name"
:
"update_query"
,
"params"
: {
"prompt"
:
"Update the salary of all employees in the IT department by 10%"
,
"sql_schema"
:
POSTGRESQL_SCHEMA
,
"database"
:
"postgresql"
,
},
},
{
"name"
:
"delete_query"
,
"params"
: {
"prompt"
:
"Delete all products with zero stock quantity"
,
"sql_schema"
:
SQLITE_SCHEMA
,
"database"
:
"sqlite"
,
},
},
{
"name"
:
"subquery_example"
,
"params"
: {
"prompt"
:
"Find all users who have never placed an order"
,
"sql_schema"
:
MYSQL_SCHEMA
,
"database"
:
"mysql"
,
},
},
{
"name"
:
"date_filtering"
,
"params"
: {
"prompt"
:
"Get all employees hired in the last year"
,
"sql_schema"
:
POSTGRESQL_SCHEMA
,
"database"
:
"postgresql"
,
},
},
]
class
TestSQLSync
:
"""Test synchronous SQL text-to-sql methods"""
sync_test_cases
=
TEST_CASES
@
pytest
.
mark
.
parametrize
(
"test_case"
,
sync_test_cases
,
ids
=
[
tc
[
"name"
]
for
tc
in
sync_test_cases
]
)
def
test_text_to_sql
(
self
,
test_case
):
"""Test synchronous text-to-sql with various inputs"""
try
:
result
=
jigsaw
.
text_to_sql
(
test_case
[
"params"
])
assert
result
[
"success"
]
assert
"sql"
in
result
assert
isinstance
(
result
[
"sql"
],
str
)
assert
len
(
result
[
"sql"
])
>
0
# Basic SQL validation - check if it contains SQL keywords
sql_lower
=
result
[
"sql"
].
lower
()
sql_keywords
=
[
"select"
,
"insert"
,
"update"
,
"delete"
,
"create"
,
"alter"
,
"drop"
,
]
assert
any
(
keyword
in
sql_lower
for
keyword
in
sql_keywords
), (
"Generated SQL should contain valid SQL keywords"
)
except
JigsawStackError
as
e
:
pytest
.
fail
(
f"Unexpected JigsawStackError in
{
test_case
[
'name'
]
}
:
{
e
}
"
)
class
TestSQLAsync
:
"""Test asynchronous SQL text-to-sql methods"""
async_test_cases
=
TEST_CASES
@
pytest
.
mark
.
parametrize
(
"test_case"
,
async_test_cases
,
ids
=
[
tc
[
"name"
]
for
tc
in
async_test_cases
]
)
@
pytest
.
mark
.
asyncio
async
def
test_text_to_sql_async
(
self
,
test_case
):
"""Test asynchronous text-to-sql with various inputs"""
try
:
result
=
await
async_jigsaw
.
text_to_sql
(
test_case
[
"params"
])
assert
result
[
"success"
]
assert
"sql"
in
result
assert
isinstance
(
result
[
"sql"
],
str
)
assert
len
(
result
[
"sql"
])
>
0
sql_lower
=
result
[
"sql"
].
lower
()
sql_keywords
=
[
"select"
,
"insert"
,
"update"
,
"delete"
,
"create"
,
"alter"
,
"drop"
,
]
assert
any
(
keyword
in
sql_lower
for
keyword
in
sql_keywords
), (
"Generated SQL should contain valid SQL keywords"
)
except
JigsawStackError
as
e
:
pytest
.
fail
(
f"Unexpected JigsawStackError in
{
test_case
[
'name'
]
}
:
{
e
}
"
)
Back
|
FazBrowse Home
|
New Git URL