Coverage for slidge/db/store.py: 89%

477 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-09-29 05:05 +0000

1from __future__ import annotations 

2 

3import hashlib 

4import logging 

5import shutil 

6import uuid 

7from collections.abc import Callable, Collection, Iterable, Iterator 

8from datetime import UTC, datetime, timedelta 

9from mimetypes import guess_extension 

10from typing import Any, ClassVar, Protocol, TypeVar 

11 

12import sqlalchemy as sa 

13import sqlalchemy.orm 

14from slixmpp import Message 

15from slixmpp.exceptions import XMPPError 

16from slixmpp.plugins.xep_0231.stanza import BitsOfBinary 

17from slixmpp.xmlstream import tostring 

18from sqlalchemy import ColumnElement, Engine, delete, event, select, update 

19from sqlalchemy.exc import InvalidRequestError 

20from sqlalchemy.orm import ( 

21 Session, 

22 attributes, 

23 joinedload, 

24 load_only, 

25 selectinload, 

26 sessionmaker, 

27 with_loader_criteria, 

28) 

29 

30from ..core import config 

31from ..util.types import MamMetadata, Sticker 

32from .models import ( 

33 ArchivedMessage, 

34 ArchivedMessageSource, 

35 Attachment, 

36 Avatar, 

37 Bob, 

38 Contact, 

39 ContactSent, 

40 DirectMessages, 

41 DirectThreads, 

42 GatewayUser, 

43 Participant, 

44 Room, 

45 Space, 

46) 

47 

48 

49class UpdatableBase(Protocol): 

50 id: ClassVar[ColumnElement[int]] 

51 user_account_id: ClassVar[ColumnElement[int]] 

52 updated: ClassVar[ColumnElement[bool]] 

53 

54 

55T = TypeVar("T", bound=UpdatableBase) 

56 

57 

58class UpdatedMixin[T]: 

59 model: type[T] = NotImplemented 

60 

61 def __init__(self, session: Session) -> None: 

62 self.reset_updated(session) 

63 

64 def get_by_pk(self, session: Session, pk: int) -> T | None: 

65 stmt = select(self.model).where(self.model.id == pk) # type:ignore[attr-defined] 

66 return session.scalar(stmt) 

67 

68 def reset_updated(self, session: Session) -> None: 

69 session.execute(update(self.model).values(updated=False)) 

70 

71 def get_for(self, session: Session, user_pk: int) -> list[T]: 

72 stmt = select(self.model).where(self.model.user_account_id == user_pk) # type:ignore[attr-defined] 

73 return list(session.scalars(stmt)) 

74 

75 

76class SlidgeStore: 

77 def __init__(self, engine: Engine) -> None: 

78 self._engine = engine 

79 self.session = sessionmaker[Any](engine) 

80 

81 self.users = UserStore(self.session) 

82 self.avatars = AvatarStore(self.session) 

83 self.id_map = IdMapStore() 

84 self.bob = BobStore() 

85 self.attachments = AttachmentStore() 

86 with self.session() as session: 

87 self.contacts = ContactStore(session) 

88 self.mam = MAMStore(session, self.session) 

89 self.rooms = RoomStore(session) 

90 self.participants = ParticipantStore(session) 

91 self.spaces = SpaceStore(session) 

92 session.commit() 

93 

94 

95class UserStore: 

96 def __init__(self, session_maker: sessionmaker[Any]) -> None: 

97 self.session = session_maker 

98 

99 def update(self, user: GatewayUser) -> None: 

100 with self.session(expire_on_commit=False) as session: 

101 # https://github.com/sqlalchemy/sqlalchemy/discussions/6473 

102 try: 

103 attributes.flag_modified(user, "legacy_module_data") 

104 attributes.flag_modified(user, "preferences") 

105 except InvalidRequestError: 

106 pass 

107 session.add(user) 

108 session.commit() 

109 

110 

111class AvatarStore: 

112 def __init__(self, session_maker: sessionmaker[Any]) -> None: 

113 self.session = session_maker 

114 

115 

116LegacyToXmppType = type[DirectMessages] | type[DirectThreads] 

117 

118 

119class IdMapStore: 

120 @staticmethod 

121 def _set( 

122 session: Session, 

123 foreign_key: int, 

124 legacy_id: str, 

125 xmpp_ids: Iterable[str], 

126 type_: LegacyToXmppType, 

127 ) -> None: 

128 kwargs = {"foreign_key": foreign_key, "legacy_id": legacy_id} 

129 ids = list( 

130 session.scalars( 

131 select(type_.id).filter( 

132 type_.foreign_key == foreign_key, type_.legacy_id == legacy_id 

133 ) 

134 ) 

135 ) 

136 if ids: 

137 log.debug("Resetting legacy ID %s", legacy_id) 

138 session.execute(delete(type_).where(type_.id.in_(ids))) 

