概要
- テンポラリテーブル作成
- SELECT/INSERT/UPDATE/DELETE文を実行
-
raise Rollback(tx)でロールバック確認 - pyinstrumentライブラリで
profile_report[<関数名>].htmlファイルに実行速度を出力
ライブラリ
# tzdataはPostgreSQLから送られてくるタイムゾーン名「Asia/Tokyo」をUTC+09:00と判定するために必要
pip install psycopg psycopg_pool tzdata pyinstrument
ソースコード
インポート文
from functools import wraps
from pyinstrument import Profiler
from psycopg.types.json import Json
from asyncio import run, SelectorEventLoop
from dataclasses import asdict, astuple, dataclass, fields
from datetime import date, timedelta, datetime
from random import choices
from string import ascii_letters, digits
from psycopg import Rollback
from psycopg_pool import AsyncConnectionPool
from psycopg.sql import SQL, Literal, Placeholder, Composable, Identifier
from psycopg.rows import DictRow, dict_row
from selectors import SelectSelector
from traceback import print_exc
from itertools import chain
from json import dumps
from pathlib import Path
main.py
# コネクションプール作成
pool = AsyncConnectionPool(
conninfo='user=postgres password=postgres',
open=False,
)
profiler = Profiler()
def profile_with_pyinstrument(func: function):
'''
関数の実行時間を pyinstrument で計測するデコレータ
・デコレータが付いたファンクションに対して計測を行う
'''
@wraps(func)
async def wrapper(*args, **kwargs):
profiler.start()
try:
return await func(*args, **kwargs)
finally:
profiler.stop()
print(profiler.output_text(True, True, flat=True))
now = datetime.now().strftime('%Y-%m-%d_%H_%M_%S')
profiler.write_html(
str(Path(__file__).parent) + f'/profile_report[{func.__name__}]_{now}.html')
return wrapper
async def create_table(conn):
'''
テンポラリテーブル作成
'''
await conn.execute('''
CREATE TEMPORARY TABLE temp_users (
id BIGINT GENERATED ALWAYS AS IDENTITY,
sub_id BIGINT DEFAULT 0,
name VARCHAR(64) NOT NULL,
email VARCHAR(254) UNIQUE,
json_data JSON NOT NULL,
memo TEXT,
del_flg BOOLEAN DEFAULT FALSE NOT NULL,
limit_at DATE NOT NULL,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY(id, sub_id)
) ON COMMIT DROP
''')
@dataclass(slots=True)
class InsertRecordTempUsers():
name: str = None
email: str = None
json_data: Json = None
memo: str = None
del_flg: bool = None
limit_at: date = None
#
# 8個全てのkeyが無くてもInsertRecordTempUsersをインスタンス化できるようにする
#
@classmethod
def from_dict(self, data: dict):
# dataclassからパラメータ一覧取得
field_names = {f.name for f in fields(self)}
# dataclass内に存在するパラメータ名を初期化
# ※dataclassに対し、引数dataに全てのカラムが無い場合でも存在するkeyだけ使って初期化可能とする
cleaned = {k: v for k, v in data.items() if k in field_names}
# dict型の場合は、Json型に変換する(json_dataに入れるため)
modified = {k: (Json(v, dumps) if type(v).__name__ != 'Json' and type(v).__name__ == 'dict' else v)
for k, v in cleaned.items()}
return self(**modified)
def new_records(nums: int = 10):
'''
テストレコード作成
'''
characters = ascii_letters + digits
records = []
for i in range(nums):
# 10桁の英数列を作る
local_address = ''.join(choices(characters, k=10))
records.append(InsertRecordTempUsers(
name='John',
email=local_address + '@*****.com',
json_data=Json({f'k{i+1}': f'v{i+1}'}),
memo='user',
del_flg=True,
limit_at=date.today() + timedelta(days=10))
)
return records
async def insert_copy_write_row(conn, records: list[InsertRecordTempUsers]):
'''
COPY文を使ったINSERT
'''
cols = []
if len(records) == 0:
return None
else:
cols = [f.name for f in fields(records[0])]
query = SQL('COPY {table}({field}) FROM STDIN').format(
table=Identifier('temp_users'),
field=SQL(',').join(map(Identifier, cols)))
if isinstance(query, Composable):
print(query.as_string(conn)) # 実行予定のSQL文を表示
# SQL実行
# COPY "temp_users"("name","email","json_data","memo","del_flg","limit_at") FROM STDIN
async with conn.cursor() as cur:
async with cur.copy(query) as copy:
for record in records:
await copy.write_row(astuple(record))
async def rollback_test(conn, nums=100_000):
# ロールバックの実行
try:
async with conn.transaction() as tx: # with句の範囲でBEGIN/ROLLBACK
await insert_copy_write_row(conn, new_records(nums))
raise Rollback(tx)
except Exception:
print_exc()
async def select(conn, fields: list = [], where: Composable = None, *, show_limit: int = 10, tail: bool = True, sql_limit: int = 1000, count_only=False):
async with conn.cursor(row_factory=dict_row) as cur:
if len(fields) > 0:
query = SQL('SELECT {fields} FROM {table}').format(
fields=SQL(',').join(map(Identifier, fields)),
table=Identifier('temp_users'))
else:
query = SQL('SELECT * FROM {}').format(Identifier('temp_users'))
if where is not None:
query += where
if count_only or (sql_limit is None or 0 >= sql_limit):
'全件レコード取得
await cur.execute(query)
if count_only:
print(f'{cur.rowcount:,}')
return
else:
result = await cur.execute(query + SQL(' LIMIT %s'), (sql_limit,))
rows: list[DictRow] | None = await result.fetchall()
# 余りにも多くの行をprintすると時間がかかるため、制限する
if show_limit > 0:
if tail:
rows = rows[-show_limit:]
else:
rows = rows[:show_limit]
if len(rows) > 0:
'結果レコードがある場合
[print(row) for row in rows]
models = [InsertRecordTempUsers.from_dict(row) for row in rows]
print(models[0].json_data)
print(models[0].json_data.obj)
else:
print('レコード0件')
async def update(conn, columns: dict[str, Composable | str], where: Composable = None, params: tuple[str | int | float] | tuple = ()):
query = SQL('UPDATE {table} SET {sets} ').format(
table=Identifier('temp_users'),
sets=SQL(',').join(map(lambda item: SQL('{key} = {value}').format(
key=Identifier(item[0]),
value=item[1] if isinstance(
item[1], Composable) else Literal(item[1]),
), columns.items())))
if where is not None:
query += where
if isinstance(query, Composable):
print(query.as_string(conn)) # 実行予定のSQL文を表示
async with conn.transaction(): # with句の範囲でBEGIN/COMMIT
if len(params) > 0:
await conn.execute(query, params)
else:
await conn.execute(query)
async def delete(conn, where: Composable = None):
query = SQL('DELETE FROM {table} ').format(
table=Identifier('temp_users'),
)
if where is not None:
query += where
if isinstance(query, Composable):
print(query.as_string(conn)) # 実行予定のSQL文を表示
async with conn.transaction(): # with句の範囲でBEGIN/COMMIT
await conn.execute(query)
@profile_with_pyinstrument
async def main():
try:
await pool.open()
async with pool.connection() as conn:
await conn.set_autocommit(False)
async with conn.transaction(): # with句の範囲でBEGIN/COMMIT
# テーブル作成
await create_table(conn)
# テーブルを多様な方法でINSERT(10万行挿入で速度確認)
# 約37秒。バルクインサート頑張って作ったのにマジのゴミです
# await insert_tuple_flatten(conn, new_records(100_000), bulk=200)
# await select(conn, count_only=True)
# 約28秒。1レコードずつ繰り返し挿入
# await insert_executemany(conn, new_records(100_000))
# await select(conn, count_only=True)
# 約5秒。COPY ... FROM STDINの使用。これ一択と言えるほど速い
await insert_copy_write_row(conn, new_records(100_000))
await select(conn, count_only=True)
# ロールバックするだけ
await rollback_test(conn, 100_000)
await select(conn, count_only=True)
# 更新
# await insert_copy_write_row(conn, new_records(5)) # 5レコード追加
await select(conn, ['json_data', 'limit_at', 'updated_at'])
update_values = {'json_data': SQL('json_data::JSONB || ') + Literal('{"special1": "v1"}'), # key=special1を追加
'limit_at': SQL('current_date + interval ') + Literal('1 days'),
'updated_at': SQL('NOW() + interval ') + Literal('1 second')}
await update(conn, update_values, SQL("WHERE json_data->>'k1' = 'v1'"))
await select(conn, ['json_data', 'limit_at', 'updated_at'], SQL("WHERE json_data->>'special1' = 'v1'"))
update_values = {'json_data': SQL(
# key=special1を削除
'json_data::JSONB - ') + Literal('special1')}
await update(conn, update_values, SQL("WHERE json_data->>'special1' = 'v1'"))
await select(conn, ['json_data', 'updated_at'], SQL("WHERE json_data->>'special1' = 'v1'"))
# 削除
await delete(conn, SQL("WHERE json_data->>'k1' = 'v1'"))
await select(conn, [], SQL("WHERE json_data->>'k1' = 'v1'"))
await delete(conn)
await select(conn)
except Exception:
print_exc()
finally:
await pool.close()
if __name__ == '__main__':
run(main(), loop_factory=lambda: SelectorEventLoop(SelectSelector()))
# run(main2(), loop_factory=lambda: SelectorEventLoop(SelectSelector()))
コメントアウトした全てのコード
@profile_with_pyinstrument
async def main2():
'''
複数の@profile_with_pyinstrumentデコレータを書いた場合の確認
'''
import asyncio
await asyncio.sleep(5)
async def insert_tuple_flatten(conn, records: list[InsertRecordTempUsers], bulk=100):
cols = []
if len(records) == 0:
return None
else:
cols = [f.name for f in fields(records[0])]
base_query = SQL('INSERT INTO {table}({field}) VALUES ').format(
table=Identifier('temp_users'),
field=SQL(',').join(map(Identifier, cols)))
for i in range(1, len(records) + 1):
if i % bulk == 1:
# 初回、またはSQLのかたまりが実行された直後の回
query = SQL('({})').format(
SQL(',').join([Placeholder()] * len(cols)))
else:
query += SQL(',({})').format(
SQL(',').join([Placeholder()] * len(cols)))
if (i > 1 and i % bulk == 0) or i == len(records):
# SQLのかたまりの上限数に達する、またはリストの末尾になった場合
# if isinstance(base_query, Composable) and isinstance(query, Composable):
# print((base_query + query).as_string(conn)) # 実行予定のSQL文を表示
if i == len(records) and i % bulk != 0:
# リストの末尾、かつSQLのかたまりの中間で終わりを迎える場合
post_idx = bulk * (i // bulk)
else:
post_idx = bulk * (i // bulk - 1) if i > bulk else 0
last_idx = i
# print(f'{post_idx=},{i=},{last_idx=}')
# SQL実行
# INSERT INTO "temp_users"("name","email","json_data","memo","del_flg","limit_at") VALUES
# (%s,%s,%s,%s,%s,%s)
# ,(%s,%s,%s,%s,%s,%s)
# ,(...)
# ,(%s,%s,%s,%s,%s,%s)
async with conn.transaction(): # with句の範囲でBEGIN/COMMIT
await conn.execute(base_query + query, tuple(chain.from_iterable(map(astuple, records[post_idx:last_idx]))))
async def insert_executemany(conn, records: list[InsertRecordTempUsers]):
cols = []
if len(records) == 0:
return None
else:
cols = [f.name for f in fields(records[0])]
query = SQL('INSERT INTO temp_users({field}) VALUES ({values})').format(
table=Identifier('temp_users'),
field=SQL(',').join(map(Identifier, cols)),
values=SQL(',').join(map(Placeholder, cols)))
if isinstance(query, Composable):
print(query.as_string(conn)) # 実行予定のSQL文を表示
# SQL実行
# INSERT INTO "temp_users"("name","email","json_data","memo","del_flg","limit_at")
# VALUES (%(name)s,%(email)s,%(json_data)s,%(memo)s,%(del_flg)s,%(limit_at)s)
async with conn.cursor() as cur:
await cur.executemany(query, map(asdict, records))