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
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-29 05:05 +0000
1from __future__ import annotations
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
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)
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)
49class UpdatableBase(Protocol):
50 id: ClassVar[ColumnElement[int]]
51 user_account_id: ClassVar[ColumnElement[int]]
52 updated: ClassVar[ColumnElement[bool]]
55T = TypeVar("T", bound=UpdatableBase)
58class UpdatedMixin[T]:
59 model: type[T] = NotImplemented
61 def __init__(self, session: Session) -> None:
62 self.reset_updated(session)
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)
68 def reset_updated(self, session: Session) -> None:
69 session.execute(update(self.model).values(updated=False))
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))
76class SlidgeStore:
77 def __init__(self, engine: Engine) -> None:
78 self._engine = engine
79 self.session = sessionmaker[Any](engine)
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()
95class UserStore:
96 def __init__(self, session_maker: sessionmaker[Any]) -> None:
97 self.session = session_maker
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()
111class AvatarStore:
112 def __init__(self, session_maker: sessionmaker[Any]) -> None:
113 self.session = session_maker
116LegacyToXmppType = type[DirectMessages] | type[DirectThreads]
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)
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)
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)
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 )
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 )
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 )
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)
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 )
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()
274 return self._get_legacy(session, foreign_key, xmpp_id, DirectMessages)
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 )
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 )
315class ContactStore(UpdatedMixin[Contact]):
316 model = Contact
318 def __init__(self, session: Session) -> None:
319 super().__init__(session)
320 session.execute(update(Contact).values(cached_presence=False))
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)
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
352class MAMStore:
353 def __init__(self, session: Session, session_maker: sessionmaker[Any]) -> None:
354 self.session = session_maker
355 self.reset_source(session)
357 @staticmethod
358 def reset_source(session: Session) -> None:
359 session.execute(
360 update(ArchivedMessage).values(source=ArchivedMessageSource.BACKFILL)
361 )
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 )
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 )
386 when = stanza["delay"]["stamp"] or datetime.now(tz=UTC)
387 stanza_id = stanza["stanza_id"]["id"]
388 assert stanza_id
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
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
401 occupant_id = stanza["occupant-id"]["id"]
403 displayed_id = stanza["displayed"]["id"] if "displayed" in stanza else None
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
417 body = stanza["body"]
418 thread = stanza["thread"]
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"]
433 payload = _serialize_children(stanza) or None
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
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
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
485 session.add(mam_msg)
487 return mam_msg
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
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()
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)
571 if source is not None:
572 q = q.where(ArchivedMessage.source == source)
574 return session.execute(q.order_by(ArchivedMessage.timestamp.desc())).scalar()
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
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()
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 )
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 )
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 )
647 ref = session.scalar(
648 select(ArchivedMessage)
649 .where(ArchivedMessage.room_id == room_pk)
650 .where(ArchivedMessage.stanza_id == stanza_id)
651 )
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))
661 pks: list[int] = []
662 stanza_ids: list[str] = []
664 for id_, sid in rows:
665 pks.append(id_)
666 stanza_ids.append(sid)
668 session.execute(
669 update(ArchivedMessage)
670 .where(ArchivedMessage.id.in_(pks))
671 .values(displayed_by_user=True)
672 )
673 return stanza_ids
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 )
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 )
698class RoomStore(UpdatedMixin[Room]):
699 model = Room
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 )
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))
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()
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
733class ParticipantStore:
734 def __init__(self, session: Session) -> None:
735 session.execute(delete(Participant))
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()
746 @staticmethod
747 def delete(session: Session, pk: int) -> None:
748 session.execute(delete(Participant).where(Participant.id == pk))
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 }
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 }
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)
775 @staticmethod
776 def __split_cid(cid: str) -> list[str]:
777 return cid.removesuffix("@bob.xmpp.org").split("+")
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]
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
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)
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 )
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
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()
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"])
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)
884class SpaceStore(UpdatedMixin[Space]):
885 model = Space
887 def __init__(self, session: Session) -> None:
888 session.execute(delete(Space))
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
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()
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 )
952 return session.execute(stmt).unique().scalar_one_or_none()
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 )
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())
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()
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 )
1006 @staticmethod
1007 def remove(session: sa.orm.Session, pks: list[int]) -> None:
1008 session.execute(delete(Attachment).where(Attachment.id.in_(pks)))
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
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
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)
1038def _serialize_children(msg: Message) -> str:
1039 return "".join(tostring(child, xmlns=msg.namespace) for child in msg.xml)
1042log = logging.getLogger(__name__)