Coverage for slidge/core/session.py: 82%

240 statements  

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

1import abc 

2import asyncio 

3import contextlib 

4import logging 

5import warnings 

6from asyncio.tasks import Task 

7from collections.abc import Coroutine 

8from typing import Any, Final, Generic, NamedTuple, Self 

9 

10import aiohttp 

11import sqlalchemy as sa 

12from slixmpp import JID, Iq, Message, Presence 

13from slixmpp.exceptions import XMPPError 

14from slixmpp.types import PresenceShows, ResourceDict 

15 

16from slidge.db.meta import JSONSerializable 

17 

18from ..command import SearchResult 

19from ..contact import LegacyContact, LegacyRoster 

20from ..db.models import Contact, GatewayUser 

21from ..group import LegacyBookmarks 

22from ..util.lock import NamedLockMixin 

23from ..util.types import ( 

24 AnyGateway, 

25 AnyMUC, 

26 AnyParticipant, 

27 AnySession, 

28 LegacyBookmarksType_co, 

29 LegacyRosterType_co, 

30 PseudoPresenceShow, 

31) 

32from ..util.util import derive_wired_class, noop_coro 

33 

34 

35class CachedPresence(NamedTuple): 

36 status: str | None 

37 show: str | None 

38 kwargs: dict[str, Any] 

39 

40 

41class BaseSession( 

42 NamedLockMixin, 

43 abc.ABC, 

44 Generic[LegacyRosterType_co, LegacyBookmarksType_co], # noqa: UP046 (mypy cannot infer variance with PEP 695 syntax) 

45): 

46 """ 

47 The session of a registered :term:`User`. 

48 

49 Represents a gateway user logged in to the legacy network and performing actions. 

50 

51 Will be instantiated automatically on slidge startup for each registered user, 

52 or upon registration for new (validated) users. 

53 

54 Must be subclassed for a functional :term:`Legacy Module`. 

55 """ 

56 

57 """ 

58 Since we cannot set the XMPP ID of messages sent by XMPP clients, we need to keep a mapping 

59 between XMPP IDs and legacy message IDs if we want to further refer to a message that was sent 

60 by the user. This also applies to 'carboned' messages, ie, messages sent by the user from 

61 the official client of a legacy network. 

62 """ 

63 

64 xmpp: AnyGateway 

65 """ 

66 The gateway instance singleton. Use it for low-level XMPP calls or custom methods that are not 

67 session-specific. 

68 

69 It is set on the session class by the gateway on startup, before any 

70 session is instantiated. Plugins may redeclare it with their own gateway 

71 class for typed access, e.g. ``xmpp: "Gateway"``. 

72 """ 

73 

74 MESSAGE_IDS_ARE_THREAD_IDS = False 

75 """ 

76 Set this to True if the legacy service uses message IDs as thread IDs, 

77 eg Mattermost, where you can only 'create a thread' by replying to the message, 

78 in which case the message ID is also a thread ID (and all messages are potential 

79 threads). 

80 """ 

81 

82 roster_cls: type[LegacyRosterType_co] 

83 """ 

84 The :class:`.LegacyRoster` subclass to use for this session's contacts. 

85 

86 Derived automatically from the first generic parameter, e.g., 

87 ``class Session(BaseSession[Roster, Bookmarks])``, which also types 

88 :attr:`.contacts`. 

89 """ 

90 bookmarks_cls: type[LegacyBookmarksType_co] 

91 """ 

92 The :class:`.LegacyBookmarks` subclass to use for this session's groups. 

93 

94 Derived automatically from the second generic parameter, e.g., 

95 ``class Session(BaseSession[Roster, Bookmarks])``, which also types 

96 :attr:`.bookmarks`. 

97 """ 

98 

99 def __init_subclass__(cls, **kwargs: object) -> None: 

100 super().__init_subclass__(**kwargs) 

101 derive_wired_class(cls, BaseSession, "roster_cls", "bookmarks_cls") 

102 

103 def __init__(self, user: GatewayUser) -> None: 

104 super().__init__() 

105 if self.MESSAGE_IDS_ARE_THREAD_IDS and not self.xmpp.THREADS: 

106 raise RuntimeError( 

107 "Threads cannot be disabled if Session.MESSAGE_IDS_ARE_THREAD_IDS == True" 

108 ) 

