Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
150 changes: 145 additions & 5 deletions backend/bracket/database.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from contextvars import ContextVar
from typing import Any

import sqlalchemy
from databases import Database
import asyncpg
from heliclockter import datetime_utc

from bracket.config import config
Expand All @@ -12,7 +14,7 @@ def datetime_decoder(value: str) -> datetime_utc:
return datetime_utc.fromisoformat(value)


async def asyncpg_init(connection: Any) -> None:
async def _init_connection(connection: asyncpg.Connection) -> None: # type: ignore[type-arg]
for timestamp_type in ("timestamp", "timestamptz"):
await connection.set_type_codec(
timestamp_type,
Expand All @@ -22,6 +24,144 @@ async def asyncpg_init(connection: Any) -> None:
)


database = Database(str(config.pg_dsn), init=asyncpg_init)
def _convert_named_params(query: str, values: dict[str, Any]) -> tuple[str, list[Any]]:
"""Convert :param style parameters to $1, $2, ... style for asyncpg."""
params: list[Any] = []
result = query
# Sort keys by length (longest first) to avoid partial replacements
sorted_keys = sorted(values.keys(), key=len, reverse=True)
# First pass: replace all :param with unique placeholders
placeholders: dict[str, str] = {}
for key in sorted_keys:
placeholder = f"\x00PARAM_{key}\x00"
placeholders[key] = placeholder
result = result.replace(f":{key}", placeholder)
# Second pass: replace placeholders with $N
for key in sorted_keys:
placeholder = placeholders[key]
if placeholder in result:
params.append(values[key])
idx = len(params)
result = result.replace(placeholder, f"${idx}", 1)
# Handle multiple occurrences of same param
while placeholder in result:
params.append(values[key])
idx = len(params)
result = result.replace(placeholder, f"${idx}", 1)
return result, params

engine = sqlalchemy.create_engine(str(config.pg_dsn))

# Context variable to track the current transaction connection
_transaction_connection: ContextVar[asyncpg.Connection | None] = ContextVar( # type: ignore[type-arg]
"_transaction_connection", default=None
)


class _Transaction:
"""Context manager for database transactions."""

def __init__(self, pool: asyncpg.Pool) -> None: # type: ignore[type-arg]
self._pool = pool
self._connection: asyncpg.Connection | None = None # type: ignore[type-arg]
self._transaction: asyncpg.connection.transaction.Transaction | None = None
self._token: Any = None

async def __aenter__(self) -> "_Transaction":
self._connection = await self._pool.acquire()
self._transaction = self._connection.transaction()
await self._transaction.start()
self._token = _transaction_connection.set(self._connection)
return self

async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None:
assert self._connection is not None
assert self._transaction is not None
assert self._token is not None
try:
if exc_type is not None:
await self._transaction.rollback()
else:
await self._transaction.commit()
finally:
_transaction_connection.reset(self._token)
await self._pool.release(self._connection)


class DatabasePool:
"""asyncpg-based database wrapper providing an interface compatible with
the old `databases.Database` usage in this codebase."""

def __init__(self, dsn: str) -> None:
self._dsn = dsn
self._pool: asyncpg.Pool | None = None # type: ignore[type-arg]

async def connect(self) -> None:
self._pool = await asyncpg.create_pool(self._dsn, init=_init_connection)

async def disconnect(self) -> None:
if self._pool is not None:
await self._pool.close()
self._pool = None

@property
def pool(self) -> asyncpg.Pool: # type: ignore[type-arg]
assert self._pool is not None, "Database pool is not initialized. Call connect() first."
return self._pool

def _get_conn(self) -> asyncpg.Connection | asyncpg.Pool: # type: ignore[type-arg]
"""Return the transaction connection if inside a transaction, otherwise the pool."""
conn = _transaction_connection.get()
if conn is not None:
return conn
return self.pool

async def fetch_one(
self, query: str, values: dict[str, Any] | None = None
) -> asyncpg.Record | None:
conn = self._get_conn()
if values:
converted_query, params = _convert_named_params(query, values)
return await conn.fetchrow(converted_query, *params)
return await conn.fetchrow(query)

async def fetch_all(
self, query: str, values: dict[str, Any] | None = None
) -> list[asyncpg.Record]:
conn = self._get_conn()
if values:
converted_query, params = _convert_named_params(query, values)
return await conn.fetch(converted_query, *params)
return await conn.fetch(query)

async def fetch_val(
self, query: str, values: dict[str, Any] | None = None, column: int = 0
) -> Any:
conn = self._get_conn()
if values:
converted_query, params = _convert_named_params(query, values)
return await conn.fetchval(converted_query, *params, column=column)
return await conn.fetchval(query, column=column)

async def execute(self, query: str, values: dict[str, Any] | None = None) -> str:
conn = self._get_conn()
if values:
converted_query, params = _convert_named_params(query, values)
return await conn.execute(converted_query, *params)
return await conn.execute(query)