139 for xmpp_id in xmpp_ids: 

140 msg = type_(xmpp_id=xmpp_id, **kwargs) 

141 session.add(msg) 

142 

143 def set_thread( 

144 self, 

145 session: Session, 

146 foreign_key: int, 

147 legacy_id: str, 

148 xmpp_id: str, 

149 group: bool, 

150 ) -> None: 

151 if group: 

152 if xmpp_id == legacy_id: 

153 # No need to store in this case. The thread column of the mam 

154 # table is already populated, and if thread_legacy_id is NULL, 

155 # we assume xmpp_id == legacy_id. 

156 return 

157 session.execute( 

158 update(ArchivedMessage) 

159 .where( 

160 ArchivedMessage.room_id == foreign_key, 

161 ArchivedMessage.stanza_id == xmpp_id, 

162 ) 

163 .values(thread_legacy_id=legacy_id) 

164 ) 

165 else: 

166 self._set(session, foreign_key, legacy_id, [xmpp_id], DirectThreads) 

167 

168 def set_msg( 

169 self, 

170 session: Session, 

171 foreign_key: int, 

172 legacy_id: str, 

173 xmpp_ids: Iterable[str], 

174 group: bool, 

175 ) -> None: 

176 if group: 

177 session.execute( 

178 update(ArchivedMessage) 

179 .where( 

180 ArchivedMessage.room_id == foreign_key, 

181 ArchivedMessage.stanza_id.in_(xmpp_ids), 

182 ) 

183 .values(legacy_id=legacy_id) 

184 ) 

185 else: 

186 # In direct messages, we need to keep the mapping metween XMPP 

187 # client-generated msg-id and legacy ID, and we also need to know 

188 # which messages have been sent by the user. 

189 self._set(session, foreign_key, legacy_id, xmpp_ids, DirectMessages) 

190 

191 def set_origin( 

192 self, session: Session, foreign_key: int, legacy_id: str, xmpp_id: str 

193 ) -> None: 

194 session.execute( 

195 update(ArchivedMessage) 

196 .where( 

197 ArchivedMessage.room_id == foreign_key, 

198 ArchivedMessage.legacy_id == legacy_id, 

199 ) 

200 .values(origin_id=xmpp_id) 

201 ) 

202 

203 def get_origin( 

204 self, session: Session, foreign_key: int, legacy_id: str 

205 ) -> list[str]: 

206 # it's unlikely type checkers will ever be able to determine that this 

207 # is filtered out for NULL values 

208 return list( # ty:ignore[invalid-return-type] 

209 session.execute( # type:ignore[arg-type] 

210 select(ArchivedMessage.origin_id).where( 

211 ArchivedMessage.room_id == foreign_key, 

212 ArchivedMessage.legacy_id == legacy_id, 

213 ArchivedMessage.origin_id.is_not(None), 

214 ) 

215 ).scalars() 

216 ) 

217 

218 @staticmethod 

219 def _get( 

220 session: Session, foreign_key: int, legacy_id: str, type_: LegacyToXmppType 

221 ) -> list[str]: 

222 return list( 

223 session.scalars( 

224 select(type_.xmpp_id).filter_by( 

225 foreign_key=foreign_key, legacy_id=str(legacy_id) 

226 ) 

227 ) 

228 ) 

229 

230 def get_xmpp( 

231 self, session: Session, foreign_key: int, legacy_id: str, group: bool 

232 ) -> list[str]: 

233 if group: 

234 return list( 

235 session.execute( 

236 select(ArchivedMessage.stanza_id).where( 

237 ArchivedMessage.room_id == foreign_key, 

238 ArchivedMessage.legacy_id == legacy_id, 

239 ) 

240 ).scalars() 

241 ) 

242 else: 

243 return self._get(session, foreign_key, legacy_id, DirectMessages) 

244 

245 @staticmethod 

246 def _get_legacy( 

247 session: Session, foreign_key: int, xmpp_id: str, type_: LegacyToXmppType 

248 ) -> str | None: 

249 return session.scalar( 

250 select(type_.legacy_id).filter_by(foreign_key=foreign_key, xmpp_id=xmpp_id) 

251 ) 

252 

253 def get_legacy( 

254 self, 

255 session: Session, 

256 foreign_key: int, 

257 xmpp_id: str, 

258 group: bool, 

259 origin: bool = False, 

260 ) -> str | None: 

261 if group: 

262 cond = ( 

263 (ArchivedMessage.origin_id == xmpp_id) 

264 if origin 

265 else (ArchivedMessage.stanza_id == xmpp_id) 

266 ) 

267 return session.execute( 

268 select(ArchivedMessage.legacy_id).where( 

269 ArchivedMessage.room_id == foreign_key, 

270 cond, 

271 ) 

272 ).scalar_one_or_none() 

273 