109 self.user = user 

110 """ 

111 The :term:`slidge user <User>`. 

112 """ 

113 self.log = logging.getLogger(user.jid.bare) 

114 

115 self.ignore_messages = set[str]() 

116 

117 self.contacts: Final[LegacyRosterType_co] = self.roster_cls(self) 

118 """This session's roster, an instance of :attr:`.roster_cls`.""" 

119 

120 self.is_logging_in = False 

121 self._logged = False 

122 self.__reset_ready() 

123 

124 self.bookmarks: Final[LegacyBookmarksType_co] = self.bookmarks_cls(self) 

125 """This session's groups, an instance of :attr:`.bookmarks_cls`.""" 

126 

127 self.thread_creation_lock = asyncio.Lock() 

128 

129 self.__cached_presence: CachedPresence | None = None 

130 

131 self.__tasks = set[asyncio.Task[Any]]() 

132 

133 @property 

134 def user_jid(self) -> JID: 

135 return self.user.jid 

136 

137 @property 

138 def user_pk(self) -> int: 

139 return self.user.id 

140 

141 @property 

142 def http(self) -> aiohttp.ClientSession: 

143 return self.xmpp.http 

144 

145 def __remove_task(self, fut: Task[Any]) -> None: 

146 self.log.debug("Removing fut %s", fut) 

147 self.__tasks.remove(fut) 

148 

149 def create_task( 

150 self, coro: Coroutine[Any, Any, Any], name: str | None = None 

151 ) -> asyncio.Task[Any]: 

152 if name is None: 

153 warnings.warn( 

154 "Calling Session.create_task without a 'name' argument will " 

155 "be deprecated in slidge >= 0.6", 

156 DeprecationWarning, 

157 ) 

158 task = self.xmpp.loop.create_task(coro, name=name) 

159 self.__tasks.add(task) 

160 self.log.debug("Creating task %s", task) 

161 task.add_done_callback(lambda _: self.__remove_task(task)) 

162 return task 

163 

164 def cancel_all_tasks(self) -> None: 

165 for task in self.__tasks: 

166 task.cancel() 

167 

168 @abc.abstractmethod 

169 async def login(self) -> str | None: 

170 """ 

171 Logs in the gateway user to the legacy network. 

172 

173 Triggered when the gateway start and on user registration. 

174 It is recommended that this function returns once the user is logged in, 

175 so if you need to await forever (for instance to listen to incoming events), 

176 it's a good idea to wrap your listener in an asyncio.Task. 

177 

178 :return: Optionally, a text to use as the gateway status, e.g., "Connected as 'dude@legacy.network'" 

179 """ 

180 raise NotImplementedError 

181 

182 async def logout(self) -> None: 

183 """ 

184 Logs out the gateway user from the legacy network. 

185 

186 Called on gateway shutdown. 

187 """ 

188 raise NotImplementedError 

189 

190 async def on_unregister(self) -> None: 

191 """Called when the user unregisters from the gateway, after their 

192 session has been terminated but before their persistent data is 

193 deleted. 

194 

195 Optionally override this if you need to clean up additional stuff; by 

196 default it just calls :meth:`.logout`. 

197 """ 

198 with contextlib.suppress(NotImplementedError): 

199 await self.logout() 

200 

201 async def on_presence( 

202 self, 

203 resource: str, 

204 show: PseudoPresenceShow, 

205 status: str, 

206 resources: dict[str, ResourceDict], 

207 merged_resource: ResourceDict | None, 

208 ) -> None: 

209 """ 

210 Called when the gateway component receives a presence, ie, when 

211 one of the user's clients goes online of offline, or changes its 

212 status. 

213 

214 :param resource: The XMPP client identifier, arbitrary string. 

215 :param show: The presence ``<show>``, if available. If the resource is 

216 just 'available' without any ``<show>`` element, this is an empty 

217 str. 

218 :param status: A status message, like a deeply profound quote, eg, 

219 "Roses are red, violets are blue, [INSERT JOKE]". 

220 :param resources: A summary of all the resources for this user. 

221 :param merged_resource: A global presence for the user account, 

222 following rules described in :meth:`merge_resources` 

223 """ 

224 raise NotImplementedError 

225 

226 async def on_search(self, form_values: dict[str, str]) -> SearchResult | None: 

