mirror of
https://github.com/element-hq/synapse.git
synced 2026-08-14 09:20:49 +00:00
Route the Rust backend's bulk write paths off psycopg2's execute_values
The Rust shim has no `psycopg2.extras`, so the bulk write helpers that Synapse
runs via `execute_values` need a shim-compatible path. Add a
`uses_psycopg2_extras` engine capability flag (True for psycopg2, False for the
Rust and SQLite engines) and branch the write paths on it:
- The bulk INSERT/DELETE helpers (`simple_insert_many_txn`,
`simple_upsert_many_txn`, `simple_delete_many_batch_txn`, and the
sliding-sync / relations / state-deletion inserts) use psycopg2's single
multi-row `execute_values` on psycopg2, and `execute_batch` (the shim's
pipelined `executemany`) on the Rust and SQLite backends.
- `execute_values` itself is now psycopg2-only (it asserts
`uses_psycopg2_extras`); no query path splices row values into SQL as
literals.
- The four VALUES-join / RETURNING queries that `execute_values(fetch=True)`
served (in `event_federation` and `end_to_end_keys`) can't be expressed by
`executemany`, so they use a parameterised `unnest(...)` form that works on
all three backends.
Also cast the sequence name in `nextval(?::text)`: the Rust driver binds typed
parameters and can't produce a `regclass` from a name, whereas a text argument
coerces (this is the only parameterised sequence call, and Postgres-only SQL, so
psycopg2 and sqlite are unaffected).
psycopg2 behaviour and performance are unchanged; the Rust and SQLite bulk paths
now match.
Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01W3G4M92AmwSSZCbmtMJU3d
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
7cc50003b1
commit
b83b3e3ede
+35
-13
@@ -500,7 +500,10 @@ class LoggingTransaction:
|
||||
More efficient than `executemany` on PostgreSQL
|
||||
"""
|
||||
|
||||
if isinstance(self.database_engine, PostgresEngine):
|
||||
if (
|
||||
isinstance(self.database_engine, PostgresEngine)
|
||||
and self.database_engine.uses_psycopg2_extras
|
||||
):
|
||||
from psycopg2.extras import execute_batch
|
||||
|
||||
# TODO: is it safe for values to be Iterable[Iterable[Any]] here?
|
||||
@@ -510,6 +513,8 @@ class LoggingTransaction:
|
||||
lambda the_sql: execute_batch(self.txn, the_sql, args), sql
|
||||
)
|
||||
else:
|
||||
# The Rust backend has no psycopg2 extras; its `executemany` is
|
||||
# already pipelined, so route there (as the sqlite path does too).
|
||||
# TODO: is it safe for values to be Iterable[Iterable[Any]] here?
|
||||
# https://docs.python.org/3/library/sqlite3.html?highlight=sqlite3#sqlite3.Cursor.executemany
|
||||
# suggests that the outer collection may be iterable, but
|
||||
@@ -524,8 +529,11 @@ class LoggingTransaction:
|
||||
template: str | None = None,
|
||||
fetch: bool = True,
|
||||
) -> list[tuple]:
|
||||
"""Corresponds to psycopg2.extras.execute_values. Only available when
|
||||
using postgres.
|
||||
"""Corresponds to psycopg2.extras.execute_values.
|
||||
|
||||
Only usable on the psycopg2 backend: the Rust backend has no psycopg2
|
||||
extras, so its callers take shim-backed paths instead — bulk statements
|
||||
through ``executemany`` and the VALUES-join queries through ``unnest()``.
|
||||
|
||||
The `fetch` parameter must be set to False if the query does not return
|
||||
rows (e.g. INSERTs).
|
||||
@@ -534,6 +542,8 @@ class LoggingTransaction:
|
||||
compose the query.
|
||||
"""
|
||||
assert isinstance(self.database_engine, PostgresEngine)
|
||||
assert self.database_engine.uses_psycopg2_extras
|
||||
|
||||
from psycopg2.extras import execute_values
|
||||
|
||||
return self._do_execute(
|
||||
@@ -1343,9 +1353,12 @@ class DatabasePool:
|
||||
if not values:
|
||||
return
|
||||
|
||||
if isinstance(txn.database_engine, PostgresEngine):
|
||||
# We use `execute_values` as it can be a lot faster than `execute_batch`,
|
||||
# but it's only available on postgres.
|
||||
if (
|
||||
isinstance(txn.database_engine, PostgresEngine)
|
||||
and txn.database_engine.uses_psycopg2_extras
|
||||
):
|
||||
# On psycopg2 a single multi-row INSERT via `execute_values` is a lot
|
||||
# faster than `execute_batch`.
|
||||
sql = "INSERT INTO %s (%s) VALUES ?" % (
|
||||
table,
|
||||
", ".join(k for k in keys),
|
||||
@@ -1353,6 +1366,7 @@ class DatabasePool:
|
||||
|
||||
txn.execute_values(sql, values, fetch=False)
|
||||
else:
|
||||
# The Rust backend and SQLite go through `executemany` (execute_batch).
|
||||
sql = "INSERT INTO %s (%s) VALUES(%s)" % (
|
||||
table,
|
||||
", ".join(k for k in keys),
|
||||
@@ -1813,9 +1827,12 @@ class DatabasePool:
|
||||
for x, y in zip(key_values, value_values):
|
||||
args.append(tuple(x) + tuple(y))
|
||||
|
||||
if isinstance(txn.database_engine, PostgresEngine):
|
||||
# We use `execute_values` as it can be a lot faster than `execute_batch`,
|
||||
# but it's only available on postgres.
|
||||
if (
|
||||
isinstance(txn.database_engine, PostgresEngine)
|
||||
and txn.database_engine.uses_psycopg2_extras
|
||||
):
|
||||
# On psycopg2 a single multi-row INSERT via `execute_values` is a lot
|
||||
# faster than `execute_batch`.
|
||||
sql = "INSERT INTO %s (%s) VALUES ? ON CONFLICT (%s) DO %s" % (
|
||||
table,
|
||||
", ".join(k for k in allnames),
|
||||
@@ -1826,6 +1843,7 @@ class DatabasePool:
|
||||
txn.execute_values(sql, args, fetch=False)
|
||||
|
||||
else:
|
||||
# The Rust backend and SQLite go through `executemany` (execute_batch).
|
||||
sql = "INSERT INTO %s (%s) VALUES (%s) ON CONFLICT (%s) DO %s" % (
|
||||
table,
|
||||
", ".join(k for k in allnames),
|
||||
@@ -1834,7 +1852,7 @@ class DatabasePool:
|
||||
latter,
|
||||
)
|
||||
|
||||
return txn.execute_batch(sql, args)
|
||||
txn.execute_batch(sql, args)
|
||||
|
||||
@overload
|
||||
async def simple_select_one(
|
||||
@@ -2580,9 +2598,12 @@ class DatabasePool:
|
||||
values: for each row, a list of values in the same order as `keys`
|
||||
"""
|
||||
|
||||
if isinstance(txn.database_engine, PostgresEngine):
|
||||
# We use `execute_values` as it can be a lot faster than `execute_batch`,
|
||||
# but it's only available on postgres.
|
||||
if (
|
||||
isinstance(txn.database_engine, PostgresEngine)
|
||||
and txn.database_engine.uses_psycopg2_extras
|
||||
):
|
||||
# On psycopg2 a single `DELETE ... IN (VALUES ?)` via `execute_values`
|
||||
# is a lot faster than `execute_batch`.
|
||||
sql = "DELETE FROM %s WHERE (%s) IN (VALUES ?)" % (
|
||||
table,
|
||||
", ".join(k for k in keys),
|
||||
@@ -2590,6 +2611,7 @@ class DatabasePool:
|
||||
|
||||
txn.execute_values(sql, values, fetch=False)
|
||||
else:
|
||||
# The Rust backend and SQLite go through `executemany` (execute_batch).
|
||||
sql = "DELETE FROM %s WHERE (%s) = (%s)" % (
|
||||
table,
|
||||
", ".join(k for k in keys),
|
||||
|
||||
@@ -1208,8 +1208,8 @@ class EndToEndKeyWorkerStore(EndToEndKeyBackgroundStore, CacheInvalidationWorker
|
||||
results: dict[str, dict[str, dict[str, JsonDict]]] = {}
|
||||
missing: list[tuple[str, str, str, int]] = []
|
||||
if isinstance(self.database_engine, PostgresEngine):
|
||||
# If we can use execute_values we can use a single batch query
|
||||
# in autocommit mode.
|
||||
# On Postgres we can claim everything in a single batch query in
|
||||
# autocommit mode.
|
||||
unfulfilled_claim_counts: dict[tuple[str, str, str], int] = {}
|
||||
for user_id, device_id, algorithm, count in query_list:
|
||||
unfulfilled_claim_counts[user_id, device_id, algorithm] = count
|
||||
@@ -1277,10 +1277,11 @@ class EndToEndKeyWorkerStore(EndToEndKeyBackgroundStore, CacheInvalidationWorker
|
||||
query_list,
|
||||
db_autocommit=True,
|
||||
)
|
||||
# Use an UPDATE FROM... RETURNING combined with a VALUES block to do
|
||||
# everything in one query. Note: this is also supported in SQLite 3.33.0,
|
||||
# (see https://www.sqlite.org/lang_update.html#update_from), but we do not
|
||||
# have an equivalent of psycopg2's execute_values to do this in one query.
|
||||
# Use an UPDATE FROM... RETURNING combined with an unnest()ed set to
|
||||
# do everything in one query. Note: UPDATE ... FROM is also supported
|
||||
# in SQLite 3.33.0 (see
|
||||
# https://www.sqlite.org/lang_update.html#update_from), but we keep the
|
||||
# per-key fallback there.
|
||||
else:
|
||||
return await self._claim_e2e_fallback_keys_simple(query_list)
|
||||
|
||||
@@ -1295,20 +1296,33 @@ class EndToEndKeyWorkerStore(EndToEndKeyBackgroundStore, CacheInvalidationWorker
|
||||
"""
|
||||
results: dict[str, dict[str, dict[str, JsonDict]]] = {}
|
||||
|
||||
sql = """
|
||||
WITH claims(user_id, device_id, algorithm, mark_as_used) AS (
|
||||
VALUES ?
|
||||
)
|
||||
UPDATE e2e_fallback_keys_json k
|
||||
SET used = used OR mark_as_used
|
||||
FROM claims
|
||||
WHERE (k.user_id, k.device_id, k.algorithm) = (claims.user_id, claims.device_id, claims.algorithm)
|
||||
RETURNING k.user_id, k.device_id, k.algorithm, k.key_id, k.key_json;
|
||||
"""
|
||||
claimed_keys = cast(
|
||||
list[tuple[str, str, str, str, str]],
|
||||
txn.execute_values(sql, query_list),
|
||||
)
|
||||
# Unnest the (user_id, device_id, algorithm, mark_as_used) tuples into the
|
||||
# `claims` set. The per-column `::` casts let the parameters bind as arrays
|
||||
# — and let the Rust driver, which prepares statements, resolve their types.
|
||||
user_ids: list[str] = []
|
||||
device_ids: list[str] = []
|
||||
algorithms: list[str] = []
|
||||
marks: list[bool] = []
|
||||
for user_id, device_id, algorithm, mark_as_used in query_list:
|
||||
user_ids.append(user_id)
|
||||
device_ids.append(device_id)
|
||||
algorithms.append(algorithm)
|
||||
marks.append(mark_as_used)
|
||||
|
||||
claimed_keys: list[tuple[str, str, str, str, str]] = []
|
||||
if user_ids:
|
||||
sql = """
|
||||
WITH claims(user_id, device_id, algorithm, mark_as_used) AS (
|
||||
SELECT * FROM unnest(?::text[], ?::text[], ?::text[], ?::boolean[])
|
||||
)
|
||||
UPDATE e2e_fallback_keys_json k
|
||||
SET used = used OR mark_as_used
|
||||
FROM claims
|
||||
WHERE (k.user_id, k.device_id, k.algorithm) = (claims.user_id, claims.device_id, claims.algorithm)
|
||||
RETURNING k.user_id, k.device_id, k.algorithm, k.key_id, k.key_json;
|
||||
"""
|
||||
txn.execute(sql, (user_ids, device_ids, algorithms, marks))
|
||||
claimed_keys = cast(list[tuple[str, str, str, str, str]], txn.fetchall())
|
||||
|
||||
seen_user_device: set[tuple[str, str]] = set()
|
||||
for user_id, device_id, algorithm, key_id, key_json in claimed_keys:
|
||||
@@ -1438,30 +1452,44 @@ class EndToEndKeyWorkerStore(EndToEndKeyBackgroundStore, CacheInvalidationWorker
|
||||
# Doing so means that keys are issued in the same order they were uploaded,
|
||||
# which reduces the chances of a client expiring its copy of a (private)
|
||||
# key while the public key is still on the server, waiting to be issued.
|
||||
sql = """
|
||||
WITH claims(user_id, device_id, algorithm, claim_count) AS (
|
||||
VALUES ?
|
||||
), ranked_keys AS (
|
||||
SELECT
|
||||
user_id, device_id, algorithm, key_id, claim_count,
|
||||
ROW_NUMBER() OVER (
|
||||
PARTITION BY (user_id, device_id, algorithm)
|
||||
ORDER BY ts_added_ms
|
||||
) AS r
|
||||
FROM e2e_one_time_keys_json
|
||||
JOIN claims USING (user_id, device_id, algorithm)
|
||||
)
|
||||
DELETE FROM e2e_one_time_keys_json k
|
||||
WHERE (user_id, device_id, algorithm, key_id) IN (
|
||||
SELECT user_id, device_id, algorithm, key_id
|
||||
FROM ranked_keys
|
||||
WHERE r <= claim_count
|
||||
)
|
||||
RETURNING user_id, device_id, algorithm, key_id, key_json;
|
||||
"""
|
||||
otk_rows = cast(
|
||||
list[tuple[str, str, str, str, str]], txn.execute_values(sql, query_list)
|
||||
)
|
||||
# Unnest the (user_id, device_id, algorithm, claim_count) tuples into the
|
||||
# `claims` set. The per-column `::` casts let the parameters bind as arrays
|
||||
# — and let the Rust driver, which prepares statements, resolve their types.
|
||||
user_ids: list[str] = []
|
||||
device_ids: list[str] = []
|
||||
algorithms: list[str] = []
|
||||
claim_counts: list[int] = []
|
||||
for user_id, device_id, algorithm, claim_count in query_list:
|
||||
user_ids.append(user_id)
|
||||
device_ids.append(device_id)
|
||||
algorithms.append(algorithm)
|
||||
claim_counts.append(claim_count)
|
||||
|
||||
otk_rows: list[tuple[str, str, str, str, str]] = []
|
||||
if user_ids:
|
||||
sql = """
|
||||
WITH claims(user_id, device_id, algorithm, claim_count) AS (
|
||||
SELECT * FROM unnest(?::text[], ?::text[], ?::text[], ?::bigint[])
|
||||
), ranked_keys AS (
|
||||
SELECT
|
||||
user_id, device_id, algorithm, key_id, claim_count,
|
||||
ROW_NUMBER() OVER (
|
||||
PARTITION BY (user_id, device_id, algorithm)
|
||||
ORDER BY ts_added_ms
|
||||
) AS r
|
||||
FROM e2e_one_time_keys_json
|
||||
JOIN claims USING (user_id, device_id, algorithm)
|
||||
)
|
||||
DELETE FROM e2e_one_time_keys_json k
|
||||
WHERE (user_id, device_id, algorithm, key_id) IN (
|
||||
SELECT user_id, device_id, algorithm, key_id
|
||||
FROM ranked_keys
|
||||
WHERE r <= claim_count
|
||||
)
|
||||
RETURNING user_id, device_id, algorithm, key_id, key_json;
|
||||
"""
|
||||
txn.execute(sql, (user_ids, device_ids, algorithms, claim_counts))
|
||||
otk_rows = cast(list[tuple[str, str, str, str, str]], txn.fetchall())
|
||||
|
||||
seen_user_device = {
|
||||
(user_id, device_id) for user_id, device_id, _, _, _ in otk_rows
|
||||
|
||||
@@ -337,18 +337,22 @@ class EventFederationWorkerStore(
|
||||
results = set()
|
||||
|
||||
if isinstance(self.database_engine, PostgresEngine):
|
||||
# We can use `execute_values` to efficiently fetch the gaps when
|
||||
# using postgres.
|
||||
sql = """
|
||||
SELECT event_id
|
||||
FROM event_auth_chains AS c, (VALUES ?) AS l(chain_id, max_seq)
|
||||
WHERE
|
||||
c.chain_id = l.chain_id
|
||||
AND sequence_number <= max_seq
|
||||
"""
|
||||
# Fetch the gaps for every chain in a single query by unnesting the
|
||||
# (chain_id, max_seq) pairs into a joinable set. The `::bigint[]`
|
||||
# casts let the parameters bind as arrays — and let the Rust driver,
|
||||
# which prepares statements, resolve their types.
|
||||
if chains:
|
||||
sql = """
|
||||
SELECT event_id
|
||||
FROM event_auth_chains AS c,
|
||||
unnest(?::bigint[], ?::bigint[]) AS l(chain_id, max_seq)
|
||||
WHERE
|
||||
c.chain_id = l.chain_id
|
||||
AND sequence_number <= max_seq
|
||||
"""
|
||||
|
||||
rows = txn.execute_values(sql, chains.items())
|
||||
results.update(r for (r,) in rows)
|
||||
txn.execute(sql, (list(chains.keys()), list(chains.values())))
|
||||
results.update(r for (r,) in txn)
|
||||
else:
|
||||
# For SQLite we just fall back to doing a noddy for loop.
|
||||
sql = """
|
||||
@@ -883,23 +887,31 @@ class EventFederationWorkerStore(
|
||||
) -> set[str]:
|
||||
result: set[str] = set()
|
||||
if isinstance(self.database_engine, PostgresEngine):
|
||||
# We can use `execute_values` to efficiently fetch the gaps when
|
||||
# using postgres.
|
||||
sql = """
|
||||
SELECT event_id
|
||||
FROM event_auth_chains AS c, (VALUES ?) AS l(chain_id, min_seq, max_seq)
|
||||
WHERE
|
||||
c.chain_id = l.chain_id
|
||||
AND min_seq < sequence_number AND sequence_number <= max_seq
|
||||
"""
|
||||
# Fetch the gaps for every chain in a single query by unnesting the
|
||||
# (chain_id, min_seq, max_seq) triples into a joinable set. The
|
||||
# `::bigint[]` casts let the parameters bind as arrays — and let the
|
||||
# Rust driver, which prepares statements, resolve their types.
|
||||
if chains:
|
||||
chain_ids: list[int] = []
|
||||
min_seqs: list[int] = []
|
||||
max_seqs: list[int] = []
|
||||
for chain_id, (min_no, max_no) in chains.items():
|
||||
chain_ids.append(chain_id)
|
||||
min_seqs.append(min_no)
|
||||
max_seqs.append(max_no)
|
||||
|
||||
args = [
|
||||
(chain_id, min_no, max_no)
|
||||
for chain_id, (min_no, max_no) in chains.items()
|
||||
]
|
||||
sql = """
|
||||
SELECT event_id
|
||||
FROM event_auth_chains AS c,
|
||||
unnest(?::bigint[], ?::bigint[], ?::bigint[])
|
||||
AS l(chain_id, min_seq, max_seq)
|
||||
WHERE
|
||||
c.chain_id = l.chain_id
|
||||
AND min_seq < sequence_number AND sequence_number <= max_seq
|
||||
"""
|
||||
|
||||
rows = txn.execute_values(sql, args)
|
||||
result.update(r for (r,) in rows)
|
||||
txn.execute(sql, (chain_ids, min_seqs, max_seqs))
|
||||
result.update(r for (r,) in txn)
|
||||
else:
|
||||
# For SQLite we just fall back to doing a noddy for loop.
|
||||
sql = """
|
||||
|
||||
@@ -138,9 +138,13 @@ class RelationsWorkerStore(SQLBaseStore):
|
||||
ON CONFLICT (room_id, thread_id)
|
||||
DO NOTHING
|
||||
"""
|
||||
if isinstance(txn.database_engine, PostgresEngine):
|
||||
if (
|
||||
isinstance(txn.database_engine, PostgresEngine)
|
||||
and txn.database_engine.uses_psycopg2_extras
|
||||
):
|
||||
txn.execute_values(sql % ("?",), rows, fetch=False)
|
||||
else:
|
||||
# The Rust backend and SQLite go through executemany.
|
||||
txn.execute_batch(sql % ("(?, ?, ?, ?, ?)",), rows)
|
||||
|
||||
# Mark the progress.
|
||||
|
||||
@@ -749,10 +749,14 @@ class SlidingSyncStore(SQLBaseStore):
|
||||
for room_id, user_id in to_update
|
||||
]
|
||||
|
||||
if isinstance(self.database_engine, PostgresEngine):
|
||||
if (
|
||||
isinstance(self.database_engine, PostgresEngine)
|
||||
and self.database_engine.uses_psycopg2_extras
|
||||
):
|
||||
sql = sql.format(value_placeholder="?")
|
||||
txn.execute_values(sql, args, fetch=False)
|
||||
else:
|
||||
# The Rust backend and SQLite go through executemany.
|
||||
sql = sql.format(value_placeholder="(?, ?, ?, ?, ?)")
|
||||
txn.execute_batch(sql, args)
|
||||
|
||||
|
||||
@@ -696,38 +696,49 @@ class UserDirectoryBackgroundUpdateStore(StateDeltasStore):
|
||||
|
||||
if isinstance(self.database_engine, PostgresEngine):
|
||||
# We weight the localpart most highly, then display name and finally
|
||||
# server name
|
||||
template = """
|
||||
# server name. Each row is (user_id, localpart, domain, display).
|
||||
rows = [
|
||||
(
|
||||
%s,
|
||||
setweight(to_tsvector('simple', %s), 'A')
|
||||
|| setweight(to_tsvector('simple', %s), 'D')
|
||||
|| setweight(to_tsvector('simple', COALESCE(%s, '')), 'B')
|
||||
p.user_id,
|
||||
get_localpart_from_id(p.user_id),
|
||||
get_domain_from_id(p.user_id),
|
||||
(
|
||||
_filter_text_for_index(p.display_name)
|
||||
if p.display_name
|
||||
else None
|
||||
),
|
||||
)
|
||||
"""
|
||||
|
||||
sql = """
|
||||
for p in profiles
|
||||
]
|
||||
if self.database_engine.uses_psycopg2_extras:
|
||||
# psycopg2: one multi-row INSERT, each row wrapped by the template.
|
||||
template = """
|
||||
(
|
||||
%s,
|
||||
setweight(to_tsvector('simple', %s), 'A')
|
||||
|| setweight(to_tsvector('simple', %s), 'D')
|
||||
|| setweight(to_tsvector('simple', COALESCE(%s, '')), 'B')
|
||||
)
|
||||
"""
|
||||
sql = """
|
||||
INSERT INTO user_directory_search(user_id, vector)
|
||||
VALUES ? ON CONFLICT (user_id) DO UPDATE SET vector=EXCLUDED.vector
|
||||
"""
|
||||
txn.execute_values(
|
||||
sql,
|
||||
[
|
||||
(
|
||||
p.user_id,
|
||||
get_localpart_from_id(p.user_id),
|
||||
get_domain_from_id(p.user_id),
|
||||
(
|
||||
_filter_text_for_index(p.display_name)
|
||||
if p.display_name
|
||||
else None
|
||||
),
|
||||
txn.execute_values(sql, rows, template=template, fetch=False)
|
||||
else:
|
||||
# The Rust backend has no execute_values; run the same per-row
|
||||
# tsvector upsert via executemany.
|
||||
sql = """
|
||||
INSERT INTO user_directory_search(user_id, vector)
|
||||
VALUES (
|
||||
?,
|
||||
setweight(to_tsvector('simple', ?), 'A')
|
||||
|| setweight(to_tsvector('simple', ?), 'D')
|
||||
|| setweight(to_tsvector('simple', COALESCE(?, '')), 'B')
|
||||
)
|
||||
for p in profiles
|
||||
],
|
||||
template=template,
|
||||
fetch=False,
|
||||
)
|
||||
ON CONFLICT (user_id) DO UPDATE SET vector=EXCLUDED.vector
|
||||
"""
|
||||
txn.execute_batch(sql, rows)
|
||||
elif isinstance(self.database_engine, Sqlite3Engine):
|
||||
values = []
|
||||
for p in profiles:
|
||||
|
||||
@@ -349,9 +349,13 @@ class StateDeletionDataStore:
|
||||
)
|
||||
for state_group in state_groups
|
||||
]
|
||||
if isinstance(txn.database_engine, PostgresEngine):
|
||||
if (
|
||||
isinstance(txn.database_engine, PostgresEngine)
|
||||
and txn.database_engine.uses_psycopg2_extras
|
||||
):
|
||||
txn.execute_values(sql % ("?",), rows, fetch=False)
|
||||
else:
|
||||
# The Rust backend and SQLite go through executemany.
|
||||
txn.execute_batch(sql % ("(?, ?)",), rows)
|
||||
|
||||
async def mark_state_groups_as_used(self, state_groups: Collection[int]) -> None:
|
||||
|
||||
@@ -43,6 +43,10 @@ class Psycopg2Engine(
|
||||
):
|
||||
"""The Postgres backend that talks to the database via psycopg2."""
|
||||
|
||||
# psycopg2's cursor is a real psycopg2 cursor, so the `psycopg2.extras`
|
||||
# helpers can be used on it directly.
|
||||
uses_psycopg2_extras: bool = True
|
||||
|
||||
def __init__(self, database_config: Mapping[str, Any]):
|
||||
super().__init__(psycopg2, database_config)
|
||||
psycopg2.extensions.register_type(psycopg2.extensions.UNICODE)
|
||||
|
||||
@@ -57,6 +57,12 @@ _RETRYABLE_PGCODES = ("40001", "40P01")
|
||||
class RustPostgresEngine(PostgresEngine[Connection, Cursor]):
|
||||
"""A :class:`PostgresEngine` that talks to the Rust backend's shim."""
|
||||
|
||||
# The shim cursor is not a psycopg2 cursor, so the `psycopg2.extras`
|
||||
# helpers can't be used on it; `LoggingTransaction.execute_batch` takes a
|
||||
# shim-backed path (its pipelined `executemany`) instead, and the callers of
|
||||
# `execute_values` fall back to `executemany` / `unnest()` for the Rust engine.
|
||||
uses_psycopg2_extras: bool = False
|
||||
|
||||
# SQL isolation-level names for each `IsolationLevel`. The shim has no
|
||||
# psycopg2-style `set_isolation_level`, so a per-transaction override is
|
||||
# applied as a `SET SESSION CHARACTERISTICS` statement (see
|
||||
|
||||
@@ -26,10 +26,11 @@ connection-level methods the database engine calls (``in_transaction``,
|
||||
``is_closed``, ``set_autocommit``) so a wrapped connection is a drop-in for the
|
||||
raw one.
|
||||
|
||||
Not handled here: ``execute_batch`` / ``execute_values`` (psycopg2 extras that
|
||||
``LoggingTransaction`` invokes directly for ``PostgresEngine``) still need a
|
||||
routing change in ``LoggingTransaction`` to reach a shim-backed implementation;
|
||||
that is a separate follow-up.
|
||||
The psycopg2 extras that ``LoggingTransaction`` uses are routed to shim-backed
|
||||
equivalents (via ``engine.uses_psycopg2_extras``): ``execute_batch`` maps onto
|
||||
the shim's pipelined ``executemany``. ``execute_values`` has no shim equivalent —
|
||||
the Rust callers use ``executemany`` (bulk statements) or ``unnest()`` (the
|
||||
VALUES-join queries) directly instead.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
@@ -100,14 +100,19 @@ class PostgresSequenceGenerator(SequenceGenerator):
|
||||
self._sequence_name = sequence_name
|
||||
|
||||
def get_next_id_txn(self, txn: Cursor) -> int:
|
||||
txn.execute("SELECT nextval(?)", (self._sequence_name,))
|
||||
# Cast the sequence name to text so the parameter is typed `text` rather
|
||||
# than inferred as `regclass`; Postgres then coerces the name, and the
|
||||
# native Rust driver (which binds typed parameters, and can't produce a
|
||||
# `regclass` from a name) doesn't need to special-case it.
|
||||
txn.execute("SELECT nextval(?::text)", (self._sequence_name,))
|
||||
fetch_res = txn.fetchone()
|
||||
assert fetch_res is not None
|
||||
return fetch_res[0]
|
||||
|
||||
def get_next_mult_txn(self, txn: Cursor, n: int) -> list[int]:
|
||||
txn.execute(
|
||||
"SELECT nextval(?) FROM generate_series(1, ?)", (self._sequence_name, n)
|
||||
"SELECT nextval(?::text) FROM generate_series(1, ?)",
|
||||
(self._sequence_name, n),
|
||||
)
|
||||
return [i for (i,) in txn]
|
||||
|
||||
|
||||
@@ -221,6 +221,14 @@ class RustDBAPIAdapterTestCase(unittest.TestCase):
|
||||
del self.conn
|
||||
self._pool.close()
|
||||
|
||||
def _logging_conn(self) -> LoggingDatabaseConnection:
|
||||
return LoggingDatabaseConnection(
|
||||
conn=self.conn,
|
||||
engine=self.engine,
|
||||
default_txn_name="test",
|
||||
server_name="test",
|
||||
)
|
||||
|
||||
def test_execute_and_fetchone(self) -> None:
|
||||
cursor = self.conn.cursor()
|
||||
# The adapter passes parameters straight through; the shim binds `$n`.
|
||||
@@ -270,13 +278,7 @@ class RustDBAPIAdapterTestCase(unittest.TestCase):
|
||||
# The whole point: a real LoggingTransaction (which converts `?` to `$n`
|
||||
# via the engine, then drives the cursor via the DBAPI2 spelling) runs
|
||||
# unchanged against the adapter.
|
||||
engine = RustPostgresEngine({})
|
||||
db_conn = LoggingDatabaseConnection(
|
||||
conn=self.conn,
|
||||
engine=engine,
|
||||
default_txn_name="test",
|
||||
server_name="test",
|
||||
)
|
||||
db_conn = self._logging_conn()
|
||||
|
||||
txn = db_conn.cursor(txn_name="test")
|
||||
txn.execute("SELECT ?::int + ?::int", (2, 3))
|
||||
@@ -346,3 +348,13 @@ class RustDBAPIAdapterTestCase(unittest.TestCase):
|
||||
before = backend_pid(conn)
|
||||
conn.reconnect()
|
||||
self.assertNotEqual(backend_pid(conn), before)
|
||||
|
||||
def test_execute_batch(self) -> None:
|
||||
# execute_batch routes to the shim's (pipelined) executemany.
|
||||
db_conn = self._logging_conn()
|
||||
txn = db_conn.cursor(txn_name="test")
|
||||
txn.execute("CREATE TEMP TABLE t (id int)")
|
||||
txn.execute_batch("INSERT INTO t VALUES (?)", [(1,), (2,), (3,)])
|
||||
txn.execute("SELECT id FROM t ORDER BY id")
|
||||
self.assertEqual(txn.fetchall(), [(1,), (2,), (3,)])
|
||||
db_conn.commit()
|
||||
|
||||
Reference in New Issue
Block a user