274 return self._get_legacy(session, foreign_key, xmpp_id, DirectMessages) 

275 

276 def get_thread( 

277 self, session: Session, foreign_key: int, xmpp_id: str, group: bool 

278 ) -> str | None: 

279 if group: 

280 row = session.execute( 

281 select( 

282 ArchivedMessage.thread, 

283 ArchivedMessage.thread_legacy_id, 

284 ).where( 

285 ArchivedMessage.thread == xmpp_id, 

286 ) 

287 ).first() 

288 if row is None: 

289 return None 

290 xmpp_id, legacy_id = row 

291 return legacy_id or xmpp_id 

292 else: 

293 return self._get_legacy( 

294 session, 

295 foreign_key, 

296 xmpp_id, 

297 DirectThreads, 

298 ) 

299 

300 @staticmethod 

301 def was_sent_by_user(session: Session, foreign_key: int, legacy_id: str) -> bool: 

302 """ 

303 Only works for direct messages 

304 """ 

305 return ( 

306 session.scalar( 

307 select(DirectMessages.id).filter_by( 

308 foreign_key=foreign_key, legacy_id=legacy_id 

309 ) 

310 ) 

311 is not None 

312 ) 

313 

314 

315class ContactStore(UpdatedMixin[Contact]): 

316 model = Contact 

317 

318 def __init__(self, session: Session) -> None: 

319 super().__init__(session) 

320 session.execute(update(Contact).values(cached_presence=False)) 

321 

322 @staticmethod 

323 def add_to_sent(session: Session, contact_pk: int, msg_id: str) -> None: 

324 if ( 

325 session.query(ContactSent.id) 

326 .where(ContactSent.contact_id == contact_pk) 

327 .where(ContactSent.msg_id == msg_id) 

328 .first() 

329 ) is not None: 

330 log.warning("Contact %s has already sent message %s", contact_pk, msg_id) 

331 return 

332 new = ContactSent(contact_id=contact_pk, msg_id=msg_id) 

333 session.add(new) 

334 

335 @staticmethod 

336 def pop_sent_up_to(session: Session, contact_pk: int, msg_id: str) -> list[str]: 

337 result = [] 

338 to_del = [] 

339 for row in session.execute( 

340 select(ContactSent) 

341 .where(ContactSent.contact_id == contact_pk) 

342 .order_by(ContactSent.id) 

343 ).scalars(): 

344 to_del.append(row.id) 

345 result.append(row.msg_id) 

346 if row.msg_id == msg_id: 

347 break 

348 session.execute(delete(ContactSent).where(ContactSent.id.in_(to_del))) 

349 return result 

350 

351 

352class MAMStore: 

353 def __init__(self, session: Session, session_maker: sessionmaker[Any]) -> None: 

354 self.session = session_maker 

355 self.reset_source(session) 

356 

357 @staticmethod 

358 def reset_source(session: Session) -> None: 

359 session.execute( 

360 update(ArchivedMessage).values(source=ArchivedMessageSource.BACKFILL) 

361 ) 

362 

363 @staticmethod 

364 def nuke_older_than(session: Session, days: int) -> None: 

365 session.execute( 

366 delete(ArchivedMessage).where( 

367 ArchivedMessage.timestamp < datetime.now(tz=UTC) - timedelta(days=days) 

368 ) 

369 ) 

370 

371 @staticmethod 

372 def add_message( 

373 session: Session, 

374 room_pk: int, 

375 stanza: Message, 

376 archive_only: bool, 

377 legacy_msg_id: str | None, 

378 from_user: bool = False, 

379 ) -> ArchivedMessage: 

380 source = ( 

381 ArchivedMessageSource.BACKFILL 

382 if archive_only 

383 else ArchivedMessageSource.LIVE 

384 ) 

385 

386 when = stanza["delay"]["stamp"] or datetime.now(tz=UTC) 

387 stanza_id = stanza["stanza_id"]["id"] 

388 assert stanza_id 

389 

390 origin_id = stanza["origin_id"]["id"] or None 

391 author_resource = stanza.get_from().resource 

392 author_affiliation = stanza["muc"]["affiliation"] or None 

393 author_role = stanza["muc"]["role"] or None 

394 

395 if from_user: 

396 author_jid_localpart = None 

397 else: 

398 author_jid = stanza["muc"]["jid"] 

399 author_jid_localpart = author_jid.user if author_jid else None 

400 

401 occupant_id = stanza["occupant-id"]["id"] 

402 

403 displayed_id = stanza["displayed"]["id"] if "displayed" in stanza else None 

404 

405 if "reactions" in stanza: 

406 reaction_id = stanza["reactions"]["id"] 

407 if reaction_id is None: 

408 log.warning("Received a reaction without an ID?: %r", stanza) 

409 reaction_id = reaction_emojis = None 

410 else: 

411 # list() because this returns a set, which would require some more 

