Coverage for slidge/core/dispatcher/util.py: 94%
122 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
1import logging
2from collections.abc import Awaitable, Callable
3from functools import wraps
4from typing import TYPE_CHECKING, TypeVar
6from slixmpp import JID, Iq, Message, Presence
7from slixmpp.exceptions import XMPPError
8from slixmpp.xmlstream import StanzaBase
10from ...contact.roster import ContactIsUser
11from ...util.types import (
12 AnyMUC,
13 AnyRecipient,
14 AnySession,
15 RecipientType,
16)
18if TYPE_CHECKING:
19 from slidge.util.types import AnyGateway
22class Ignore(BaseException):
23 pass
26class DispatcherMixin:
27 __slots__: list[str] = []
28 xmpp: "AnyGateway"
30 def __init__(self, xmpp: "AnyGateway") -> None:
31 self.xmpp = xmpp # type:ignore[misc]
33 async def _get_session(
34 self,
35 stanza: Message | Presence | Iq,
36 timeout: int | None = 10,
37 wait_for_ready: bool = True,
38 logged: bool = False,
39 ) -> AnySession:
40 xmpp = self.xmpp
41 if stanza.get_from().server == xmpp.boundjid.bare:
42 log.debug("Ignoring echo")
43 raise Ignore
44 if (
45 isinstance(stanza, Message)
46 and stanza.get_type() == "chat"
47 and stanza.get_to() == xmpp.boundjid.bare
48 ):
49 log.debug("Ignoring message to component")
50 raise Ignore
51 session = await self._get_session_from_jid(
52 stanza.get_from(), timeout, wait_for_ready, logged
53 )
54 if isinstance(stanza, Message) and _ignore(session, stanza):
55 raise Ignore
56 return session
58 async def _get_session_from_jid(
59 self,
60 jid: JID,
61 timeout: int | None = 10,
62 wait_for_ready: bool = True,
63 logged: bool = False,
64 ) -> AnySession:
65 session = self.xmpp.get_session_from_jid(jid)
66 if session is None:
67 raise XMPPError("registration-required")
68 if logged:
69 session.raise_if_not_logged()
70 if wait_for_ready:
71 await session.wait_for_ready(timeout)
72 return session
74 async def get_muc_from_stanza(self, iq: Iq | Message | Presence) -> AnyMUC:
75 ito = iq.get_to()
76 if ito == self.xmpp.boundjid.bare:
77 raise XMPPError("bad-request", text="This is only handled for MUCs")
79 session = await self._get_session(iq, logged=True)
80 muc = await session.bookmarks.by_jid(ito)
81 return muc # type:ignore[no-any-return]
83 def _xmpp_msg_id_to_legacy(
84 self,
85 xmpp_id: str,
86 recipient: AnyRecipient,
87 origin: bool = False,
88 ) -> str:
89 with self.xmpp.store.session() as orm:
90 sent = self.xmpp.store.id_map.get_legacy(
91 orm, recipient.stored.id, xmpp_id, recipient.is_group, origin
92 )
93 if sent is not None:
94 return sent
96 return xmpp_id
98 async def _get_recipient_and_thread(
99 self, msg: Message
100 ) -> tuple[AnyRecipient, str | None]:
101 session = await self._get_session(msg)
102 e: AnyRecipient = await get_recipient(session, msg)
103 legacy_thread = await self._xmpp_to_legacy_thread(session, msg, e)
104 return e, legacy_thread
106 async def _xmpp_to_legacy_thread(
107 self, session: AnySession, msg: Message, recipient: RecipientType
108 ) -> str | None:
109 if not self.xmpp.THREADS:
110 return None
112 xmpp_thread = msg["thread"]
113 if not xmpp_thread:
114 return None
116 if session.MESSAGE_IDS_ARE_THREAD_IDS:
117 return self._xmpp_msg_id_to_legacy(xmpp_thread, recipient)
119 with session.xmpp.store.session() as orm:
120 legacy_thread_str = session.xmpp.store.id_map.get_thread(
121 orm, recipient.stored.id, xmpp_thread, recipient.is_group
122 )
123 if legacy_thread_str is not None:
124 return legacy_thread_str
125 async with session.thread_creation_lock:
126 legacy_thread = await recipient.create_thread(xmpp_thread)
127 with session.xmpp.store.session() as orm:
128 session.xmpp.store.id_map.set_thread(
129 orm,
130 recipient.stored.id,
131 str(legacy_thread),
132 xmpp_thread,
133 recipient.is_group,
134 )
135 orm.commit()
136 return legacy_thread
139def _ignore(session: AnySession, msg: Message) -> bool:
140 i = msg.get_id()
141 if i.startswith("slidge-carbon-"):
142 return True
143 if i not in session.ignore_messages:
144 return False
145 session.log.debug("Ignored sent carbon: %s", i)
146 session.ignore_messages.remove(i)
147 return True
150async def get_recipient(session: AnySession, m: Message) -> AnyRecipient:
151 session.raise_if_not_logged()
152 if m.get_type() == "groupchat":
153 muc = await session.bookmarks.by_jid(m.get_to())
154 r = m.get_from().resource
155 if r not in muc.get_user_resources():
156 session.create_task(muc.kick_resource(r), name=f"kick of {r} from {muc}")
157 raise XMPPError("not-acceptable", "You are not connected to this chat")
158 return muc # type:ignore[no-any-return]
159 else:
160 return await session.contacts.by_jid(m.get_to()) # type:ignore[no-any-return]
163SelfType = TypeVar("SelfType")
164StanzaType = TypeVar("StanzaType", bound=StanzaBase)
165HandlerType = Callable[[SelfType, StanzaType], Awaitable[None]]
168def exceptions_to_xmpp_errors[SelfType, StanzaType: StanzaBase](
169 cb: HandlerType[SelfType, StanzaType],
170) -> HandlerType[SelfType, StanzaType]:
171 @wraps(cb)
172 async def wrapped(self: SelfType, stanza: StanzaType) -> None:
173 try:
174 await cb(self, stanza)
175 except Ignore:
176 pass
177 except XMPPError:
178 raise
179 except NotImplementedError as e:
180 log.debug("NotImplementedError raised in %s", cb)
181 tb = e.__traceback__
182 assert tb is not None
183 while tb.tb_next is not None:
184 tb = tb.tb_next
185 frame = tb.tb_frame
186 method_name = frame.f_code.co_name
187 self_obj = frame.f_locals.get("self")
188 if self_obj is not None:
189 method_name = f"{self_obj.__class__.__name__}.{method_name}"
190 raise XMPPError(
191 "feature-not-implemented",
192 f"{method_name} is not implemented by the legacy module",
193 clear=False,
194 )
195 except ContactIsUser:
196 raise XMPPError(
197 "bad-request", "Actions with your bridged self are not allowed."
198 )
199 except Exception as e:
200 log.error(
201 "Failed to handle incoming stanza: %s - %s", self, stanza, exc_info=e
202 )
203 raise XMPPError("internal-server-error", str(e))
205 return wrapped
208log = logging.getLogger(__name__)