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

1"""Upload attachments to an HTTP upload service (:xep:`0363`).""" 

2 

3from __future__ import annotations 

4 

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 

23 

24import aiohttp 

25from PIL.Image import Image 

26from slixmpp import JID, Iq 

27 

28from ..db.models import Attachment 

29from ..util.types import AvatarMetadata, LegacyAttachment 

30from ..util.util import fix_namespaces, fix_suffix 

31from . import config 

32 

33if TYPE_CHECKING: 

34 from ..db.avatar import AvatarType 

35 from .gateway import BaseGateway 

36 from .session import BaseSession 

37 

38 

39class AttachmentUploader: 

40 """Turns a :class:`.LegacyAttachment` into a URL that XMPP clients can fetch.""" 

41 

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 """ 

48 

49 def __init__(self, xmpp: BaseGateway, session: BaseSession | None) -> None: 

50 self.xmpp = xmpp 

51 self.session = session 

52 

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. 

59 

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 

71 

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 ) 

99 

100 def record(self, stored: Attachment) -> None: 

101 """Remember an uploaded attachment, so that it is not uploaded twice. 

102 

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() 

111 

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 

116 

117 async def get_url(self, att: LegacyAttachment, stored: Attachment) -> str: 

118 att = _ensure_name(att) 

119 

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 ) 

126 

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 ) 

131 

132 att.legacy_file_id = stored.legacy_file_id 

133 

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 ) 

139 

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() 

152 

153 if stored.legacy_file_id: 

154 stored.url = new_url 

155 

156 return new_url 

157 

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 } 

195 

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() 

201 

202 return get 

203 

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] 

234 

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 

245 

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) 

269 

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) 

275 

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) 

311 

312 if att.size is None: 

313 att.size = destination.stat().st_size 

314 

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) 

318 

319 url = f"{file_id}/{uu}/{att.name}" 

320 return destination, url 

321 

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_) 

329 

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 ) 

356 

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 

378 

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 ) 

402 

403 

404class _AttachmentWithName(LegacyAttachment): 

405 name: str 

406 path: Path | None 

407 

408 

409class _AttachmentWithMetadata(_AttachmentWithName): 

410 size: int 

411 content_type: str 

412 

413 

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) 

423 

424 

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 

449 

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" 

454 

455 return cast(_AttachmentWithMetadata, att) 

456 

457 

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]]: 

464 

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") 

487 

488 

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) 

495 

496 

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 

504 

505 

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 

515 

516 

517def _iter_io(fp: IO[bytes], chunk_size: int = 2**18) -> Iterator[bytes]: 

518 return iter(functools.partial(fp.read, chunk_size), b"") 

519 

520 

521_TEMP_ROOT = Path(tempfile.gettempdir()) 

522_TEMP_PREFIX = "slidge-async-iterator-download" 

523 

524log = logging.getLogger(__name__)