412 # SQLAlchemy plumbing. A list is fine! 

413 reaction_emojis = list(stanza["reactions"]["values"]) or None 

414 else: 

415 reaction_id = reaction_emojis = None 

416 

417 body = stanza["body"] 

418 thread = stanza["thread"] 

419 

420 del stanza["delay"] 

421 del stanza["markable"] 

422 del stanza["store"] 

423 del stanza["chat_state"] 

424 del stanza["muc"] 

425 del stanza["origin_id"] 

426 del stanza["stanza_id"] 

427 del stanza["occupant-id"] 

428 del stanza["displayed"] 

429 del stanza["reactions"] 

430 del stanza["body"] 

431 del stanza["thread"] 

432 

433 payload = _serialize_children(stanza) or None 

434 

435 existing = session.execute( 

436 select(ArchivedMessage).where( 

437 ArchivedMessage.room_id == room_pk, 

438 ArchivedMessage.stanza_id == stanza_id, 

439 ) 

440 ).scalar() 

441 if existing is None and legacy_msg_id is not None: 

442 existing = session.execute( 

443 select(ArchivedMessage).where( 

444 ArchivedMessage.room_id == room_pk, 

445 ArchivedMessage.legacy_id == legacy_msg_id, 

446 ) 

447 ).scalar() 

448 elif payload is None and displayed_id is not None and occupant_id is not None: 

449 existing_displayed_marker = session.scalar( 

450 select(ArchivedMessage).where( 

451 ArchivedMessage.room_id == room_pk, 

452 ArchivedMessage.occupant_id == occupant_id, 

453 ArchivedMessage.displayed_id == displayed_id, 

454 ) 

455 ) 

456 if existing_displayed_marker is not None: 

457 log.debug("Received a duplicate displayed marker, ignoring it") 

458 return existing_displayed_marker 

459 

460 if existing is None: 

461 mam_msg = ArchivedMessage(stanza_id=stanza_id, room_id=room_pk) 

462 else: 

463 log.debug("Updating message %s in room %s", stanza_id, room_pk) 

464 mam_msg = existing 

465 

466 mam_msg.stanza_id = stanza_id 

467 mam_msg.timestamp = when 

468 mam_msg.payload = payload 

469 mam_msg.origin_id = origin_id 

470 mam_msg.author_resource = author_resource 

471 mam_msg.author_affiliation = author_affiliation 

472 mam_msg.author_role = author_role 

473 mam_msg.author_jid_localpart = author_jid_localpart 

474 mam_msg.room_id = room_pk 

475 mam_msg.source = source 

476 mam_msg.legacy_id = legacy_msg_id 

477 mam_msg.occupant_id = occupant_id 

478 mam_msg.from_user = from_user 

479 mam_msg.displayed_id = displayed_id 

480 mam_msg.reaction_emojis = reaction_emojis 

481 mam_msg.reaction_id = reaction_id 

482 mam_msg.body = body 

483 mam_msg.thread = thread 

484 

485 session.add(mam_msg) 

486 

487 return mam_msg 

488 

489 @staticmethod 

490 def get_messages( 

491 session: Session, 

492 room_pk: int, 

493 start_date: datetime | None = None, 

494 end_date: datetime | None = None, 

495 before_id: str | None = None, 

496 after_id: str | None = None, 

497 ids: Collection[str] = (), 

498 last_page_n: int | None = None, 

499 author_resource: str | None = None, 

500 flip: bool = False, 

501 ) -> Iterator[ArchivedMessage]: 

502 q = select(ArchivedMessage).where(ArchivedMessage.room_id == room_pk) 

503 if start_date is not None: 

504 q = q.where(ArchivedMessage.timestamp >= start_date) 

505 if end_date is not None: 

506 q = q.where(ArchivedMessage.timestamp <= end_date) 

507 if before_id is not None: 

508 stamp = session.execute( 

509 select(ArchivedMessage.timestamp).where( 

510 ArchivedMessage.stanza_id == before_id, 

511 ArchivedMessage.room_id == room_pk, 

512 ) 

513 ).scalar_one_or_none() 

514 if stamp is None: 

515 raise XMPPError( 

516 "item-not-found", 

517 f"Message {before_id} not found", 

518 ) 

519 q = q.where(ArchivedMessage.timestamp < stamp) 

520 if after_id is not None: 

521 stamp = session.execute( 

522 select(ArchivedMessage.timestamp).where( 

523 ArchivedMessage.stanza_id == after_id, 

524 ArchivedMessage.room_id == room_pk, 

525 ) 

526 ).scalar_one_or_none() 

527 if stamp is None: 

528 raise XMPPError( 

529 "item-not-found", 

530 f"Message {after_id} not found", 

531 ) 

532 q = q.where(ArchivedMessage.timestamp > stamp) 

533 if ids: 