def transaction(self) -> _Transaction:
return _Transaction(self.pool)

@asynccontextmanager
async def acquire(self) -> AsyncIterator[asyncpg.Connection]: # type: ignore[type-arg]
async with self.pool.acquire() as conn:
yield conn

async def __aenter__(self) -> "DatabasePool":
return self

async def __aexit__(self, *args: Any) -> None:
pass


database = DatabasePool(str(config.pg_dsn))
5 changes: 3 additions & 2 deletions backend/bracket/routes/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,6 @@
from bracket.database import database
from bracket.models.db.tournament import Tournament
from bracket.models.db.user import UserInDB, UserPublic
from bracket.schema import tournaments
from bracket.sql.tournaments import sql_get_tournament_by_endpoint_name
from bracket.sql.users import get_user, get_user_access_to_club, get_user_access_to_tournament
from bracket.utils.db import fetch_all_parsed
Expand Down Expand Up @@ -140,7 +139,9 @@ async def user_authenticated_or_public_dashboard(
pass

tournaments_fetched = await fetch_all_parsed(
database, Tournament, tournaments.select().where(tournaments.c.id == tournament_id)
database, Tournament,
"SELECT * FROM tournaments WHERE id = :tournament_id",
{"tournament_id": tournament_id},
)
if len(tournaments_fetched) < 1 or not tournaments_fetched[0].dashboard_public:
raise HTTPException(
Expand Down
29 changes: 15 additions & 14 deletions backend/bracket/routes/courts.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@
)
from bracket.routes.models import CourtsResponse, SingleCourtResponse, SuccessResponse
from bracket.routes.util import disallow_archived_tournament
from bracket.schema import courts
from bracket.sql.courts import get_all_courts_in_tournament, sql_delete_court, update_court
from bracket.sql.stages import get_full_tournament_details
from bracket.utils.db import fetch_one_parsed
Expand Down Expand Up @@ -50,9 +49,8 @@ async def update_court_by_id(
await fetch_one_parsed(
database,
Court,
courts.select().where(
(courts.c.id == court_id) & (courts.c.tournament_id == tournament_id)
),
"SELECT * FROM courts WHERE id = :court_id AND tournament_id = :tournament_id",
{"court_id": court_id, "tournament_id": tournament_id},
)
)
)
Expand Down Expand Up @@ -94,22 +92,25 @@ async def create_court(
existing_courts = await get_all_courts_in_tournament(tournament_id)
check_requirement(existing_courts, user, "max_courts")

last_record_id = await database.execute(
query=courts.insert(),
values=CourtToInsert(
**court_body.model_dump(),
created=datetime_utc.now(),
tournament_id=tournament_id,
).model_dump(),
insertable = CourtToInsert(
**court_body.model_dump(),
created=datetime_utc.now(),
tournament_id=tournament_id,
)
values = insertable.model_dump()
columns = ", ".join(values.keys())
placeholders = ", ".join(f":{k}" for k in values.keys())
last_record_id = await database.fetch_val(
query=f"INSERT INTO courts ({columns}) VALUES ({placeholders}) RETURNING id",
values=values,
)
return SingleCourtResponse(
data=assert_some(
await fetch_one_parsed(
database,
Court,
courts.select().where(
(courts.c.id == last_record_id) & (courts.c.tournament_id == tournament_id)
),
"SELECT * FROM courts WHERE id = :court_id AND tournament_id = :tournament_id",
{"court_id": last_record_id, "tournament_id": tournament_id},
)
)
)
14 changes: 6 additions & 8 deletions backend/bracket/routes/players.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@
SuccessResponse,
)
from bracket.routes.util import disallow_archived_tournament
from bracket.schema import players
from bracket.sql.players import (
get_all_players_in_tournament,
get_player_count,
Expand Down Expand Up @@ -54,20 +53,19 @@ async def update_player_by_id(
_: UserPublic = Depends(user_authenticated_for_tournament),
__: Tournament = Depends(disallow_archived_tournament),
) -> SinglePlayerResponse:
values = player_body.model_dump()
set_clause = ", ".join(f"{k} = :{k}" for k in values)
await database.execute(
query=players.update().where(
(players.c.id == player_id) & (players.c.tournament_id == tournament_id)
),
values=player_body.model_dump(),
query=f"UPDATE players SET {set_clause} WHERE id = :player_id AND tournament_id = :tournament_id",
values={**values, "player_id": player_id, "tournament_id": tournament_id},
)
return SinglePlayerResponse(
data=assert_some(
await fetch_one_parsed(
database,
Player,
players.select().where(
(players.c.id == player_id) & (players.c.tournament_id == tournament_id)
),
"SELECT * FROM players WHERE id = :player_id AND tournament_id = :tournament_id",
{"player_id": player_id, "tournament_id": tournament_id},
)
)
)
Expand Down
Loading
Loading