Coverage for slidge/db/alembic/versions/f91200fe83e4_mam_body_and_thread_columns.py: 54%
52 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-29 05:05 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-29 05:05 +0000
1"""mam: body and thread columns
3Revision ID: f91200fe83e4
4Revises: 2d5d506dcbbb
5Create Date: 2026-09-15 19:59:56.769232
7"""
9import logging
10from collections.abc import Sequence
11from xml.etree import ElementTree as ET
13import sqlalchemy as sa
14from alembic import op
15from slixmpp import Message
17from slidge.db.store import _serialize_children
19# revision identifiers, used by Alembic.
20revision: str = "f91200fe83e4"
21down_revision: str | None = "2d5d506dcbbb"
22branch_labels: str | Sequence[str] | None = None
23depends_on: str | Sequence[str] | None = None
26def upgrade() -> None:
27 with op.batch_alter_table("mam", schema=None) as batch_op:
28 batch_op.add_column(sa.Column("body", sa.String(), nullable=True))
29 batch_op.add_column(sa.Column("thread", sa.String(), nullable=True))
31 conn = op.get_bind()
32 q = sa.select(mam_tbl.c.id, mam_tbl.c.payload).where(mam_tbl.c.payload.isnot(None))
34 update_stmt = mam_tbl.update().where(mam_tbl.c.id == sa.bindparam("_id"))
36 count_q = sa.select(sa.func.count()).select_from(q.subquery())
37 total = conn.execute(count_q).scalar_one()
39 log.info("Migrating %s rows...", total)
41 result = conn.execute(q, execution_options={"yield_per": 1000})
43 n = 0
44 for i, batch in enumerate(result.partitions()):
45 log.debug("Batch %s", i)
46 params = []
47 for row in batch:
48 mam_id, payload = row
49 msg = Message()
50 try:
51 xml = ET.fromstring(f"<x xmlns='{msg.namespace}'>{payload}</x>")
52 except Exception:
53 log.exception("Skipping row %s", mam_id)
54 continue
55 for child in xml:
56 msg.append(child)
57 thread = msg["thread"] or None
58 body = msg["body"] or None
59 del msg["thread"]
60 del msg["body"]
61 payload = _serialize_children(msg) or None
62 if not (body or thread):
63 continue
64 n += 1
65 params.append(
66 {
67 "_id": mam_id,
68 "thread": thread,
69 "body": body,
70 "payload": payload,
71 }
72 )
74 if params:
75 conn.execute(update_stmt, params)
77 log.info("Optimized storage of %s rows...", n)
80def downgrade() -> None:
81 raise RuntimeError("downgrades are not supported")
84mam_tbl = sa.table(
85 "mam",
86 sa.column("id", sa.Integer),
87 sa.column("payload", sa.String),
88 sa.column("body", sa.String),
89 sa.column("thread", sa.String),
90)
92log = logging.getLogger("MAM body/thread migrations")