227 """ 

228 Triggered when the user uses Jabber Search (:xep:`0055`) on the component 

229 

230 Form values is a dict in which keys are defined in :attr:`.BaseGateway.SEARCH_FIELDS` 

231 

232 :param form_values: search query, defined for a specific plugin by overriding 

233 in :attr:`.BaseGateway.SEARCH_FIELDS` 

234 :return: 

235 """ 

236 raise NotImplementedError 

237 

238 async def on_avatar( 

239 self, 

240 bytes_: bytes | None, 

241 hash_: str | None, 

242 type_: str | None, 

243 width: int | None, 

244 height: int | None, 

245 ) -> None: 

246 """ 

247 Triggered when the user uses modifies their avatar via :xep:`0084`. 

248 

249 :param bytes_: The data of the avatar. According to the spec, this 

250 should always be a PNG, but some implementations do not respect 

251 that. If `None` it means the user has unpublished their avatar. 

252 :param hash_: The SHA1 hash of the avatar data. This is an identifier of 

253 the avatar. 

254 :param type_: The MIME type of the avatar. 

255 :param width: The width of the avatar image. 

256 :param height: The height of the avatar image. 

257 """ 

258 raise NotImplementedError 

259 

260 async def on_leave_space(self, space_legacy_id: str) -> None: 

261 """ 

262 Triggered when the user sends a request to leave a :xep:`0503` space. 

263 

264 :param space_legacy_id: The legacy ID of the space to leave 

265 """ 

266 raise NotImplementedError 

267 

268 async def on_preferences( 

269 self, previous: dict[str, Any], new: dict[str, Any] 

270 ) -> None: 

271 """ 

272 This is called when the user updates their preferences. 

273 

274 Override this if you need set custom preferences field and need to trigger 

275 something when a preference has changed. 

276 """ 

277 raise NotImplementedError 

278 

279 def __reset_ready(self) -> None: 

280 self.ready = self.xmpp.loop.create_future() 

281 

282 @property 

283 def logged(self) -> bool: 

284 return self._logged 

285 

286 @logged.setter 

287 def logged(self, v: bool) -> None: 

288 self.is_logging_in = False 

289 self._logged = v 

290 if self.ready.done(): 

291 if v: 

292 return 

293 self.__reset_ready() 

294 self.shutdown(logout=False) 

295 with self.xmpp.store.session() as orm: 

296 self.xmpp.store.mam.reset_source(orm) 

297 self.xmpp.store.rooms.reset_updated(orm) 

298 self.xmpp.store.contacts.reset_updated(orm) 

299 orm.commit() 

300 else: 

301 if v: 

302 self.ready.set_result(True) 

303 

304 def __repr__(self) -> str: 

305 return f"<Session of {self.user_jid}>" 

306 

307 def shutdown(self, logout: bool = True) -> asyncio.Task[None]: 

308 for m in self.bookmarks: 

309 m.shutdown() 

310 with self.xmpp.store.session() as orm: 

311 for localpart in orm.execute( 

312 sa.select(Contact.jid_localpart).filter_by( 

313 user=self.user, is_friend=True 

314 ) 

315 ).scalars(): 

316 pres = self.xmpp.make_presence( 

317 pfrom=f"{localpart}@{self.xmpp.boundjid.bare}", 

318 pto=self.user_jid, 

319 ptype="unavailable", 

320 pstatus="Gateway has shut down.", 

321 ) 

322 pres.send() 

323 if logout: 

324 return self.xmpp.loop.create_task( 

325 self.__logout(), name=f"logout of {self.user}" 

326 ) 

327 else: 

328 return self.xmpp.loop.create_task(noop_coro(), name="noop") 

329 

330 async def __logout(self) -> None: 

331 try: 

332 await self.logout() 

333 except NotImplementedError: 

334 pass 

335 except KeyboardInterrupt: 

336 pass 

337 

338 def raise_if_not_logged(self) -> None: 

339 if not self.logged: 

340 raise XMPPError( 

341 "internal-server-error", 

342 text="You are not logged to the legacy network", 

343 ) 

344 

345 @classmethod 

346 def _from_user_or_none(cls, user: GatewayUser | None) -> Self: 

347 if user is None: 

348 log.debug("user not found") 

349 raise XMPPError(text="User not found", condition="subscription-required") 

350 

