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

1"""mam: body and thread columns 

2 

3Revision ID: f91200fe83e4 

4Revises: 2d5d506dcbbb 

5Create Date: 2026-09-15 19:59:56.769232 

6 

7""" 

8 

9import logging 

10from collections.abc import Sequence 

11from xml.etree import ElementTree as ET 

12 

13import sqlalchemy as sa 

14from alembic import op 

15from slixmpp import Message 

16 

17from slidge.db.store import _serialize_children 

18 

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 

24 

25 

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)) 

30 

31 conn = op.get_bind() 

32 q = sa.select(mam_tbl.c.id, mam_tbl.c.payload).where(mam_tbl.c.payload.isnot(None)) 

33 

34 update_stmt = mam_tbl.update().where(mam_tbl.c.id == sa.bindparam("_id")) 

35 

36 count_q = sa.select(sa.func.count()).select_from(q.subquery()) 

37 total = conn.execute(count_q).scalar_one() 

38 

39 log.info("Migrating %s rows...", total) 

40 

41 result = conn.execute(q, execution_options={"yield_per": 1000}) 

42 

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 ) 

73 

74 if params: 

75 conn.execute(update_stmt, params) 

76 

77 log.info("Optimized storage of %s rows...", n) 

78 

79 

80def downgrade() -> None: 

81 raise RuntimeError("downgrades are not supported") 

82 

83 

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) 

91 

92log = logging.getLogger("MAM body/thread migrations")