534 q = q.filter(ArchivedMessage.stanza_id.in_(ids)) 

535 if author_resource is not None: 

536 q = q.where(ArchivedMessage.author_resource == author_resource) 

537 if flip: 

538 q = q.order_by(ArchivedMessage.timestamp.desc()) 

539 else: 

540 q = q.order_by(ArchivedMessage.timestamp.asc()) 

541 msgs = list(session.execute(q).scalars()) 

542 if ids and len(msgs) != len(ids): 

543 raise XMPPError( 

544 "item-not-found", 

545 "One of the requested messages IDs could not be found " 

546 "with the given constraints.", 

547 ) 

548 if last_page_n is not None: 

549 msgs = msgs[:last_page_n] if flip else msgs[-last_page_n:] 

550 yield from msgs 

551 

552 @staticmethod 

553 def get_first( 

554 session: Session, room_pk: int, with_legacy_id: bool = False 

555 ) -> ArchivedMessage | None: 

556 q = ( 

557 select(ArchivedMessage) 

558 .where(ArchivedMessage.room_id == room_pk) 

559 .order_by(ArchivedMessage.timestamp.asc()) 

560 ) 

561 if with_legacy_id: 

562 q = q.filter(ArchivedMessage.legacy_id.isnot(None)) 

563 return session.execute(q).scalar() 

564 

565 @staticmethod 

566 def get_last( 

567 session: Session, room_pk: int, source: ArchivedMessageSource | None = None 

568 ) -> ArchivedMessage | None: 

569 q = select(ArchivedMessage).where(ArchivedMessage.room_id == room_pk) 

570 

571 if source is not None: 

572 q = q.where(ArchivedMessage.source == source) 

573 

574 return session.execute(q.order_by(ArchivedMessage.timestamp.desc())).scalar() 

575 

576 def get_first_and_last(self, session: Session, room_pk: int) -> list[MamMetadata]: 

577 r = [] 

578 first = self.get_first(session, room_pk) 

579 if first is not None: 

580 r.append(MamMetadata(first.stanza_id, first.timestamp)) 

581 last = self.get_last(session, room_pk) 

582 if last is not None: 

583 r.append(MamMetadata(last.stanza_id, last.timestamp)) 

584 return r 

585 

586 @staticmethod 

587 def get_most_recent_with_legacy_id( 

588 session: Session, room_pk: int, source: ArchivedMessageSource | None = None 

589 ) -> ArchivedMessage | None: 

590 q = ( 

591 select(ArchivedMessage) 

592 .where(ArchivedMessage.room_id == room_pk) 

593 .where(ArchivedMessage.legacy_id.isnot(None)) 

594 ) 

595 if source is not None: 

596 q = q.where(ArchivedMessage.source == source) 

597 return session.execute(q.order_by(ArchivedMessage.timestamp.desc())).scalar() 

598 

599 @staticmethod 

600 def get_least_recent_with_legacy_id_after( 

601 session: Session, 

602 room_pk: int, 

603 after_id: str, 

604 source: ArchivedMessageSource = ArchivedMessageSource.LIVE, 

605 ) -> ArchivedMessage | None: 

606 return session.scalar( 

607 select(ArchivedMessage) 

608 .where( 

609 ArchivedMessage.room_id == room_pk, 

610 ArchivedMessage.legacy_id.isnot(None), 

611 ArchivedMessage.source == source, 

612 ArchivedMessage.timestamp 

613 > ( 

614 select(sa.func.max(ArchivedMessage.timestamp)) 

615 .where( 

616 ArchivedMessage.room_id == room_pk, 

617 ArchivedMessage.legacy_id == after_id, 

618 ) 

619 .scalar_subquery() 

620 ), 

621 ) 

622 .order_by(ArchivedMessage.timestamp.asc(), ArchivedMessage.id.desc()) 

623 .limit(1) 

624 ) 

625 

626 @staticmethod 

627 def get_by_legacy_id( 

628 session: Session, room_pk: int, legacy_id: str 

629 ) -> ArchivedMessage | None: 

630 return ( 

631 session.query(ArchivedMessage) 

632 .filter(ArchivedMessage.room_id == room_pk) 

633 .filter(ArchivedMessage.legacy_id == legacy_id) 

634 .first() 

635 ) 

636 

637 @staticmethod 

638 def pop_unread_up_to(session: Session, room_pk: int, stanza_id: str) -> list[str]: 

639 q = ( 

640 select(ArchivedMessage.id, ArchivedMessage.stanza_id) 

641 .where(ArchivedMessage.room_id == room_pk) 

642 .where(~ArchivedMessage.displayed_by_user) 

643 .where(ArchivedMessage.legacy_id.is_not(None)) 

644 .order_by(ArchivedMessage.timestamp.asc()) 

645 ) 

646 

