Coverage for slidge/core/dispatcher/registration.py: 58%

74 statements  

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

1from __future__ import annotations 

2 

3import logging 

4from typing import TYPE_CHECKING 

5 

6from slixmpp import JID, Iq 

7from slixmpp.exceptions import XMPPError 

8 

9from slidge.db.meta import JSONSerializable 

10 

11from ...db import GatewayUser 

12from .. import config 

13from .util import DispatcherMixin 

14 

15if TYPE_CHECKING: 

16 from slidge.util.types import AnyGateway 

17 

18 

19class RegistrationMixin(DispatcherMixin): 

20 __slots__: list[str] = [] 

21 

22 def __init__(self, xmpp: AnyGateway) -> None: 

23 super().__init__(xmpp) 

24 xmpp.plugin["xep_0077"].api.register( 

25 self.xmpp.make_registration_form, # type:ignore[arg-type] 

26 "make_registration_form", 

27 ) 

28 xmpp.plugin["xep_0077"].api.register(self._user_get, "user_get") # type:ignore[arg-type] 

29 xmpp.plugin["xep_0077"].api.register(self._user_validate, "user_validate") # type:ignore[arg-type] 

30 xmpp.plugin["xep_0077"].api.register(self._user_modify, "user_modify") # type:ignore[arg-type] 

31 # kept for slixmpp internal API compat 

32 # TODO: either fully use slixmpp internal API or rewrite registration without it at all 

33 xmpp.plugin["xep_0077"].api.register(lambda *a: None, "user_remove") 

34 

35 xmpp.add_event_handler("user_register", self._on_user_register) 

36 xmpp.add_event_handler("user_unregister", self._on_user_unregister) 

37 

38 def get_user(self, jid: JID) -> GatewayUser | None: 

39 session = self.xmpp.get_session_from_jid(jid) 

40 if session is None: 

41 return None 

42 return session.user 

43 

44 async def _user_get( 

45 self, 

46 _gateway_jid: JID, 

47 _node: str, 

48 ifrom: JID, 

49 iq: Iq, 

50 ) -> GatewayUser | None: 

51 if ifrom is None: 

52 ifrom = iq.get_from() 

53 return self.get_user(ifrom) 

54 

55 async def _user_validate( 

56 self, 

57 _gateway_jid: JID, 

58 _node: str, 

59 ifrom: JID, 

60 iq: Iq, 

61 ) -> None: 

62 xmpp = self.xmpp 

63 log.debug("User validate: %s", ifrom.bare) 

64 form_dict = {f.var: iq.get(f.var) for f in xmpp.REGISTRATION_FIELDS} 

65 xmpp.raise_if_not_allowed_jid(ifrom) 

66 try: 

67 legacy_module_data = await xmpp.user_prevalidate(ifrom, form_dict) 

68 except ValueError as e: 

69 raise XMPPError("bad-request", str(e)) 

70 except XMPPError: 

71 raise 

72 except Exception as e: # noqa: BLE001 

73 raise XMPPError("internal-server-error", str(e)) 

74 if legacy_module_data is None: 

75 legacy_module_data = form_dict 

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

77 user = GatewayUser( 

78 jid=ifrom.bare, 

79 legacy_module_data=legacy_module_data, 

80 ) 

81 orm.add(user) 

82 orm.commit() 

83 log.info("New user: %s", user) 

84 

85 async def _user_modify( 

86 self, 

87 _gateway_jid: JID, 

88 _node: str, 

89 ifrom: JID, 

90 form_dict: JSONSerializable, 

91 ) -> None: 

92 try: 

93 await self.xmpp.user_prevalidate(ifrom, form_dict) 

94 except ValueError as e: 

95 raise XMPPError("bad-request", str(e)) 

96 except XMPPError: 

97 raise 

98 except Exception as e: # noqa: BLE001 

99 raise XMPPError("internal-server-error", str(e)) 

100 log.debug("Modify user: %s", ifrom) 

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

102 user = orm.query(GatewayUser).one_or_none() 

103 if user is None: 

104 raise XMPPError("internal-server-error", "User not found") 

105 user.legacy_module_data.update(form_dict) 

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

107 

108 async def _on_user_register(self, iq: Iq) -> None: 

109 session = await self._get_session(iq, wait_for_ready=False) 

110 for jid in config.ADMINS: 

111 self.xmpp.send_message( 

112 mto=jid, 

113 mbody=f"{iq.get_from()} has registered", 

114 mtype="chat", 

115 mfrom=self.xmpp.boundjid.bare, 

116 ) 

117 session.send_gateway_message(self.xmpp.WELCOME_MESSAGE) 

118 await self.xmpp.login_wrap(session) 

119 

120 async def _on_user_unregister(self, iq: Iq) -> None: 

121 await self.xmpp.kill_session(iq.get_from()) 

122 

123 

124log = logging.getLogger(__name__)