Skip to content

Commit 73b32e1

Browse files
Merge pull request #838 from pyathena-dev/feat/835-insertmanyvalues
Batch SQLAlchemy executemany inserts into multi-row INSERT statements
2 parents 2a107ad + 9af65f2 commit 73b32e1

4 files changed

Lines changed: 321 additions & 2 deletions

File tree

‎docs/sqlalchemy.md‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -756,6 +756,34 @@ Athena `FLOAT` and `REAL` are the same 32-bit floating-point type, which keeps a
756756
Use `Double` (SQLAlchemy 2.0+) for 64-bit values.
757757
`Float(precision)` does not change the Athena type.
758758

759+
## Bulk inserts
760+
761+
With SQLAlchemy 2.0, an insert executed with a list of parameter sets runs as multi-row `INSERT INTO ... VALUES (...), (...)` statements of up to 100 rows each, instead of one query per row.
762+
This applies to Core `insert()` with a list of parameters and to ORM flushes that insert several objects.
763+
`CursorResult.rowcount` is the total number of inserted rows, or -1 if Athena does not report a count.
764+
If a statement fails, the rows of the earlier statements remain inserted.
765+
766+
Athena limits a query to 262,144 bytes.
767+
For wide rows, lower the number of rows per statement for the engine or for one execution.
768+
To run one query per row, disable the batching:
769+
770+
```python
771+
from sqlalchemy import create_engine
772+
773+
engine = create_engine(
774+
"awsathena+rest://:@athena.us-west-2.amazonaws.com:443/default?s3_staging_dir=s3://YOUR_S3_BUCKET/path/to/",
775+
insertmanyvalues_page_size=20,
776+
)
777+
778+
with engine.begin() as conn:
779+
conn.execution_options(insertmanyvalues_page_size=5).execute(table.insert(), rows)
780+
781+
engine_per_row = create_engine(
782+
"awsathena+rest://:@athena.us-west-2.amazonaws.com:443/default?s3_staging_dir=s3://YOUR_S3_BUCKET/path/to/",
783+
use_insertmanyvalues=False,
784+
)
785+
```
786+
759787
## Complex data types
760788

761789
### STRUCT type support

‎pyathena/sqlalchemy/base.py‎

Lines changed: 30 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313