647 ref = session.scalar( 

648 select(ArchivedMessage) 

649 .where(ArchivedMessage.room_id == room_pk) 

650 .where(ArchivedMessage.stanza_id == stanza_id) 

651 ) 

652 

653 if ref is None: 

654 log.debug( 

655 "(pop unread in muc): message not found, returning all MAM messages." 

656 ) 

657 rows = session.execute(q) 

658 else: 

659 rows = session.execute(q.where(ArchivedMessage.timestamp <= ref.timestamp)) 

660 

661 pks: list[int] = [] 

662 stanza_ids: list[str] = [] 

663 

664 for id_, sid in rows: 

665 pks.append(id_) 

666 stanza_ids.append(sid) 

667 

668 session.execute( 

669 update(ArchivedMessage) 

670 .where(ArchivedMessage.id.in_(pks)) 

671 .values(displayed_by_user=True) 

672 ) 

673 return stanza_ids 

674 

675 @staticmethod 

676 def is_displayed_by_user( 

677 session: Session, room_jid_localpart: str, legacy_msg_id: str 

678 ) -> bool: 

679 return any( 

680 session.execute( 

681 select(ArchivedMessage.displayed_by_user) 

682 .join(Room) 

683 .where(Room.jid_localpart == room_jid_localpart) 

684 .where(ArchivedMessage.legacy_id == legacy_msg_id) 

685 ).scalars() 

686 ) 

687 

688 @staticmethod 

689 def get_body(session: Session, room_pk: int, legacy_msg_id: str) -> str | None: 

690 return session.scalar( 

691 select(ArchivedMessage.body).where( 

692 ArchivedMessage.room_id == room_pk, 

693 ArchivedMessage.legacy_id == legacy_msg_id, 

694 ) 

695 ) 

696 

697 

698class RoomStore(UpdatedMixin[Room]): 

699 model = Room 

700 

701 def reset_updated(self, session: Session) -> None: 

702 super().reset_updated(session) 

703 session.execute( 

704 update(Room).values( 

705 subject_setter=None, 

706 user_resources=None, 

707 history_filled=False, 

708 participants_filled=False, 

709 ) 

710 ) 

711 

712 @staticmethod 

713 def get_all(session: Session, user_pk: int) -> Iterator[Room]: 

714 yield from session.scalars(select(Room).where(Room.user_account_id == user_pk)) 

715 

716 @staticmethod 

717 def get(session: Session, user_pk: int, legacy_id: str) -> Room: 

718 return session.execute( 

719 select(Room) 

720 .where(Room.user_account_id == user_pk) 

721 .where(Room.legacy_id == legacy_id) 

722 ).scalar_one() 

723 

724 @staticmethod 

725 def nick_available(session: Session, room_pk: int, nickname: str) -> bool: 

726 return ( 

727 session.execute( 

728 select(Participant.id).filter_by(room_id=room_pk, nickname=nickname) 

729 ) 

730 ).one_or_none() is None 

731 

732 

733class ParticipantStore: 

734 def __init__(self, session: Session) -> None: 

735 session.execute(delete(Participant)) 

736 

737 @staticmethod 

738 def get_all( 

739 session: Session, room_pk: int, user_included: bool = True 

740 ) -> Iterator[Participant]: 

741 query = select(Participant).where(Participant.room_id == room_pk) 

742 if not user_included: 

743 query = query.where(~Participant.is_user) 

744 yield from session.scalars(query).unique() 

745 

746 @staticmethod 

747 def delete(session: Session, pk: int) -> None: 

748 session.execute(delete(Participant).where(Participant.id == pk)) 

749 

750 

751class BobStore: 

752 _ATTR_MAP: ClassVar[dict[str, str]] = { 

753 "sha-1": "sha_1", 

754 "sha1": "sha_1", 

755 "sha-256": "sha_256", 

756 "sha256": "sha_256", 

757 "sha-512": "sha_512", 

758 "sha512": "sha_512", 

759 } 

760 

761 _ALG_MAP: ClassVar[dict[str, Callable[[bytes], hashlib._Hash]]] = { 

762 "sha_1": hashlib.sha1, 

763 "sha_256": hashlib.sha256, 

764 "sha_512": hashlib.sha512, 

765 } 

766 

767 def __init__(self) -> None: 

768 if (config.HOME_DIR / "slidge_stickers").exists(): 

769 shutil.move( 

770 config.HOME_DIR / "slidge_stickers", config.HOME_DIR / "bob_store" 

771 ) 

772 self.root_dir = config.HOME_DIR / "bob_store" 

773 self.root_dir.mkdir(exist_ok=True) 

774 

775 @staticmethod 

776 def __split_cid(cid: str) -> list[str]: 

777 return cid.removesuffix("@bob.xmpp.org").split("+") 

778 

779 def __get_condition(self, cid: str) -> ColumnElement[bool]: 

