-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_h2odb.py
More file actions
451 lines (355 loc) · 19.5 KB
/
Copy pathtest_h2odb.py
File metadata and controls
451 lines (355 loc) · 19.5 KB
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
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
"""Tests for H2O Database Benchmark implementation.
Copyright 2026 Joe Harris / BenchBox Project
Licensed under the MIT License. See LICENSE file in the project root for details.
"""
import sqlite3
import tempfile
from pathlib import Path
from unittest.mock import Mock, patch
import pytest
from benchbox import H2ODB
from benchbox.core.h2odb.benchmark import H2OBenchmark as H2ODBBenchmark
from benchbox.utils.path_utils import resolve_benchmark_runs_dir
from .fixtures.benchmark_test_mixin import BenchmarkTestMixin
pytestmark = pytest.mark.medium
@pytest.mark.h2odb
class TestH2ODB:
"""Test the H2O Database Benchmark implementation."""
@pytest.fixture
def h2odb(self, small_scale_factor: float, temp_dir: Path) -> H2ODB:
"""Create an H2ODB H2ODBBenchmark instance for testing."""
return H2ODB(scale_factor=small_scale_factor, output_dir=temp_dir)
# Runs to completion but far exceeds the medium lane's 60s cap: timed out
# on a clean GitHub runner (develop-post-merge run 28706929881, medium-test
# job) and measured ~257s in a 4-core container (2026-07-05). The marker
# overrides the CLI --timeout=60 (pytest-timeout marker precedence) so the
# test keeps running in the medium lane instead of being killed mid-work.
@pytest.mark.timeout(600)
def test_generate_data(self, h2odb: H2ODB) -> None:
"""Test that data generation produces expected files."""
data_paths = h2odb.generate_data()
# Check that output includes expected files (returns dict)
expected_tables = ["trips"] # H2ODB has one main table
assert isinstance(data_paths, dict), "generate_data should return a dictionary"
for table in expected_tables:
assert table in data_paths, f"Table {table} not found in generated data"
# Verify files exist
for _table_name, path in data_paths.items():
assert Path(path).exists(), f"Generated file {path} does not exist"
def test_get_queries(self, h2odb: H2ODB) -> None:
"""Test that all H2ODBBenchmark queries can be retrieved."""
queries = h2odb.get_queries()
# H2ODB has 10 analytical queries
expected_queries = ["Q1", "Q2", "Q3", "Q4", "Q5", "Q6", "Q7", "Q8", "Q9", "Q10"]
assert len(queries) == 10
for query_id in expected_queries:
assert query_id in queries, f"Query {query_id} not found"
assert isinstance(queries[query_id], str)
assert queries[query_id].strip()
# Check that queries contain expected SQL elements
query_sql = queries[query_id].upper()
assert "SELECT" in query_sql
assert "FROM" in query_sql
assert "TRIPS" in query_sql # All queries should use the trips table
# Check for analytical functions that are common in H2O H2ODBBenchmarks
if "GROUP BY" in query_sql or "COUNT" in query_sql or "SUM" in query_sql or "AVG" in query_sql:
# This is an analytical query, which is expected for H2O H2ODBBenchmark
pass
def test_get_query(self, h2odb: H2ODB) -> None:
"""Test retrieving a specific query."""
query1 = h2odb.get_query("Q1")
assert isinstance(query1, str)
assert "SELECT" in query1.upper()
assert "TRIPS" in query1.upper()
# Test a different query
query5 = h2odb.get_query("Q5")
assert isinstance(query5, str)
assert "SELECT" in query5.upper()
assert "TRIPS" in query5.upper()
def test_translate_query(self, h2odb: H2ODB, sql_dialect: str) -> None:
"""Test translating a query to different SQL dialects."""
h2odb.get_query("Q1")
translated_query = h2odb.translate_query("Q1", dialect=sql_dialect)
assert isinstance(translated_query, str)
assert "SELECT" in translated_query.upper()
def test_invalid_query_id(self, h2odb: H2ODB) -> None:
"""Test that requesting an invalid query raises an exception."""
with pytest.raises(ValueError):
h2odb.get_query("Q11") # Only Q1-Q10 exist
with pytest.raises(ValueError):
h2odb.get_query("Q0") # Q0 doesn't exist
def test_get_query_params(self, h2odb: H2ODB) -> None:
"""Test that H2O queries don't accept parameters."""
# Test with default (no parameters)
param_query = h2odb.get_query("Q1")
assert isinstance(param_query, str)
assert "SELECT" in param_query.upper()
# Test that passing parameters raises an error (H2O queries are static)
with pytest.raises(ValueError, match="H2O DB queries are static and don't accept parameters"):
h2odb.get_query("Q1", params={"param1": "value1"})
def test_get_schema(self, h2odb: H2ODB) -> None:
"""Test retrieving the H2ODB schema."""
schema = h2odb.get_schema()
# Check that schema is a dictionary and trips table is present
assert isinstance(schema, dict), "Schema should be a dictionary"
assert "trips" in schema, "trips table not found in schema"
# Check trips table structure
trips_table = schema["trips"]
column_names = [col["name"] for col in trips_table["columns"]]
# Expected columns based on NYC taxi data structure (using actual names from implementation)
expected_columns = [
"vendor_id",
"pickup_datetime",
"dropoff_datetime",
"passenger_count",
"trip_distance",
"pickup_longitude",
"pickup_latitude",
"dropoff_longitude",
"dropoff_latitude",
"payment_type",
"fare_amount",
"extra",
"mta_tax",
"tip_amount",
"tolls_amount",
"total_amount",
]
for column in expected_columns:
assert column in column_names, f"Column {column} not found in trips table"
def test_get_create_tables_sql(self, h2odb: H2ODB) -> None:
"""Test retrieving SQL to create H2ODB tables."""
sql = h2odb.get_create_tables_sql()
assert isinstance(sql, str)
assert "CREATE TABLE" in sql
assert "trips" in sql
# Check for expected column types (using actual column names)
assert "vendor_id" in sql
assert "pickup_datetime" in sql
assert "fare_amount" in sql
def test_H2ODBBenchmark_properties(self, h2odb: H2ODB) -> None:
"""Test H2ODBBenchmark-specific properties."""
# Test that H2ODB has exactly one table (the trips table)
schema = h2odb.get_schema()
assert len(schema) == 1
# Test that all queries are analytical in nature
queries = h2odb.get_queries()
assert len(queries) == 10
# Check that queries test different analytical aspects
all_queries_text = " ".join(queries.values()).upper()
# Should contain various analytical operations
analytical_keywords = ["COUNT", "SUM", "AVG", "GROUP BY", "ORDER BY"]
found_keywords = [kw for kw in analytical_keywords if kw in all_queries_text]
assert len(found_keywords) >= 3, "H2O queries should contain various analytical operations"
def test_taxi_data_focus(self, h2odb: H2ODB) -> None:
"""Test that the H2ODBBenchmark focuses on taxi/trip data analysis."""
schema = h2odb.get_schema()
trips_table = schema["trips"]
column_names = [col["name"] for col in trips_table["columns"]]
# Should have taxi-specific columns
taxi_columns = [
"pickup_datetime",
"dropoff_datetime",
"trip_distance",
"fare_amount",
"tip_amount",
"passenger_count",
]
for col in taxi_columns:
assert col in column_names, f"Taxi-specific column {col} not found"
# Queries should reference taxi-related concepts
queries = h2odb.get_queries()
all_queries_text = " ".join(queries.values()).upper()
# Should contain references to taxi operations
taxi_keywords = ["TRIP", "FARE", "TIP", "PASSENGER", "PICKUP", "DROPOFF"]
found_keywords = [kw for kw in taxi_keywords if kw in all_queries_text]
assert len(found_keywords) >= 2, "H2O queries should reference taxi/trip concepts"
def test_analytical_query_patterns(self, h2odb: H2ODB) -> None:
"""Test that queries follow expected analytical patterns."""
queries = h2odb.get_queries()
# Should have queries that test different analytical patterns
has_aggregation = False
has_grouping = False
for _query_id, query_sql in queries.items():
query_upper = query_sql.upper()
if any(agg in query_upper for agg in ["COUNT", "SUM", "AVG", "MIN", "MAX"]):
has_aggregation = True
if "GROUP BY" in query_upper:
has_grouping = True
if "ORDER BY" in query_upper:
pass
if "WHERE" in query_upper:
pass
assert has_aggregation, "H2O H2ODBBenchmark should include aggregation queries"
assert has_grouping, "H2O H2ODBBenchmark should include grouping queries"
# Note: not all analytical queries need sorting and filtering, so we don't assert those
@pytest.mark.h2odb
class TestH2ODBBenchmarkDirectly(BenchmarkTestMixin):
"""Test H2ODBBenchmark class directly for better coverage."""
benchmark_class = H2ODBBenchmark
sample_query_id = "Q1"
sample_table = "trips"
sample_sql = "SELECT COUNT(*) FROM trips"
sample_csv_filename = "trips.csv"
sample_csv_content = "2023-01-01 12:00:00|2023-01-01 12:30:00|5.5|15.50|3.00|2|40.7589|-73.9851|40.7614|-73.9776\n"
get_query_passes_params = False
@pytest.fixture
def h2odb_benchmark(self, small_scale_factor: float, temp_dir: Path) -> H2ODBBenchmark:
"""Create an H2ODBBenchmark instance for testing."""
return H2ODBBenchmark(scale_factor=small_scale_factor, output_dir=temp_dir)
@pytest.fixture
def benchmark_instance(self, h2odb_benchmark: H2ODBBenchmark) -> H2ODBBenchmark:
"""Alias for the mixin's benchmark_instance fixture."""
return h2odb_benchmark
def test_init_with_default_output_dir(self) -> None:
"""Test H2ODBBenchmark initialization with default output directory."""
h2odb = H2ODBBenchmark(scale_factor=0.01)
assert h2odb.scale_factor == 0.01
# Default path follows: benchmark_runs/datagen/{benchmark}_sf{formatted_sf}
# 0.01 formats as "sf001" per format_scale_factor()
assert h2odb.output_dir == resolve_benchmark_runs_dir() / "datagen" / "h2odb_sf001"
assert h2odb._name == "H2O Database Benchmark"
assert h2odb._version == "1.0"
def test_generate_data_validation(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test data generation input validation."""
# Test unsupported output format
with pytest.raises(ValueError, match="Unsupported output format"):
h2odb_benchmark.generate_data(output_format="json")
# Test invalid table names
with pytest.raises(ValueError, match="Invalid table names"):
h2odb_benchmark.generate_data(tables=["invalid_table"])
def test_load_data_to_database_batch_processing(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test load_data_to_database with batch processing."""
# Set up mock data with many rows to test batching
with tempfile.TemporaryDirectory() as temp_dir:
csv_file = Path(temp_dir) / "trips.csv"
# a large CSV file to test batch processing
rows = []
for i in range(15000): # More than batch_size (10000)
rows.append(
f"2023-01-01 12:00:00|2023-01-01 12:30:00|{i}.5|15.50|3.00|2|40.7589|-73.9851|40.7614|-73.9776"
)
csv_file.write_text("\n".join(rows) + "\n")
h2odb_benchmark.tables = {"trips": str(csv_file)}
mock_connection = Mock()
mock_connection.executescript = Mock()
mock_connection.executemany = Mock()
mock_connection.commit = Mock()
with patch.object(H2ODBBenchmark, "get_create_tables_sql") as mock_get_sql:
mock_get_sql.return_value = "CREATE TABLE trips (...);"
h2odb_benchmark.load_data_to_database(mock_connection)
# Should be called multiple times due to batching
assert mock_connection.executemany.call_count >= 2
def test_run_H2ODBBenchmark_default_queries(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test run_benchmark with default queries."""
mock_connection = Mock()
with patch.object(h2odb_benchmark, "execute_query") as mock_execute:
mock_execute.return_value = [("result1",), ("result2",)]
with patch.object(h2odb_benchmark.query_manager, "get_all_queries") as mock_get_all:
mock_get_all.return_value = {
"Q1": "SELECT COUNT(*) FROM trips",
"Q2": "SELECT AVG(fare_amount) FROM trips",
}
result = h2odb_benchmark.run_benchmark(mock_connection, iterations=1)
assert result["benchmark"] == "H2O Database Benchmark"
assert result["scale_factor"] == h2odb_benchmark.scale_factor
assert result["iterations"] == 1
assert len(result["queries"]) == 2
assert "Q1" in result["queries"]
assert "Q2" in result["queries"]
def test_run_H2ODBBenchmark_timing_calculation(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test run_ H2ODBBenchmark timing calculations."""
mock_connection = Mock()
with patch.object(H2ODBBenchmark, "execute_query") as mock_execute:
mock_execute.return_value = [("result1",), ("result2",)]
mock_clock = Mock(side_effect=[0.0, 0.5, 1.0, 1.2]) # Two iterations
with (
patch("benchbox.core.h2odb.benchmark.mono_time", mock_clock),
patch("benchbox.utils.clock.mono_time", mock_clock),
):
result = h2odb_benchmark.run_benchmark(mock_connection, queries=["Q1"], iterations=2)
query_result = result["queries"]["Q1"]
assert abs(query_result["avg_time"] - 0.35) < 0.01 # (0.5 + 0.2) / 2
assert abs(query_result["min_time"] - 0.2) < 0.01
assert abs(query_result["max_time"] - 0.5) < 0.01
def test_run_H2ODBBenchmark_with_exceptions(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test run_ H2ODBBenchmark handling exceptions during execution."""
mock_connection = Mock()
with patch.object(H2ODBBenchmark, "execute_query") as mock_execute:
mock_execute.side_effect = Exception("Database error")
result = h2odb_benchmark.run_benchmark(mock_connection, queries=["Q1"], iterations=1)
query_result = result["queries"]["Q1"]
assert query_result["iterations"][0]["success"] is False
assert query_result["iterations"][0]["error"] == "Database error"
assert query_result["avg_time"] == 0
def test_sqlite_integration(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test actual SQLite integration for load_data_to_database."""
# Generate some test data for trips table - all 22 columns
with tempfile.TemporaryDirectory() as temp_dir:
csv_file = Path(temp_dir) / "trips.csv"
# Data format: vendor_id|pickup_datetime|dropoff_datetime|passenger_count|trip_distance|pickup_longitude|pickup_latitude|rate_code_id|store_and_fwd_flag|dropoff_longitude|dropoff_latitude|payment_type|fare_amount|extra|mta_tax|tip_amount|tolls_amount|improvement_surcharge|total_amount|pickup_location_id|dropoff_location_id|congestion_surcharge
csv_file.write_text(
"1|2023-01-01 12:00:00|2023-01-01 12:30:00|2|5.5|-73.9851|40.7589|1|N|-73.9776|40.7614|1|15.50|0.50|0.50|3.00|0.00|0.30|19.80|161|239|2.50\n2|2023-01-01 13:00:00|2023-01-01 13:15:00|1|2.1|-73.9934|40.7505|1|N|-73.9889|40.7490|2|8.50|0.00|0.50|1.50|0.00|0.30|10.80|230|186|2.50\n"
)
h2odb_benchmark.tables = {"trips": str(csv_file)}
# in-memory SQLite database
conn = sqlite3.connect(":memory:")
try:
# This should work without errors
h2odb_benchmark.load_data_to_database(conn, tables=["trips"])
# Verify data was loaded
cursor = conn.cursor()
cursor.execute("SELECT COUNT(*) FROM trips")
count = cursor.fetchone()[0]
assert count == 2
cursor.execute("SELECT pickup_datetime, trip_distance FROM trips ORDER BY pickup_datetime")
rows = cursor.fetchall()
assert len(rows) == 2
assert rows[0][0] == "2023-01-01 12:00:00" # pickup_datetime
assert float(rows[0][1]) == 5.5 # trip_distance
finally:
conn.close()
def test_schema_methods_delegation(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test schema-related methods delegate correctly."""
# Test get_schema
schema = h2odb_benchmark.get_schema()
assert isinstance(schema, dict)
assert "trips" in schema
# Test get_create_tables_sql
sql = h2odb_benchmark.get_create_tables_sql()
assert isinstance(sql, str)
assert "CREATE TABLE" in sql
def test_get_all_queries_method(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test get_all_queries method."""
with patch.object(h2odb_benchmark.query_manager, "get_all_queries") as mock_get_all:
mock_get_all.return_value = {"Q1": "SELECT ...", "Q2": "SELECT ..."}
result = h2odb_benchmark.get_all_queries()
mock_get_all.assert_called_once()
assert result == {"Q1": "SELECT ...", "Q2": "SELECT ..."}
def test_csv_parsing_edge_cases(self, h2odb_benchmark: H2ODBBenchmark) -> None:
"""Test CSV parsing with various edge cases."""
with tempfile.TemporaryDirectory() as temp_dir:
csv_file = Path(temp_dir) / "trips.csv"
# Test with pipe delimiters and typical H2O data format
csv_file.write_text(
"2023-01-01 12:00:00|2023-01-01 12:30:00|5.5|15.50|3.00|2|40.7589|-73.9851|40.7614|-73.9776\n"
)
h2odb_benchmark.tables = {"trips": str(csv_file)}
# mock connection that tracks the data being inserted
mock_connection = Mock()
mock_connection.executescript = Mock()
mock_connection.executemany = Mock()
mock_connection.commit = Mock()
with patch.object(H2ODBBenchmark, "get_create_tables_sql") as mock_get_sql:
mock_get_sql.return_value = "CREATE TABLE trips (...);"
h2odb_benchmark.load_data_to_database(mock_connection)
# Verify executemany was called with the right data
call_args = mock_connection.executemany.call_args
assert call_args is not None
sql, data = call_args[0]
assert "INSERT INTO trips VALUES" in sql
assert len(data) == 1
# Check that the datetime and numeric fields are parsed correctly
assert data[0][0] == "2023-01-01 12:00:00" # pickup_datetime
assert data[0][1] == "2023-01-01 12:30:00" # dropoff_datetime
assert data[0][2] == "5.5" # trip_distance