1414
from sqlalchemy import exc, schema, types, util
1515
from sqlalchemy.engine import Engine, reflection
16-
from sqlalchemy.engine.default import DefaultDialect
16+
from sqlalchemy.engine.default import DefaultDialect, DefaultExecutionContext
1717
from sqlalchemy.engine.interfaces import ExecutionContext
1818
from sqlalchemy.sql.compiler import (
1919
DDLCompiler,
@@ -153,6 +153,12 @@ class AthenaDialect(DefaultDialect):
153153
supports_default_values: bool = False
154154
supports_empty_insert: bool = False
155155
supports_multivalues_insert: bool = True
156+
# Render executemany inserts as multi-row INSERT statements. Athena has no
157+
# RETURNING, so batching must also apply to inserts without it. The page
158+
# size keeps typical rows well below Athena's 262,144-byte query limit.
159+
use_insertmanyvalues: bool = True
160+
use_insertmanyvalues_wo_returning: bool = True
161+
insertmanyvalues_page_size: int = 100
156162
supports_sane_rowcount: bool = True
157163
supports_sane_multi_rowcount: bool = True
158164
supports_native_decimal: bool = True
@@ -708,6 +714,19 @@ def get_indexes(
708714
return [] # pragma: no cover
709715

710716
def do_execute(self, cursor, statement, parameters, context=None):
717+
"""Execute a statement with the DB API cursor.
718+
719+
SQLAlchemy calls this once per page of an "insertmanyvalues" insert.
720+
For those pages, the execution context's row count accumulates the
721+
cursor row counts, so that ``CursorResult.rowcount`` reports the total
722+
like ``Cursor.executemany``, or -1 if any page has an unknown count.
723+
724+
Args:
725+
cursor: The DB API cursor.
726+
statement: The SQL statement.
727+
parameters: The statement parameters.
728+
context: The SQLAlchemy execution context, if any.
729+
"""
711730
on_start_query_execution = None
712731
if isinstance(context, ExecutionContext):
713732
execution_options = context.execution_options
@@ -719,6 +738,16 @@ def do_execute(self, cursor, statement, parameters, context=None):
719738
else:
720739
cursor.execute(statement, parameters)
721740

741+
# An executemany context reaches do_execute only for insertmanyvalues
742+
# pages; other executemany statements go through do_executemany.
743+
if isinstance(context, DefaultExecutionContext) and context.executemany:
744+
total = context._rowcount
745+
count = cursor.rowcount
746+
if total is None:
747+
context._rowcount = count
748+
else:
749+
context._rowcount = total + count if total >= 0 and count >= 0 else -1
750+
722751
def do_rollback(self, dbapi_connection: PoolProxiedConnection) -> None:
723752
# No transactions for Athena
724753
pass # pragma: no cover

‎tests/pyathena/aio/sqlalchemy/test_base.py‎

Lines changed: 62 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,12 @@
1+
from datetime import datetime
2+
from decimal import Decimal
3+
14
import pytest
25
import sqlalchemy
36
from sqlalchemy import cast, literal, select, text, types
4-
from sqlalchemy.sql.schema import MetaData, Table
7+
from sqlalchemy.sql.schema import Column, MetaData, Table
58

9+
from pyathena.sqlalchemy.types import AthenaArray
610
from tests import ENV
711
from tests.pyathena.util import throttle_metadata_api
812

@@ -229,3 +233,60 @@ async def test_executemany_failure(self, async_engine, executemany_table):
229233
assert rows == [(1, 11), (2, 21), (3, 30)]
230234
result = await conn.execute(statement, {"group_id": "2"})
231235
assert result.rowcount == 1
236+
237+
async def test_insertmanyvalues(self, async_engine):
238+
_, conn = async_engine
239+
table_name = "insertmanyvalues_async"
240+
table = Table(
241+
table_name,
242+
MetaData(schema=ENV.schema),
243+
Column("id", types.Integer),
244+
Column("name", types.String),
245+
Column("data", types.LargeBinary),
246+
Column("ts", types.DateTime),
247+
Column("amount", types.Numeric(10, 3)),
248+
Column("tags", AthenaArray(types.Integer)),
249+
awsathena_location=f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/",
250+
awsathena_tblproperties={"table_type": "ICEBERG"},
251+
)
252+
rows = [
253+
{
254+
"id": 1,
255+
"name": "it's",
256+
"data": b"\x00\x01",
257+
"ts": datetime(2026, 1, 2, 3, 4, 5, 123000),
258+
"amount": Decimal("1.5"),
259+
"tags": [1, None],
260+
},
261+
{
262+
"id": 2,
263+
"name": None,
264+
"data": None,
265+
"ts": datetime(2026, 1, 2, 3, 4, 5, 123456),
266+
"amount": Decimal("12.345"),
267+
"tags": None,
268+
},
269+
{"id": 3, "name": "c", "data": b"", "ts": None, "amount": None, "tags": []},
270+
{"id": 4, "name": "d", "data": b"d", "ts": None, "amount": None, "tags": [4]},
271+
{"id": 5, "name": "e", "data": b"e", "ts": None, "amount": None, "tags": [5, 5]},
272+
]
273+
274+
statements = []
275+
276+
def record(conn, cursor, statement, parameters, context, executemany):
277+
statements.append(statement)
278+
279+
await conn.run_sync(table.create)
280+
sqlalchemy.event.listen(conn.sync_connection, "before_cursor_execute", record)
281+
try:
282+
result = await conn.execute(
283+
table.insert(), rows, execution_options={"insertmanyvalues_page_size": 2}
284+
)
285+
finally:
286+
sqlalchemy.event.remove(conn.sync_connection, "before_cursor_execute", record)
287+
288+
# One event per page; a DB API executemany would fire a single event.
289+
assert len(statements) == 3
290+
assert result.rowcount == 5
291+
actual = (await conn.execute(select(table).order_by(table.c.id))).mappings().all()
292+
assert [dict(row) for row in actual] == rows

‎tests/pyathena/sqlalchemy/test_base.py‎

Lines changed: 201 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
from pyathena.converter import DefaultTypeConverter
2323
from pyathena.cursor import Cursor
2424
from pyathena.error import DatabaseError, OperationalError
25+
from pyathena.formatter import DefaultParameterFormatter
2526
from pyathena.sqlalchemy.base import AthenaDialect
2627
from pyathena.sqlalchemy.types import (
2728
TINYINT,
@@ -60,6 +61,48 @@ def unique_s3tables_table_name(base: str) -> str:
6061
return f"{base}_{uuid.uuid4().hex[:8]}"
6162

6263

64+
def recording_engine(rowcounts=None, **kwargs):
65+
"""Create an engine whose DB API connection records statements offline.
66+
67+
Args:
68+
rowcounts: Row counts reported by successive cursor calls. By default,
69+
a call reports the number of rows it inserted.
70+
**kwargs: Additional keyword arguments for ``create_engine``.
71+
72+
Returns:
73+
A tuple of the engine and the list of recorded
74+
``(method, operation, parameters)`` calls.
75+
"""
76+
calls = []
77+
counts = iter(rowcounts or ())
78+
79+
class RecordingCursor:
80+
description = None
81+
rowcount = -1
82+
83+
def execute(self, operation, parameters=None, **_):
84+
calls.append(("execute", operation, parameters))
85+
rows = sum(1 for key in parameters if key.startswith("id"))
86+
self.rowcount = next(counts, rows)
87+
88+
def executemany(self, operation, seq_of_parameters, **_):
89+
calls.append(("executemany", operation, seq_of_parameters))
90+
self.rowcount = next(counts, len(seq_of_parameters))
91+
92+
def close(self):
93+
pass
94+
95+
connection = SimpleNamespace(
96+
cursor=RecordingCursor, close=lambda: None, commit=lambda: None, rollback=lambda: None
97+
)
98+
engine = create_engine(
99+
"awsathena+rest://athena.us-west-2.amazonaws.com/default",
100+
creator=lambda: connection,
101+
**kwargs,
102+
)
103+
return engine, calls
104+
105+
63106
class TestAthenaDialect:
64107
def test_columns_from_information_schema(self):
65108
# Rows arrive unordered, and Athena reports a missing comment as NULL.
@@ -490,6 +533,108 @@ def list_table_metadata(**kwargs):
490533
("pyathena_table_metadata", "other_catalog", "default", table_name): listed
491534
}
492535

536+
def test_insertmanyvalues_pages(self):
537+
engine, calls = recording_engine()
538+
table = Table("t", MetaData(), Column("id", types.Integer), Column("name", types.String))
539+
540+
with engine.connect() as conn:
541+
result = conn.execute(
542+
table.insert(), [{"id": i, "name": f"name {i}"} for i in range(250)]
543+
)
544+
545+
# 100 rows per statement by default, and the total row count.
546+
assert [(method, len(parameters) // 2) for method, _, parameters in calls] == [
547+
("execute", 100),
548+
("execute", 100),
549+
("execute", 50),
550+
]
551+
assert result.rowcount == 250
552+
_, operation, parameters = calls[-1]
553+
assert (
554+
DefaultParameterFormatter()
555+
.format(operation, parameters)
556+
.startswith("INSERT INTO t (id, name) VALUES (200, 'name 200'), (201, 'name 201'), ")
557+
)
558+
559+
def test_insertmanyvalues_unknown_rowcount(self):
560+
engine, _ = recording_engine(rowcounts=[100, -1, 50])
561+
table = Table("t", MetaData(), Column("id", types.Integer))
562+
563+
with engine.connect() as conn:
564+
result = conn.execute(table.insert(), [{"id": i} for i in range(250)])
565+
566+
assert result.rowcount == -1
567+
568+
@pytest.mark.parametrize("configure", ["engine", "execution_options"])
569+
def test_insertmanyvalues_page_size(self, configure):
570+
engine_kwargs = {"insertmanyvalues_page_size": 2} if configure == "engine" else {}
571+
engine, calls = recording_engine(**engine_kwargs)
572+
table = Table("t", MetaData(), Column("id", types.Integer))
573+
574+
with engine.connect() as conn:
575+
if configure == "execution_options":
576+
conn = conn.execution_options(insertmanyvalues_page_size=2)
577+
result = conn.execute(table.insert(), [{"id": i} for i in range(5)])
578+
579+
assert [len(parameters) for _, _, parameters in calls] == [2, 2, 1]
580+
assert result.rowcount == 5
581+
582+
def test_insertmanyvalues_disabled(self):
583+
engine, calls = recording_engine(use_insertmanyvalues=False)
584+
table = Table("t", MetaData(), Column("id", types.Integer))
585+
586+
with engine.connect() as conn:
587+
result = conn.execute(table.insert(), [{"id": i} for i in range(3)])
588+
589+
((method, operation, parameters),) = calls
590+
assert method == "executemany"
591+
assert operation == "INSERT INTO t (id) VALUES (%(id)s)"
592+
assert parameters == [{"id": 0}, {"id": 1}, {"id": 2}]
593+
assert result.rowcount == 3
594+
595+
def test_insertmanyvalues_formats_rows(self):
596+
engine, calls = recording_engine()
597+
table = Table(
598+
"t",
599+
MetaData(),
600+
Column("id", types.Integer),
601+
Column("name", types.String),
602+
Column("data", types.LargeBinary),
603+
Column("ts", types.DateTime),
604+
Column("amount", types.Numeric(10, 3)),
605+
Column("tags", AthenaArray(types.Integer)),
606+
)
607+
rows = [
608+
{
609+
"id": 1,
610+
"name": "it's 100%",
611+
"data": b"\x00\x01",
612+
"ts": datetime(2026, 1, 2, 3, 4, 5, 123000),
613+
"amount": Decimal("1.5"),
614+
"tags": [1, None],
615+
},
616+
{
617+
"id": 2,
618+
"name": None,
619+
"data": None,
620+
"ts": datetime(2026, 1, 2, 3, 4, 5, 123456),
621+
"amount": None,
622+
"tags": None,
623+
},
624+
]
625+
626+
with engine.connect() as conn:
627+
conn.execute(table.insert(), rows)
628+
629+
((_, operation, parameters),) = calls
630+
assert DefaultParameterFormatter().format(operation, parameters) == (
631+
"INSERT INTO t (id, name, data, ts, amount, tags) VALUES "
632+
"(1, 'it''s 100%', X'0001', TIMESTAMP '2026-01-02 03:04:05.123', DECIMAL '1.5', "
633+
"CAST(ARRAY[1, null] AS ARRAY(INTEGER))), "
634+
"(2, null, null, TIMESTAMP '2026-01-02 03:04:05.123456', null, "
635+
"CAST(null AS ARRAY(INTEGER)))"
636+
)
637+
493638

494639
class TestSQLAlchemyAthena:
495640
@pytest.mark.parametrize(
@@ -2951,6 +3096,62 @@ def test_insert_from_select_cte_follows_insert_one(self, engine):
29513096
)
29523097
assert actual == [(2, "bar")]
29533098

3099+
def test_insertmanyvalues(self, engine):
3100+
engine, conn = engine
3101+
table_name = "insertmanyvalues"
3102+
table = Table(
3103+
table_name,
3104+
MetaData(schema=ENV.schema),
3105+
Column("id", types.Integer),
3106+
Column("name", types.String),
3107+
Column("data", types.LargeBinary),
3108+
Column("ts", types.DateTime),
3109+
Column("amount", types.Numeric(10, 3)),
3110+
Column("tags", AthenaArray(types.Integer)),
3111+
awsathena_location=f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/",
3112+
awsathena_tblproperties={"table_type": "ICEBERG"},
3113+
)
3114+
rows = [
3115+
{
3116+
"id": 1,
3117+
"name": "it's",
3118+
"data": b"\x00\x01",
3119+
"ts": datetime(2026, 1, 2, 3, 4, 5, 123000),
3120+
"amount": Decimal("1.5"),
3121+
"tags": [1, None],
3122+
},
3123+
{
3124+
"id": 2,
3125+
"name": None,
3126+
"data": None,
3127+
"ts": datetime(2026, 1, 2, 3, 4, 5, 123456),
3128+
"amount": Decimal("12.345"),
3129+
"tags": None,
3130+
},
3131+
{"id": 3, "name": "c", "data": b"", "ts": None, "amount": None, "tags": []},
3132+
{"id": 4, "name": "d", "data": b"d", "ts": None, "amount": None, "tags": [4]},
3133+
{"id": 5, "name": "e", "data": b"e", "ts": None, "amount": None, "tags": [5, 5]},
3134+
]
3135+
statements = []
3136+
3137+
def record(conn, cursor, statement, parameters, context, executemany):
3138+
statements.append(statement)
3139+
3140+
table.create(bind=conn)
3141+
sqlalchemy.event.listen(conn, "before_cursor_execute", record)
3142+
try:
3143+
result = conn.execution_options(insertmanyvalues_page_size=2).execute(
3144+
table.insert(), rows
3145+
)
3146+
finally:
3147+
sqlalchemy.event.remove(conn, "before_cursor_execute", record)
3148+
3149+
# Rows with different literal precisions share each multi-row statement.
3150+
assert len(statements) == 3
3151+
assert result.rowcount == 5
3152+
actual = conn.execute(sqlalchemy.select(table).order_by(table.c.id)).mappings().all()
3153+
assert [dict(row) for row in actual] == rows
3154+
29543155
def test_get_view_definition(self, engine):
29553156
engine, conn = engine
29563157
insp = sqlalchemy.inspect(engine)

0 commit comments

Comments
 (0)