Coverage for slidge/core/attachment_upload.py: 86%
300 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
1"""Upload attachments to an HTTP upload service (:xep:`0363`)."""
3from __future__ import annotations
5import contextlib
6import datetime
7import functools
8import hashlib
9import io
10import logging
11import os
12import shutil
13import stat
14import tempfile
15import warnings
16from base64 import b64encode
17from collections.abc import AsyncIterator, Iterator
18from mimetypes import guess_type
19from pathlib import Path
20from typing import IO, TYPE_CHECKING, Literal, cast
21from urllib.parse import quote as urlquote
22from uuid import uuid4
24import aiohttp
25from PIL.Image import Image
26from slixmpp import JID, Iq
28from ..db.models import Attachment
29from ..util.types import AvatarMetadata, LegacyAttachment
30from ..util.util import fix_namespaces, fix_suffix
31from . import config
33if TYPE_CHECKING:
34 from ..db.avatar import AvatarType
35 from .gateway import BaseGateway
36 from .session import BaseSession
39class AttachmentUploader:
40 """Turns a :class:`.LegacyAttachment` into a URL that XMPP clients can fetch."""
42 xmpp: BaseGateway
43 session: BaseSession | None
44 """
45 The session uplaoding attachments, or :const:`None` for the gateway component,
46 which is not bound to a :term:`User`.
47 """
49 def __init__(self, xmpp: BaseGateway, session: BaseSession | None) -> None:
50 self.xmpp = xmpp
51 self.session = session
53 @contextlib.asynccontextmanager
54 async def dedup_lock(
55 self,
56 attachment: LegacyAttachment | Path | str,
57 ) -> AsyncIterator[None]:
58 """Take a lock on an attachment with a given name.
60 Prevents races which download the same attachment several times."""
61 session = self.session
62 if (
63 session is None
64 or not isinstance(attachment, LegacyAttachment)
65 or attachment.legacy_file_id is None
66 ):
67 yield
68 else:
69 async with session.lock(("attachment", attachment.legacy_file_id)):
70 yield
72 async def get_stored(self, attachment: LegacyAttachment) -> Attachment:
73 """
74 Fetch the :class:`.Attachment` already uploaded for this user, if any,
75 or a new (transient) one.
76 """
77 session = self.session
78 if attachment.legacy_file_id is not None and session is not None:
79 with self.xmpp.store.session() as orm:
80 stored = (
81 orm.query(Attachment)
82 .filter_by(
83 legacy_file_id=str(attachment.legacy_file_id),
84 user_account_id=session.user_pk,
85 )
86 .one_or_none()
87 )
88 if stored is not None:
89 if not await self.__valid_url(session, stored.url):
90 stored.url = None # type:ignore
91 return stored
92 return Attachment(
93 user_account_id=None if session is None else session.user_pk,
94 legacy_file_id=None
95 if attachment.legacy_file_id is None
96 else str(attachment.legacy_file_id),
97 url=attachment.url if config.USE_ATTACHMENT_ORIGINAL_URLS else None,
98 )
100 def record(self, stored: Attachment) -> None:
101 """Remember an uploaded attachment, so that it is not uploaded twice.
103 No-op without a session, since attachments are stored per user.
104 """
105 # TODO: we need a separate mechanism to record gateway component attachments.
106 if self.session is None:
107 return
108 with self.xmpp.store.session(expire_on_commit=False) as orm:
109 orm.add(stored)
110 orm.commit()
112 @staticmethod
113 async def __valid_url(session: BaseSession, url: str) -> bool:
114 async with session.http.head(url) as r:
115 return r.status < 400
117 async def get_url(self, att: LegacyAttachment, stored: Attachment) -> str:
118 att = _ensure_name(att)
120 if len(att.name) > config.ATTACHMENT_MAXIMUM_FILE_NAME_LENGTH:
121 log.debug("Trimming long filename: %s", att.name)
122 base, ext = os.path.splitext(att.name)
123 att.name = (
124 base[: config.ATTACHMENT_MAXIMUM_FILE_NAME_LENGTH - len(ext)] + ext
125 )
127 if config.FIX_FILENAME_SUFFIX_MIME_TYPE and isinstance(att.path, Path):
128 att.name, att.content_type = fix_suffix(
129 att.path, att.content_type, att.name
130 )
132 att.legacy_file_id = stored.legacy_file_id
134 if att.date is None and att.path is not None:
135 assert isinstance(att.path, Path)
136 att.date = datetime.datetime.fromtimestamp(
137 att.path.stat().st_mtime, tz=datetime.UTC
138 )
140 hasher = hashlib.sha256() if att.sha256 is None else None
141 if config.NO_UPLOAD_PATH:
142 att.path, new_url = await self.__no_upload(
143 att, stored.legacy_file_id, hasher
144 )
145 new_url = (
146 (config.NO_UPLOAD_URL_PREFIX or "") + "/message/" + urlquote(new_url)
147 )
148 else:
149 new_url = await self.__upload(att, hasher=hasher)
150 if hasher is not None:
151 att.sha256 = b64encode(hasher.digest()).decode()
153 if stored.legacy_file_id:
154 stored.url = new_url
156 return new_url
158 async def __upload(
159 self,
160 att: _AttachmentWithName,
161 purpose: Literal["message", "profile"] = "message",
162 hasher: hashlib._Hash | None = None,
163 ) -> str:
164 assert config.UPLOAD_SERVICE
165 att = await _ensure_metadata(att, self.xmpp.http)
166 iq_slot = await self.__request_upload_slot(
167 config.UPLOAD_SERVICE,
168 att.name,
169 att.size,
170 att.content_type,
171 purpose=purpose,
172 )
173 if iq_slot["type"] == "error":
174 # COMPAT: In theory, __request_upload_slot() raises IqError, but
175 # there is a bug in prosody's mod_privilege where the outer
176 # IQ type is (illegally) not set to error when the inner IQ
177 # has type='error'.
178 raise RuntimeError(f"Error while requesting upload slot: {iq_slot}")
179 slot = iq_slot.get_plugin("http_upload_slot", check=True)
180 if slot is None:
181 raise RuntimeError(f"No upload slot in this IQ: {iq_slot}")
182 put = slot["put"]["url"]
183 assert isinstance(put, str)
184 if not put:
185 raise RuntimeError(f"Cannot find a PUT URL in: {slot}")
186 get = slot["get"]["url"]
187 assert isinstance(get, str)
188 if not get:
189 raise RuntimeError(f"Cannot find a GET URL in: {slot}")
190 headers = {
191 "Content-Length": str(att.size),
192 "Content-Type": att.content_type,
193 **{header["name"]: header["value"] for header in slot["put"]["headers"]},
194 }
196 async with (
197 _get_data(att, self.xmpp.http, hasher) as data,
198 self.xmpp.http.put(slot["put"]["url"], data=data, headers=headers) as resp,
199 ):
200 resp.raise_for_status()
202 return get
204 async def __request_upload_slot(
205 self,
206 upload_service: JID | str,
207 filename: str,
208 size: int,
209 content_type: str,
210 *,
211 purpose: Literal["message", "profile"] = "message",
212 ) -> Iq:
213 iq_request = self.xmpp.make_iq_get(ito=upload_service)
214 request = iq_request["http_upload_request"]
215 request["filename"] = filename
216 request["size"] = str(size)
217 request["content-type"] = content_type
218 if purpose != "message":
219 request.enable(purpose)
220 session = self.session
221 if session is not None:
222 iq_request.set_from(session.user_jid)
223 try:
224 return await self.xmpp.plugin["xep_0356"].send_privileged_iq(iq_request)
225 except Exception as e: # noqa: BLE001
226 warnings.warn(
227 "Could not request upload slot on behalf of "
228 f"{session.user_jid}: {e}."
229 "Falling back to not using privileges."
230 )
231 fix_namespaces(iq_request.xml, "jabber:client", "jabber:component:accept")
232 iq_request.set_from(config.UPLOAD_REQUESTER or self.xmpp.boundjid)
233 return await iq_request.send() # type:ignore[no-any-return]
235 async def __no_upload(
236 self,
237 att: _AttachmentWithName,
238 legacy_file_id: str | None,
239 hasher: hashlib._Hash | None = None,
240 ) -> tuple[Path, str]:
241 file_id = uuid4().hex if legacy_file_id is None else legacy_file_id
242 assert config.NO_UPLOAD_PATH is not None
243 assert config.NO_UPLOAD_URL_PREFIX is not None
244 destination_dir = Path(config.NO_UPLOAD_PATH) / "message" / file_id
246 if destination_dir.exists():
247 log.debug("Dest dir exists: %s", destination_dir)
248 files = [f for f in destination_dir.glob("**/*") if f.is_file()]
249 if len(files) == 1:
250 log.debug(
251 "Found the legacy attachment '%s' at '%s'",
252 legacy_file_id,
253 files[0],
254 )
255 name = files[0].name
256 uu = files[0].parent.name # anti-obvious url trick, see below
257 return files[0], f"{file_id}/{uu}/{name}"
258 else:
259 log.warning(
260 (
261 "There are several or zero files in %s, "
262 "slidge doesn't know which one to pick among %s. "
263 "Removing the dir."
264 ),
265 destination_dir,
266 files,
267 )
268 shutil.rmtree(destination_dir)
270 log.debug("Did not find a file in: %s", destination_dir)
271 # let's use a UUID to avoid URLs being too obvious
272 uu = str(uuid4())
273 destination_dir = destination_dir / uu
274 destination_dir.mkdir(parents=True)
276 assert att.name
277 destination = destination_dir / att.name
278 if att.path:
279 assert isinstance(att.path, Path)
280 if hasher is not None:
281 with att.path.open("rb") as fp:
282 for chunk in _iter_io(fp):
283 hasher.update(chunk)
284 try:
285 destination.hardlink_to(att.path)
286 except OSError as e:
287 if is_temp_path(att.path):
288 shutil.copy2(att.path, destination)
289 else:
290 log.debug("Could not hardlink: %s, attempting symlink", e)
291 try:
292 destination.symlink_to(att.path)
293 except OSError as e:
294 log.debug("Could not symlink: %s, copying data", e)
295 shutil.copy2(att.path, destination)
296 elif att.data:
297 destination.write_bytes(att.data)
298 if hasher is not None:
299 hasher.update(att.data)
300 else:
301 with destination.open("wb") as f:
302 async with _get_data(att, self.xmpp.http, hasher) as data:
303 if isinstance(data, AsyncIterator):
304 async for chunk in data:
305 f.write(chunk)
306 elif isinstance(data, bytes):
307 f.write(data)
308 else:
309 for chunk in _iter_io(data):
310 f.write(chunk)
312 if att.size is None:
313 att.size = destination.stat().st_size
315 if config.NO_UPLOAD_FILE_READ_OTHERS:
316 log.debug("Changing perms of %s", destination)
317 destination.chmod(destination.stat().st_mode | stat.S_IROTH)
319 url = f"{file_id}/{uu}/{att.name}"
320 return destination, url
322 async def upload_avatar(
323 self, avatar: AvatarType, img: Image, hash_: str
324 ) -> AvatarMetadata | None:
325 if config.NO_UPLOAD_PATH:
326 return await self.__no_upload_avatar(avatar, img, hash_)
327 else:
328 return await self.__upload_avatar(avatar, img, hash_)
330 async def __no_upload_avatar(
331 self, avatar: AvatarType, img: Image, hash_: str
332 ) -> AvatarMetadata | None:
333 assert config.NO_UPLOAD_PATH
334 assert img.format
335 format = img.format.lower()
336 dest = (Path(config.NO_UPLOAD_PATH) / "profile" / hash_ / "avatar").with_suffix(
337 "." + format
338 )
339 url = f"{config.NO_UPLOAD_URL_PREFIX}/profile/{hash_}/{dest.name}"
340 if not dest.exists():
341 dest.parent.mkdir(exist_ok=True, parents=True)
342 if avatar.path:
343 shutil.copy2(avatar.path, dest)
344 elif avatar.data:
345 dest.write_bytes(avatar.data)
346 else:
347 img.save(dest)
348 return AvatarMetadata(
349 url=url,
350 width=img.width,
351 height=img.height,
352 bytes=dest.stat().st_size,
353 id=hashlib.sha1(dest.read_bytes()).hexdigest(),
354 type=format,
355 )
357 async def __upload_avatar(
358 self, avatar: AvatarType, img: Image, hash_: str
359 ) -> AvatarMetadata | None:
360 assert config.UPLOAD_SERVICE is not None
361 iq_or_info = await self.xmpp.plugin["xep_0030"].get_info(
362 JID(config.UPLOAD_SERVICE), cached=True
363 )
364 if isinstance(iq_or_info, Iq):
365 features = iq_or_info["disco_info"].get_features()
366 else:
367 features = iq_or_info.get_features()
368 if (
369 self.xmpp.plugin["xep_0363"].stanza.ProfilePurpose.namespace + "#profile"
370 not in features
371 ):
372 warnings.warn(
373 f"The upload service {config.UPLOAD_SERVICE} does not support "
374 "the 'profile' purpose, avatar data can only be served in-band.",
375 UserWarning,
376 )
377 return None
379 att = await _ensure_metadata(
380 _AttachmentWithName(
381 data=avatar.data,
382 url=avatar.url,
383 path=avatar.path,
384 name=avatar.path.name if avatar.path else "avatar",
385 ),
386 self.xmpp.http,
387 )
388 hasher = hashlib.sha1()
389 try:
390 url = await self.__upload(att, purpose="profile", hasher=hasher)
391 except Exception:
392 log.exception("Could not upload avatar")
393 return None
394 return AvatarMetadata(
395 url=url,
396 width=img.width,
397 height=img.height,
398 bytes=att.size,
399 id=hasher.hexdigest(),
400 type=att.content_type.removeprefix("image/"),
401 )
404class _AttachmentWithName(LegacyAttachment):
405 name: str
406 path: Path | None
409class _AttachmentWithMetadata(_AttachmentWithName):
410 size: int
411 content_type: str
414def _ensure_name(att: LegacyAttachment) -> _AttachmentWithName:
415 if not att.name:
416 if att.path:
417 att.name = Path(att.path).name
418 elif att.url:
419 att.name = att.url.split("/")[-1]
420 else:
421 att.name = "unnamed-file"
422 return cast(_AttachmentWithName, att)
425async def _ensure_metadata(
426 att: _AttachmentWithName, http: aiohttp.ClientSession
427) -> _AttachmentWithMetadata:
428 if att.size is None:
429 if att.data:
430 att.size = len(att.data)
431 elif att.stream:
432 att.stream.seek(0, io.SEEK_END)
433 att.size = att.stream.tell()
434 att.stream.seek(0)
435 elif att.path:
436 assert isinstance(att.path, Path)
437 att.size = att.path.stat().st_size
438 elif att.url:
439 async with http.head(att.url) as resp:
440 att.size = resp.content_length
441 elif att.aio_stream:
442 warnings.warn("A size should be passed with async iterators")
443 tmp_dir = Path(tempfile.mkdtemp(prefix=_TEMP_PREFIX))
444 with (tmp_dir / att.name).open("wb") as fp:
445 async for chunk in att.aio_stream:
446 fp.write(chunk)
447 att.path = Path(fp.name)
448 att.size = att.path.stat().st_size
450 if not att.content_type:
451 att.content_type, _encoding = guess_type(att.name)
452 if not att.content_type:
453 att.content_type = "application/octet-stream"
455 return cast(_AttachmentWithMetadata, att)
458@contextlib.asynccontextmanager
459async def _get_data(
460 att: LegacyAttachment,
461 http: aiohttp.ClientSession,
462 hasher: hashlib._Hash | None = None,
463) -> AsyncIterator[bytes | IO[bytes] | AsyncIterator[bytes]]:
465 if att.data is not None:
466 if hasher is not None:
467 hasher.update(att.data)
468 yield att.data
469 elif att.stream is not None:
470 _hash(att.stream, hasher)
471 yield att.stream
472 elif att.path is not None:
473 assert isinstance(att.path, Path)
474 with att.path.open("rb") as fp:
475 _hash(fp, hasher)
476 yield fp
477 elif att.aio_stream is not None:
478 # The aiostream may already have been consumed if a size wasn't passed.
479 # But in this case the `path` attribute is not None, cf above.
480 yield _hash_and_yield_async(att.aio_stream, hasher)
481 elif att.url is not None:
482 async with http.get(att.url) as resp_get:
483 resp_get.raise_for_status()
484 yield _hash_and_yield_async(resp_get.content.iter_any(), hasher)
485 else:
486 raise RuntimeError("NEVER")
489def _hash(fp: IO[bytes], hasher: hashlib._Hash | None = None) -> None:
490 if hasher is not None:
491 fp.seek(0)
492 for chunk in _iter_io(fp):
493 hasher.update(chunk)
494 fp.seek(0)
497async def _hash_and_yield_async(
498 source: AsyncIterator[bytes], hasher: hashlib._Hash | None = None
499) -> AsyncIterator[bytes]:
500 async for chunk in source:
501 if hasher is not None:
502 hasher.update(chunk)
503 yield chunk
506def is_temp_path(path: Path, async_iterator_download_only: bool = False) -> bool:
507 try:
508 rel = path.relative_to(_TEMP_ROOT)
509 except ValueError:
510 return False
511 if async_iterator_download_only:
512 return rel.parts[0].startswith(_TEMP_PREFIX)
513 else:
514 return True
517def _iter_io(fp: IO[bytes], chunk_size: int = 2**18) -> Iterator[bytes]:
518 return iter(functools.partial(fp.read, chunk_size), b"")
521_TEMP_ROOT = Path(tempfile.gettempdir())
522_TEMP_PREFIX = "slidge-async-iterator-download"
524log = logging.getLogger(__name__)