351 session = _sessions.get(user.jid.bare) 

352 if session is None: 

353 _sessions[user.jid.bare] = session = cls(user) 

354 assert isinstance(session, cls) 

355 return session 

356 

357 @classmethod 

358 def from_user(cls, user: GatewayUser) -> Self: 

359 return cls._from_user_or_none(user) 

360 

361 @classmethod 

362 def from_stanza(cls, s: Message | Iq | Presence) -> Self: 

363 # """ 

364 # Get a user's :class:`.LegacySession` using the "from" field of a stanza 

365 # 

366 # Meant to be called from :class:`BaseGateway` only. 

367 # 

368 # :param s: 

369 # :return: 

370 # """ 

371 return cls.from_jid(s.get_from()) 

372 

373 @classmethod 

374 def from_jid(cls, jid: JID) -> Self: 

375 # """ 

376 # Get a user's :class:`.LegacySession` using its jid 

377 # 

378 # Meant to be called from :class:`BaseGateway` only. 

379 # 

380 # :param jid: 

381 # :return: 

382 # """ 

383 session = _sessions.get(jid.bare) 

384 if session is not None: 

385 assert isinstance(session, cls) 

386 return session 

387 with cls.xmpp.store.session() as orm: 

388 user = orm.query(GatewayUser).filter_by(jid=jid.bare).one_or_none() 

389 return cls._from_user_or_none(user) 

390 

391 @classmethod 

392 async def kill_by_jid(cls, jid: JID) -> None: 

393 # """ 

394 # Terminate a user session. 

395 # 

396 # Meant to be called from :class:`BaseGateway` only. 

397 # 

398 # :param jid: 

399 # :return: 

400 # """ 

401 log.debug("Killing session of %s", jid) 

402 for user_jid, session in _sessions.items(): 

403 if user_jid == jid.bare: 

404 break 

405 else: 

406 log.debug("Did not find a session for %s", jid) 

407 return 

408 for c in session.contacts: 

409 c.unsubscribe() 

410 for m in session.bookmarks: 

411 m.shutdown() 

412 

413 try: 

414 session = _sessions.pop(jid.bare) 

415 except KeyError: 

416 log.warning("User not found during unregistration") 

417 return 

418 

419 session.cancel_all_tasks() 

420 

421 await session.on_unregister() 

422 with cls.xmpp.store.session() as orm: 

423 orm.delete(session.user) 

424 orm.commit() 

425 

426 def __ack(self, msg: Message) -> None: 

427 if not self.xmpp.PROPER_RECEIPTS: 

428 self.xmpp.delivery_receipt.ack(msg) 

429 

430 def send_gateway_status( 

431 self, 

432 status: str | None = None, 

433 show: PresenceShows | None = None, 

434 **kwargs: Any, # noqa 

435 ) -> None: 

436 """ 

437 Send a presence from the gateway to the user. 

438 

439 Can be used to indicate the user session status, ie "SMS code required", "connected", … 

440 

441 :param status: A status message 

442 :param show: Presence stanza 'show' element. I suggest using "dnd" to show 

443 that the gateway is not fully functional 

444 """ 

445 self.__cached_presence = CachedPresence(status, show, kwargs) 

446 self.xmpp.send_presence( 

447 pto=self.user_jid.bare, pstatus=status, pshow=show, **kwargs 

448 ) 

449 

450 def send_cached_presence(self, to: JID) -> None: 

451 if not self.__cached_presence: 

452 self.xmpp.send_presence(pto=to, ptype="unavailable") 

453 return 

454 self.xmpp.send_presence( 

455 pto=to, 

456 pstatus=self.__cached_presence.status, 

457 pshow=self.__cached_presence.show, 

458 **self.__cached_presence.kwargs, 

459 ) 

460 

461 def send_gateway_message( 

462 self, 

463 text: str, 

464 **msg_kwargs: Any, # noqa 

465 ) -> None: 

466 """ 

467 Send a message from the gateway component to the user. 

468 

469 Can be used to indicate the user session status, ie "SMS code required", "connected", … 

470 

471 :param text: A text 

472 """ 

473 self.xmpp.send_text(text, mto=self.user_jid, **msg_kwargs) 

474 

475 def send_gateway_invite( 

476 self, 

477 muc: AnyMUC | JID | str, 

478 reason: str | None = None, 

479 password: str | None = None, 

480 ) -> None: 