780 alg_name, digest = self.__split_cid(cid) 

781 attr = self._ATTR_MAP.get(alg_name) 

782 if attr is None: 

783 log.warning("Unknown hash algorithm: %s", alg_name) 

784 raise ValueError 

785 return getattr(Bob, attr) == digest # type:ignore[no-any-return] 

786 

787 def get(self, session: Session, cid: str) -> Bob | None: 

788 try: 

789 return session.query(Bob).filter(self.__get_condition(cid)).scalar() # type:ignore[no-any-return] 

790 except ValueError: 

791 log.warning("Cannot get Bob with CID: %s", cid) 

792 return None 

793 

794 def get_sticker(self, session: Session, cid: str) -> Sticker | None: 

795 bob = self.get(session, cid) 

796 if bob is None: 

797 return None 

798 return self.__sticker_from_bob(bob) 

799 

800 def __sticker_from_bob(self, bob: Bob) -> Sticker: 

801 return Sticker( 

802 self.root_dir / bob.file_name, 

803 bob.content_type, 

804 {h: getattr(bob, h) for h in self._ALG_MAP}, 

805 ) 

806 

807 def get_bob( 

808 self, session: Session, _jid: object, _node: object, _ifrom: object, cid: str 

809 ) -> BitsOfBinary | None: 

810 stored = self.get(session, cid) 

811 if stored is None: 

812 return None 

813 bob = BitsOfBinary() 

814 bob["data"] = (self.root_dir / stored.file_name).read_bytes() 

815 if stored.content_type is not None: 

816 bob["type"] = stored.content_type 

817 bob["cid"] = cid 

818 return bob 

819 

820 def del_bob( 

821 self, session: Session, _jid: object, _node: object, _ifrom: object, cid: str 

822 ) -> None: 

823 try: 

824 file_name = session.scalar( 

825 delete(Bob).where(self.__get_condition(cid)).returning(Bob.file_name) 

826 ) 

827 except ValueError: 

828 log.warning("Cannot delete Bob with CID: %s", cid) 

829 return 

830 if file_name is None: 

831 log.warning("No BoB with CID: %s", cid) 

832 return 

833 (self.root_dir / file_name).unlink() 

834 

835 def set_bob( 

836 self, 

837 session: Session, 

838 _jid: object, 

839 _node: object, 

840 _ifrom: object, 

841 bob: BitsOfBinary, 

842 ) -> Sticker | None: 

843 return self.set_sticker(session, bob["cid"], bob["data"], bob["type"]) 

844 

845 def set_sticker( 

846 self, 

847 session: Session, 

848 cid: str, 

849 bytes_: bytes, 

850 content_type: str | None, 

851 ) -> Sticker | None: 

852 try: 

853 alg_name, digest = self.__split_cid(cid) 

854 except ValueError: 

855 log.warning("Invalid CID provided: %s", cid) 

856 return None 

857 attr = self._ATTR_MAP.get(alg_name) 

858 if attr is None: 

859 log.warning("Cannot set Bob: Unknown algorithm type: %s", alg_name) 

860 return None 

861 existing = self.get(session, cid) 

862 if existing: 

863 log.debug("Bob already exists") 

864 return None 

865 path = self.root_dir / uuid.uuid4().hex 

866 if content_type is None: 

867 try: 

868 import magic 

869 except ImportError: 

870 content_type = "application/octet-stream" 

871 else: 

872 content_type = magic.from_buffer(bytes_, mime=True) 

873 path = path.with_suffix(guess_extension(content_type) or "") 

874 path.write_bytes(bytes_) 

875 hashes = {k: v(bytes_).hexdigest() for k, v in self._ALG_MAP.items()} 

876 if hashes[attr] != digest: 

877 path.unlink(missing_ok=True) 

878 raise ValueError("Provided CID does not match calculated hash") 

879 row = Bob(file_name=path.name, content_type=content_type, **hashes) 

880 session.add(row) 

881 return self.__sticker_from_bob(row) 

882 

883 

884class SpaceStore(UpdatedMixin[Space]): 

885 model = Space 

886 

887 def __init__(self, session: Session) -> None: 

888 session.execute(delete(Space)) 

889 

890 @staticmethod 

891 def add_or_get(session: Session, user_pk: int, legacy_id: str) -> Space: 

892 space = session.execute( 

893 select(Space) 

894 .where(Space.user_account_id == user_pk) 

895 .where(Space.legacy_id == legacy_id) 

896 .options( 

897 joinedload(Space.avatar), 

898 joinedload(Space.banner), 

899 ) 

900 ).scalar_one_or_none() 

901 if space is None: 

902 space = Space( 

903 user_account_id=user_pk, 

904 legacy_id=legacy_id, 

905 avatar=None, 

906 banner=None, 

907 rooms=[], 

908 ) 

909 session.add(space) 

910 session.commit() 

