|
22 | 22 | from pyathena.converter import DefaultTypeConverter |
23 | 23 | from pyathena.cursor import Cursor |
24 | 24 | from pyathena.error import DatabaseError, OperationalError |
| 25 | +from pyathena.formatter import DefaultParameterFormatter |
25 | 26 | from pyathena.sqlalchemy.base import AthenaDialect |
26 | 27 | from pyathena.sqlalchemy.types import ( |
27 | 28 | TINYINT, |
@@ -60,6 +61,48 @@ def unique_s3tables_table_name(base: str) -> str: |
60 | 61 | return f"{base}_{uuid.uuid4().hex[:8]}" |
61 | 62 |
|
62 | 63 |
|
| 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 | + |
63 | 106 | class TestAthenaDialect: |
64 | 107 | def test_columns_from_information_schema(self): |
65 | 108 | # Rows arrive unordered, and Athena reports a missing comment as NULL. |
@@ -490,6 +533,108 @@ def list_table_metadata(**kwargs): |
490 | 533 | ("pyathena_table_metadata", "other_catalog", "default", table_name): listed |
491 | 534 | } |
492 | 535 |
|
| 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 | + |
493 | 638 |
|
494 | 639 | class TestSQLAlchemyAthena: |
495 | 640 | @pytest.mark.parametrize( |
@@ -2951,6 +3096,62 @@ def test_insert_from_select_cte_follows_insert_one(self, engine): |
2951 | 3096 | ) |
2952 | 3097 | assert actual == [(2, "bar")] |
2953 | 3098 |
|
| 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 | + |
2954 | 3155 | def test_get_view_definition(self, engine): |
2955 | 3156 | engine, conn = engine |
2956 | 3157 | insp = sqlalchemy.inspect(engine) |
|
0 commit comments