481 """ 

482 Send an invitation to join a MUC, emanating from the gateway component. 

483 

484 The send is deferred until the session is fully initialised 

485 (``bookmarks.ready``), so that clients can actually join the room when 

486 they receive the invitation. 

487 

488 :param muc: 

489 :param reason: 

490 :param password: 

491 """ 

492 self.xmpp.invite_to( 

493 muc, 

494 reason=reason, 

495 password=password, 

496 session=self, 

497 mto=self.user_jid, 

498 ) 

499 

500 async def input(self, text: str, **msg_kwargs: Any) -> str: # noqa 

501 """ 

502 Request user input via direct messages from the gateway component. 

503 

504 Wraps call to :meth:`.BaseSession.input` 

505 

506 :param text: The prompt to send to the user 

507 :param msg_kwargs: Extra attributes 

508 :return: 

509 """ 

510 return await self.xmpp.input(self.user_jid, text, **msg_kwargs) 

511 

512 async def send_qr(self, text: str) -> None: 

513 """ 

514 Sends a QR code generated from 'text' via HTTP Upload and send the URL to 

515 ``self.user`` 

516 

517 :param text: Text to encode as a QR code 

518 """ 

519 await self.xmpp.send_qr(text, mto=self.user_jid) 

520 

521 async def get_contact_or_group_or_participant( 

522 self, jid: JID, create: bool = True 

523 ) -> "LegacyContact | AnyMUC | AnyParticipant | None": 

524 contact: LegacyContact | None = self.contacts.by_jid_only_if_exists(jid) 

525 if contact is not None: 

526 return contact 

527 if (muc := self.bookmarks.by_jid_only_if_exists(JID(jid.bare))) is not None: 

528 return await self.__get_muc_or_participant(muc, jid) 

529 else: 

530 muc = None 

531 

532 if not create: 

533 return None 

534 

535 try: 

536 contact = await self.contacts.by_jid(jid) 

537 except XMPPError: 

538 if muc is None: 

539 try: 

540 muc = await self.bookmarks.by_jid(jid) 

541 except XMPPError: 

542 return None 

543 return await self.__get_muc_or_participant(muc, jid) 

544 return contact 

545 

546 @staticmethod 

547 async def __get_muc_or_participant( 

548 muc: AnyMUC, jid: JID 

549 ) -> "AnyMUC | AnyParticipant | None": 

550 if nick := jid.resource: 

551 return await muc.get_participant(nick, create=False, fill_first=True) 

552 return muc 

553 

554 async def wait_for_ready(self, timeout: float | None = 10) -> None: 

555 # """ 

556 # Wait until session, contacts and bookmarks are ready 

557 # 

558 # (slidge internal use) 

559 # 

560 # :param timeout: 

561 # :return: 

562 # """ 

563 try: 

564 await asyncio.wait_for(asyncio.shield(self.ready), timeout) 

565 await asyncio.wait_for(asyncio.shield(self.contacts.ready), timeout) 

566 await asyncio.wait_for(asyncio.shield(self.bookmarks.ready), timeout) 

567 except TimeoutError: 

568 raise XMPPError( 

569 "recipient-unavailable", 

570 "Legacy session is not fully initialized, retry later", 

571 ) 

572 

573 def legacy_module_data_update(self, data: JSONSerializable) -> None: 

574 user = self.user 

575 user.legacy_module_data.update(data) 

576 self.xmpp.store.users.update(user) 

577 

578 def legacy_module_data_set(self, data: JSONSerializable) -> None: 

579 user = self.user 

580 user.legacy_module_data = data 

581 self.xmpp.store.users.update(user) 

582 

583 def legacy_module_data_clear(self) -> None: 

584 user = self.user 

585 user.legacy_module_data.clear() 

586 self.xmpp.store.users.update(user) 

587 

588 

589# References to `BaseSession` need for this to be defined. 

590# Does not satisfy `roster_cls = type[Roster]` for subclasses. 

591BaseSession.roster_cls = LegacyRoster # type:ignore[misc] 

592BaseSession.bookmarks_cls = LegacyBookmarks # type:ignore[misc] 

593 

594# keys = user.jid.bare 

595_sessions: dict[str, AnySession] = {} 

596log = logging.getLogger(__name__)