Skip to content

Commit 39699ed

Browse files
Fix evalbench core robustness (#248)
* fix(core): improve robustness and database support - evalbench.py: Support dict-based dataset_config extraction. - postgres.py: Enable UNIX domain socket connections for local auth. - oneshotorchestrator.py: Add database name mapping/overrides for multi-engine evaluations. - analyzer.py: Fix KeyError and ZeroDivisionError during reporting. * refactor(analyzer): gracefully handle empty results without short-circuiting * refactor(analyzer): restore LLM metrics logging * fix(databases): enhance db registry and driver robustness - databases/__init__.py: Register SpannerDB and MongoDB in factory. - mysql.py: Improve connection pooling and Cloud SQL handling. - sqlite.py: Implement copy-on-write for temporary databases to support file-based datasets. * fix(databases): enhance db registry and driver robustness - databases/__init__.py: Register SpannerDB and MongoDB in factory. - mysql.py: Improve connection pooling and Cloud SQL handling. - sqlite.py: Implement copy-on-write for temporary databases to support file-based datasets. * fix(dataset): correct progress calculation for multi-dialect evaluations * Fix python test failures and pycodestyle warnings - Updated `evalbench/test/mongodb_test.py` to match the expected data format for `insert_data`. - Fixed `batch_execute` in `evalbench/databases/spanner.py` to correctly use `self.database.update_ddl` instead of just executing standard queries, as Spanner snapshots do not support DDL commands. - Updated `evalbench/test/spanner_test.py` to use `batch_execute` for DDL statements (CREATE TABLE, DROP TABLE). - Modified `.pycodestyle` to ignore `W504` (line break after binary operator), which conflicts with `W503` (line break before binary operator). * Fix trailing whitespace and formatting errors flagged by pycodestyle - Used autopep8 and sed to clean up all W291 and W293 errors in the python source. * Fix comment formatting in spanner.py --------- Co-authored-by: IsmailMehdi <IsmailMehdi@gmail.com>
1 parent aec92c9 commit 39699ed

54 files changed

Lines changed: 465 additions & 208 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎evalbench/client/eval_client.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,8 @@ def __init__(self, endpoint: str):
4444

4545
# 4. Composite Credentials
4646
# Combine the SSL channel with the token-based call credentials
47-
composite_creds = grpc.composite_channel_credentials(channel_creds, call_creds)
47+
composite_creds = grpc.composite_channel_credentials(
48+
channel_creds, call_creds)
4849
self.channel = grpc.aio.secure_channel(address, composite_creds)
4950

5051
self.stub = eval_service_pb2_grpc.EvalServiceStub(self.channel)

‎evalbench/databases/__init__.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,8 @@ def get_database(db_config, db_name) -> DB:
1616
# - It will override the provided default database_name
1717
# - This is useful as the default db may be "postgres" or a default only used for setup
1818
if db_name:
19-
db_config["database_name"] = db_name
19+
suffix = db_config.get("db_name_suffix", "")
20+
db_config["database_name"] = f"{db_name}{suffix}"
2021

2122
if db_config["db_type"] == "postgres":
2223
return PGDB(db_config)

‎evalbench/databases/bigquery.py‎

Lines changed: 32 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -68,30 +68,37 @@ def _run_execute(query: str, eval_query: Optional[str] = None, rollback=False):
6868
error = None
6969
query_replaced = query.replace("{{dataset}}", self.db_name)
7070
if eval_query is not None:
71-
eval_query_replaced = eval_query.replace("{{dataset}}", self.db_name)
71+
eval_query_replaced = eval_query.replace(
72+
"{{dataset}}", self.db_name)
7273
try:
7374
if rollback:
7475
try:
7576
initial_query = "SELECT 1;"
7677
job_config = QueryJobConfig(create_session=True)
77-
init_job = self.client.query(initial_query, job_config=job_config)
78+
init_job = self.client.query(
79+
initial_query, job_config=job_config)
7880
init_job.result()
7981
session_id = init_job.session_info.session_id
80-
conn_props = [ConnectionProperty(key="session_id", value=session_id)]
82+
conn_props = [ConnectionProperty(
83+
key="session_id", value=session_id)]
8184

8285
self.client.query(
8386
"BEGIN TRANSACTION;",
84-
job_config=QueryJobConfig(connection_properties=conn_props)
87+
job_config=QueryJobConfig(
88+
connection_properties=conn_props)
8589
).result()
8690

87-
result = self._execute_queries(query_replaced, job_config=QueryJobConfig(connection_properties=conn_props))
91+
result = self._execute_queries(
92+
query_replaced, job_config=QueryJobConfig(connection_properties=conn_props))
8893

8994
if eval_query:
90-
eval_result = self._execute_queries(eval_query_replaced, job_config=QueryJobConfig(connection_properties=conn_props))
95+
eval_result = self._execute_queries(
96+
eval_query_replaced, job_config=QueryJobConfig(connection_properties=conn_props))
9197

9298
self.client.query(
9399
"ROLLBACK TRANSACTION;",
94-
job_config=QueryJobConfig(connection_properties=conn_props)
100+
job_config=QueryJobConfig(
101+
connection_properties=conn_props)
95102
).result()
96103

97104
except Exception as e:
@@ -102,7 +109,8 @@ def _run_execute(query: str, eval_query: Optional[str] = None, rollback=False):
102109
if 'session_id' in locals():
103110
self.client.query(
104111
"CALL BQ.ABORT_SESSION();",
105-
job_config=QueryJobConfig(connection_properties=conn_props)
112+
job_config=QueryJobConfig(
113+
connection_properties=conn_props)
106114
).result()
107115
if not rollback:
108116
result = self._execute_queries(query_replaced)
@@ -113,9 +121,11 @@ def _run_execute(query: str, eval_query: Optional[str] = None, rollback=False):
113121
except (GoogleAPICallError, Exception) as e:
114122
error = str(e)
115123
if "resources exceeded" in error:
116-
raise ResourceExhaustedError(f"BigQuery resources exhausted: {e}") from e
124+
raise ResourceExhaustedError(
125+
f"BigQuery resources exhausted: {e}") from e
117126
elif "quota exceeded" in error:
118-
raise ResourceExhaustedError(f"BigQuery quota exceeded: {e}") from e
127+
raise ResourceExhaustedError(
128+
f"BigQuery quota exceeded: {e}") from e
119129
else:
120130
print(error)
121131

@@ -140,9 +150,11 @@ def get_metadata(self) -> dict:
140150
try:
141151
for table in self.client.list_tables(self.db_name):
142152
schema = self.client.get_table(table.reference).schema
143-
metadata[table.table_id] = [{"name": f.name, "type": f.field_type} for f in schema]
153+
metadata[table.table_id] = [
154+
{"name": f.name, "type": f.field_type} for f in schema]
144155
except Exception as e:
145-
print(f"Error while fetching metadata for dataset '{self.db_name}': {e}")
156+
print(
157+
f"Error while fetching metadata for dataset '{self.db_name}': {e}")
146158
return metadata
147159

148160
#####################################################
@@ -155,7 +167,8 @@ def generate_ddl(self, schema: DatabaseSchema) -> List[str]:
155167
ddl_statements = []
156168
try:
157169
for table in schema.tables:
158-
columns = ", ".join([f"{col.name} {col.type}" for col in table.columns])
170+
columns = ", ".join(
171+
[f"{col.name} {col.type}" for col in table.columns])
159172
ddl_statements.append(
160173
f"CREATE TABLE `{self.project_id}.{self.db_name}.{table.name}` ({columns})"
161174
)
@@ -191,7 +204,8 @@ def drop_all_tables(self):
191204
self.client.delete_table(full_table_id)
192205

193206
except Exception as e:
194-
raise RuntimeError(f"Failed to drop tables in dataset {self.db_name}: {e}")
207+
raise RuntimeError(
208+
f"Failed to drop tables in dataset {self.db_name}: {e}")
195209

196210
def _is_float(self, value) -> bool:
197211
try:
@@ -204,11 +218,13 @@ def _get_column_name_to_type_mapping(self, sql_statements: List[str]) -> Dict[st
204218
schema_mapping = {}
205219

206220
for statement in sql_statements:
207-
table_match = re.search(r'CREATE TABLE\s+`{{dataset}}\.(\w+)`', statement)
221+
table_match = re.search(
222+
r'CREATE TABLE\s+`{{dataset}}\.(\w+)`', statement)
208223
if not table_match:
209224
continue
210225
table_name = table_match.group(1)
211-
column_section_match = re.search(r'\(\n(.*?)\n\)', statement, re.DOTALL)
226+
column_section_match = re.search(
227+
r'\(\n(.*?)\n\)', statement, re.DOTALL)
212228
if not column_section_match:
213229
continue
214230
columns_raw = column_section_match.group(1).split(",\n")

‎evalbench/databases/bigtable.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -114,7 +114,8 @@ def get_metadata(self) -> dict:
114114
{"name": cf, "type": COLUMN_FAMILY_TYPE} for cf in column_families
115115
]
116116
except Exception:
117-
logging.error(f"Failed to get metadata for table {table.table_id}")
117+
logging.error(
118+
f"Failed to get metadata for table {table.table_id}")
118119
return db_metadata
119120

120121
def generate_ddl(

‎evalbench/databases/db.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -91,7 +91,8 @@ def setup_tmp_users(self):
9191
self.dql_user = "tmp_dql_user_" + generate_key()
9292
self.dml_user = "tmp_dml_user_" + generate_key()
9393
self.tmp_user_password = generate_key()
94-
self.create_tmp_users(self.dql_user, self.dml_user, self.tmp_user_password)
94+
self.create_tmp_users(self.dql_user, self.dml_user,
95+
self.tmp_user_password)
9596
self.tmp_users.extend([self.dql_user, self.dml_user])
9697

9798
def delete_tmp_users(self, users) -> None:

‎evalbench/databases/mongodb.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -149,7 +149,8 @@ def generate_ddl(
149149
for table in schema.tables:
150150
ddl.append(f"Collection: {table.name}")
151151
if table.columns:
152-
col_descs = [f"{col.name} ({col.type})" for col in table.columns]
152+
col_descs = [
153+
f"{col.name} ({col.type})" for col in table.columns]
153154
ddl.append(f" Fields: {', '.join(col_descs)}")
154155
return ddl
155156

‎evalbench/databases/mysql.py‎

Lines changed: 71 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
import sqlparse
33
from sqlalchemy import text, MetaData
44
from sqlalchemy.engine.base import Connection
5+
import pymysql
56
import logging
67
from .db import DB
78
from google.cloud.sql.connector import Connector
@@ -47,37 +48,81 @@ class MySQLDB(DB):
4748

4849
def __init__(self, db_config):
4950
super().__init__(db_config)
50-
self.connector = Connector()
51+
52+
# Auto-deduce use_cloud_sql: format is PROJECT:REGION:INSTANCE (2 colons)
53+
self.use_cloud_sql = db_config.get("use_cloud_sql")
54+
if self.use_cloud_sql is None:
55+
self.use_cloud_sql = (self.db_path.count(":") == 2)
56+
57+
self.connector = Connector() if self.use_cloud_sql else None
5158

5259
def get_conn():
53-
conn = self.connector.connect(
54-
self.db_path,
55-
"pymysql",
56-
user=self.username,
57-
password=self.password,
58-
db=self.db_name,
59-
)
60-
return conn
60+
"""Callable for sqlalchemy 'creator' parameter."""
61+
if self.use_cloud_sql:
62+
return self.connector.connect(
63+
self.db_path,
64+
"pymysql",
65+
user=self.username,
66+
password=self.password,
67+
db=self.db_name,
68+
)
69+
else:
70+
# Local/Direct connection
71+
host = self.db_path
72+
port = 3306
73+
if ":" in self.db_path:
74+
parts = self.db_path.split(":")
75+
host = parts[0]
76+
port = int(parts[1])
77+
78+
return pymysql.connect(
79+
host=host,
80+
port=port,
81+
user=self.username,
82+
password=self.password or "",
83+
database=self.db_name
84+
)
6185

62-
def get_engine_args():
63-
common_args = {
64-
"creator": get_conn,
65-
"connect_args": {"command_timeout": 60, "multi_statements": True},
86+
def get_engine_config():
87+
"""Returns (db_url, engine_args)"""
88+
args = {
89+
"connect_args": {},
6690
}
91+
url = ""
92+
93+
if self.use_cloud_sql:
94+
args["creator"] = get_conn
95+
# Cloud SQL needs explicit command_timeout and multi_statements
96+
args["connect_args"]["command_timeout"] = 60
97+
args["connect_args"]["multi_statements"] = True
98+
url = "mysql+pymysql://"
99+
else:
100+
# Standard local connection via URL
101+
# SQLAlchemy parses this URL and loads the driver internally
102+
args["connect_args"]["connect_timeout"] = 60
103+
104+
password_part = f":{self.password}" if self.password else ""
105+
url = f"mysql+pymysql://{self.username}{password_part}@{self.db_path}/{self.db_name}"
106+
107+
password_part = f":{self.password}" if self.password else ""
108+
url = f"mysql+pymysql://{self.username}{password_part}@{self.db_path}/{self.db_name}"
109+
67110
if "is_tmp_db" in db_config:
68-
common_args["pool_size"] = 1
69-
common_args["pool_recycle"] = 300
111+
args["pool_size"] = 1
112+
args["pool_recycle"] = 300
70113
else:
71-
common_args["pool_size"] = 50
72-
common_args["pool_recycle"] = 300
73-
return common_args
114+
args["pool_size"] = 50
115+
args["pool_recycle"] = 300
116+
return url, args
74117

75-
self.engine = sqlalchemy.create_engine("mysql+pymysql://", **get_engine_args())
118+
db_url, engine_args = get_engine_config()
119+
self.engine = sqlalchemy.create_engine(db_url, **engine_args)
76120

77121
def close_connections(self):
78122
try:
79123
self.engine.dispose()
80-
self.connector.close()
124+
if self.connector:
125+
self.connector.close()
81126
except Exception:
82127
logging.warning(
83128
f"Failed to close connections. This may result in idle unused connections."
@@ -143,7 +188,8 @@ def _run_execute(query: str, eval_query: Optional[str] = None, rollback=False):
143188
result = self._execute_queries(connection, query)
144189

145190
if eval_query:
146-
eval_result = self._execute_queries(connection, eval_query)
191+
eval_result = self._execute_queries(
192+
connection, eval_query)
147193

148194
if batch_commands and len(batch_commands) > 0:
149195
for command in batch_commands:
@@ -181,7 +227,8 @@ def get_metadata(self) -> dict:
181227
for table in metadata.tables.values():
182228
columns = []
183229
for column in table.columns:
184-
columns.append({"name": column.name, "type": str(column.type)})
230+
columns.append(
231+
{"name": column.name, "type": str(column.type)})
185232
db_metadata[table.name] = columns
186233
except Exception:
187234
pass
@@ -203,7 +250,8 @@ def generate_ddl(
203250
columns = ", ".join(
204251
[f"{column.name} {column.type}" for column in table.columns]
205252
)
206-
create_statements.append(f"CREATE TABLE `{table.name}` ({columns});")
253+
create_statements.append(
254+
f"CREATE TABLE `{table.name}` ({columns});")
207255
return create_statements
208256

209257
def create_tmp_database(self, database_name: str):

0 commit comments

Comments
 (0)