911 return space 

912 

913 @staticmethod 

914 def get_all(session: Session, user_pk: int) -> Iterable[Space]: 

915 return session.execute( 

916 select(Space).where(Space.user_account_id == user_pk) 

917 ).scalars() 

918 

919 @staticmethod 

920 def get_by_legacy_id( 

921 session: Session, 

922 user_pk: int, 

923 legacy_id: str, 

924 affiliations: bool = False, 

925 images: bool = False, 

926 room_legacy_id_filter: Iterable[str] | None = None, 

927 ) -> Space | None: 

928 stmt = ( 

929 select(Space) 

930 .where(Space.user_account_id == user_pk) 

931 .where(Space.legacy_id == legacy_id) 

932 ) 

933 if affiliations: 

934 stmt = stmt.options( 

935 joinedload(Space.owners), 

936 joinedload(Space.creator), 

937 ) 

938 if images: 

939 stmt = stmt.options( 

940 joinedload(Space.avatar), 

941 joinedload(Space.banner), 

942 ) 

943 if room_legacy_id_filter is not None: 

944 stmt = stmt.options( 

945 selectinload(Space.rooms), 

946 with_loader_criteria( 

947 Room, 

948 Room.legacy_id.in_(room_legacy_id_filter), 

949 ), 

950 ) 

951 

952 return session.execute(stmt).unique().scalar_one_or_none() 

953 

954 @staticmethod 

955 def get_unupdated(session: Session, user_pk: int) -> list[Space]: 

956 return list( 

957 session.execute( 

958 select(Space) 

959 .where(Space.user_account_id == user_pk) 

960 .where(Space.updated.is_(False)) 

961 .options( 

962 joinedload(Space.avatar), 

963 joinedload(Space.banner), 

964 ) 

965 ).scalars() 

966 ) 

967 

968 @staticmethod 

969 def get_rooms( 

970 session: Session, 

971 user_pk: int, 

972 legacy_id: str, 

973 room_legacy_ids: Iterable[str] = (), 

974 ) -> list[Room]: 

975 q = ( 

976 select(Room) 

977 .join(Room.space) 

978 .where(Room.user_account_id == user_pk) 

979 .where(Space.legacy_id == legacy_id) 

980 .options(load_only(Room.jid_localpart, Room.name)) 

981 ) 

982 if room_legacy_ids: 

983 q = q.where(Room.legacy_id.in_(room_legacy_ids)) 

984 return list(session.execute(q).scalars()) 

985 

986 @staticmethod 

987 def exists(session: Session, user_pk: int, legacy_id: str) -> bool: 

988 return session.execute( 

989 select( 

990 sa.exists() 

991 .where(Space.user_account_id == user_pk) 

992 .where(Space.legacy_id == legacy_id) 

993 ) 

994 ).scalar_one() 

995 

996 

997class AttachmentStore: 

998 @staticmethod 

999 def get_all(session: sa.orm.Session) -> list[Attachment]: 

1000 return list( 

1001 session.execute( 

1002 select(Attachment).options(load_only(Attachment.id, Attachment.url)) 

1003 ).scalars() 

1004 ) 

1005 

1006 @staticmethod 

1007 def remove(session: sa.orm.Session, pks: list[int]) -> None: 

1008 session.execute(delete(Attachment).where(Attachment.id.in_(pks))) 

1009 

1010 

1011@event.listens_for(sa.orm.Session, "after_flush") 

1012def _check_avatar_orphans(session: Session, flush_context: sa.ExecutionContext) -> None: 

1013 if not session.deleted: 

1014 return 

1015 

1016 potentially_orphaned = set() 

1017 for obj in session.deleted: 

1018 if isinstance(obj, (Contact, Room)) and obj.avatar_id: 

1019 potentially_orphaned.add(obj.avatar_id) 

1020 if not potentially_orphaned: 

1021 return 

1022 

1023 result = session.execute( 

1024 sa.delete(Avatar).where( 

1025 sa.and_( 

1026 Avatar.id.in_(potentially_orphaned), 

1027 sa.not_(sa.exists().where(Contact.avatar_id == Avatar.id)), 

1028 sa.not_(sa.exists().where(Room.avatar_id == Avatar.id)), 

1029 sa.not_(sa.exists().where(Space.avatar_id == Avatar.id)), 

1030 sa.not_(sa.exists().where(Space.banner_id == Avatar.id)), 

1031 ) 

1032 ) 

1033 ) 

1034 deleted_count = result.rowcount # type:ignore[attr-defined] 

1035 log.debug("Auto-deleted %s orphaned avatars", deleted_count) 

1036 

1037 

1038def _serialize_children(msg: Message) -> str: 

1039 return "".join(tostring(child, xmlns=msg.namespace) for child in msg.xml) 

1040 

1041 

1042log = logging.getLogger(__name__)