Coverage for slidge/db/avatar.py: 87%

230 statements  

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

1import asyncio 

2import hashlib 

3import io 

4import logging 

5from collections.abc import AsyncIterator 

6from concurrent.futures import ThreadPoolExecutor 

7from contextlib import asynccontextmanager 

8from http import HTTPStatus 

9from pathlib import Path 

10from typing import Literal 

11 

12import aiohttp 

13from multidict import CIMultiDictProxy 

14from PIL.Image import Image 

15from PIL.Image import open as open_image 

16from sqlalchemy import select 

17 

18from ..core import config 

19from ..core.attachment_upload import AttachmentUploader 

20from ..util.lock import NamedLockMixin 

21from ..util.types import AnySession, AvatarMetadata 

22from ..util.types import Avatar as AvatarType 

23from .models import Avatar 

24from .store import AvatarStore 

25 

26_AVATAR_FETCH_CONCURRENCY = 8 

27_AVATAR_DOWNLOAD_TIMEOUT = 30 

28 

29 

30class CachedAvatar: 

31 def __init__(self, stored: Avatar, root_dir: Path) -> None: 

32 self.stored = stored 

33 self._root = root_dir 

34 

35 @property 

36 def pk(self) -> int | None: 

37 return self.stored.id 

38 

39 @property 

40 def hash(self) -> str | None: 

41 """ 

42 SHA1 of avatar in PNG format, stored in self._root, meant to be 

43 served in-band. `None` if the avatar is only available through HTTP. 

44 """ 

45 return self.stored.hash 

46 

47 @property 

48 def height(self) -> int: 

49 return self.stored.height 

50 

51 @property 

52 def width(self) -> int: 

53 return self.stored.width 

54 

55 @property 

56 def etag(self) -> str | None: 

57 return self.stored.etag 

58 

59 @property 

60 def last_modified(self) -> str | None: 

61 return self.stored.last_modified 

62 

63 @property 

64 def data(self) -> bytes: 

65 return self.path.read_bytes() 

66 

67 @property 

68 def path(self) -> Path: 

69 assert self.hash 

70 return (self._root / self.hash).with_suffix(".png") 

71 

72 

73class NotModified(Exception): 

74 pass 

75 

76 

77class AvatarDownloadError(Exception): 

78 def __init__(self, url: str, status: int | None = None, message: str = "") -> None: 

79 self.url = url 

80 self.status = status 

81 self.message = message 

82 if status is None: 

83 super().__init__(f"{message} ({url})" if message else url) 

84 else: 

85 super().__init__(f"{status} {message} ({url})") 

86 

87 

88class OOBStoreError(Exception): 

89 pass 

90 

91 

92class AvatarCache(NamedLockMixin): 

93 dir: Path 

94 http: aiohttp.ClientSession 

95 store: AvatarStore 

96 

97 def __init__(self) -> None: 

98 self._thread_pool = ThreadPoolExecutor(config.AVATAR_RESAMPLING_THREADS) 

99 self._download_semaphore = asyncio.BoundedSemaphore(_AVATAR_FETCH_CONCURRENCY) 

100 super().__init__() 

101 

102 def from_stored(self, stored: Avatar) -> CachedAvatar: 

103 return CachedAvatar(stored, self.dir) 

104 

105 def set_dir(self, path: Path) -> None: 

106 self.dir = path 

107 self.dir.mkdir(exist_ok=True) 

108 for f in path.glob("*"): 

109 if f.suffix != ".png": 

110 # FIXME: remove this before 1.0.0! 

111 # slidge v0.5.0 wrote useless non-PNG files in here, this 

112 # cleans them up 

113 f.unlink() 

114 log.debug("Checking avatar files") 

115 with self.store.session(expire_on_commit=False) as orm: 

116 for stored in orm.query(Avatar).all(): 

117 avatar = CachedAvatar(stored, path) 

118 if avatar.hash is None or (avatar.hash and avatar.path.exists()): 

119 continue 

120 log.warning( 

121 "Removing avatar %s from store because %s does not exist", 

122 avatar.hash, 

123 avatar.path, 

124 ) 

125 orm.delete(stored) 

126 orm.commit() 

127 

128 def close(self) -> None: 

129 self._thread_pool.shutdown(cancel_futures=True) 

130 

131 def __get_http_headers( 

132 self, cached: CachedAvatar | Avatar | None = None 

133 ) -> dict[str, str]: 

134 headers = {} 

135 if ( 

136 cached 

137 and cached.hash 

138 and (self.dir / cached.hash).with_suffix(".png").exists() 

139 ): 

140 if last_modified := cached.last_modified: 

141 headers["If-Modified-Since"] = last_modified 

142 if etag := cached.etag: 

143 headers["If-None-Match"] = etag 

144 return headers 

145 

146 @asynccontextmanager 

147 async def _fetch( 

148 self, 

149 url: str, 

150 headers: dict[str, str], 

151 method: Literal["GET", "HEAD"], 

152 ) -> AsyncIterator[aiohttp.ClientResponse]: 

153 """Common bits of logic for HTTP requests fetching avatar (meta)data.""" 

154 async with self._download_semaphore: 

155 try: 

156 async with self.http.request( 

157 method, 

158 url, 

159 headers=headers, 

160 timeout=aiohttp.ClientTimeout(total=_AVATAR_DOWNLOAD_TIMEOUT), 

161 ) as response: 

162 yield response 

163 except aiohttp.ClientResponseError as e: 

164 raise AvatarDownloadError(url, e.status, e.message) from e 

165 except (aiohttp.ClientError, TimeoutError) as e: 

166 raise AvatarDownloadError(url, message=str(e)) from e 

167 

168 async def __download_if_modified( 

169 self, 

170 url: str, 

171 headers: dict[str, str], 

172 ) -> tuple[CIMultiDictProxy[str], bytes]: 

173 """ 

174 Download avatar only if it has been modified compared to what we have 

175 in cache. 

176 

177 :return: HTTP response headers, data 

178 :raise: NotModified if fetching was not necessary 

179 """ 

180 async with self._fetch(url, headers, "GET") as response: 

181 if response.status == HTTPStatus.NOT_MODIFIED: 

182 log.debug("Using avatar cache for %s", url) 

183 raise NotModified 

184 response.raise_for_status() 

185 data = await response.read() 

186 return response.headers, data 

187 

188 async def __is_modified(self, url: str, headers: dict[str, str]) -> bool: 

189 async with self._fetch(url, headers, "HEAD") as response: 

190 response.raise_for_status() 

191 if response.status == HTTPStatus.NOT_MODIFIED: 

192 return False 

193 cached_last_modified = headers.get("If-Modified-Since") 

194 if not cached_last_modified: 

195 return True 

196 response_last_modified = response.headers.get("last-modified") 

197 if not response_last_modified: 

198 return True 

199 return cached_last_modified != response_last_modified 

200 

201 async def url_modified(self, url: str) -> bool: 

202 with self.store.session() as orm: 

203 cached = orm.query(Avatar).filter_by(url=url).one_or_none() 

204 if cached is None: 

205 return True 

206 headers = self.__get_http_headers(cached) 

207 return await self.__is_modified(url, headers) 

208 

209 @staticmethod 

210 async def __open_image(avatar: AvatarType) -> Image: 

211 if avatar.data is not None: 

212 return open_image(io.BytesIO(avatar.data)) 

213 elif avatar.path is not None: 

214 return open_image(avatar.path) 

215 raise TypeError("Avatar must be bytes or a Path", avatar) 

216 

217 async def get( 

218 self, 

219 avatar: AvatarType, 

220 session: AnySession | None = None, 

221 convert: bool = True, 

222 ) -> CachedAvatar: 

223 if avatar.unique_id is not None: 

224 with self.store.session() as orm: 

225 stored = ( 

226 orm.query(Avatar) 

227 .filter_by(legacy_id=str(avatar.unique_id)) 

228 .one_or_none() 

229 ) 

230 if stored is not None: 

231 return self.from_stored(stored) 

232 

233 if avatar.url is not None: 

234 return await self.__fetch_url_if_not_cached( 

235 avatar, session=session, convert=convert 

236 ) 

237 

238 return await self.__process( 

239 avatar, await self.__open_image(avatar), session=session, convert=convert 

240 ) 

241 

242 async def __fetch_url_if_not_cached( 

243 self, 

244 avatar: AvatarType, 

245 session: AnySession | None = None, 

246 convert: bool = True, 

247 ) -> CachedAvatar: 

248 assert avatar.url is not None 

249 async with self.lock(avatar.unique_id or avatar.url): 

250 with self.store.session() as orm: 

251 if avatar.unique_id is None: 

252 stored = orm.query(Avatar).filter_by(url=avatar.url).one_or_none() 

253 else: 

254 stored = ( 

255 orm.query(Avatar) 

256 .filter_by(legacy_id=str(avatar.unique_id)) 

257 .one_or_none() 

258 ) 

259 if stored is not None: 

260 return self.from_stored(stored) 

261 

262 try: 

263 response_headers, data = await self.__download_if_modified( 

264 avatar.url, self.__get_http_headers(stored) 

265 ) 

266 except NotModified: 

267 assert stored is not None 

268 return self.from_stored(stored) 

269 

270 return await self.__process( 

271 avatar, 

272 open_image(io.BytesIO(data)), 

273 response_headers, 

274 session, 

275 convert=convert, 

276 img_bytes=data, 

277 ) 

278 

279 async def __process( 

280 self, 

281 avatar: AvatarType, 

282 img: Image, 

283 response_headers: CIMultiDictProxy[str] | None = None, 

284 session: AnySession | None = None, 

285 convert: bool = True, 

286 img_bytes: bytes | None = None, 

287 ) -> CachedAvatar: 

288 if convert: 

289 too_big = (size := config.AVATAR_SIZE) and any(x > size for x in img.size) 

290 if too_big: 

291 await asyncio.get_event_loop().run_in_executor( 

292 self._thread_pool, img.thumbnail, (size, size) 

293 ) 

294 img_bytes = _get_png_bytes(img) 

295 log.debug("Resampled image to %s", img.size) 

296 else: 

297 img_bytes = ( 

298 _get_any_bytes(img_bytes, avatar) 

299 if img.format == "PNG" 

300 else _get_png_bytes(img) 

301 ) 

302 else: 

303 img_bytes = _get_any_bytes(img_bytes, avatar) 

304 

305 hash_ = hashlib.sha1(img_bytes).hexdigest() 

306 http_metadata = await self.__upload(session, avatar, img, hash_, len(img_bytes)) 

307 if convert: 

308 # convert means that the avatar must be converted in order to be 

309 # served in band. 

310 # currently, it is only False for space avatars (that have no protocol 

311 # for in-band serving). 

312 file_path = (self.dir / hash_).with_suffix(".png") 

313 if file_path.exists(): 

314 log.warning("Overwriting %s", file_path) 

315 with file_path.open("wb") as file: 

316 file.write(img_bytes) 

317 with self.store.session(expire_on_commit=False) as orm: 

318 stored = orm.execute( 

319 select(Avatar).where(Avatar.hash == hash_) 

320 ).scalar() 

321 

322 if stored is not None: 

323 if ( 

324 avatar.unique_id is not None 

325 and str(avatar.unique_id) != stored.legacy_id 

326 ): 

327 log.warning( 

328 "Updating the 'unique' hash of an avatar, was '%s', is now '%s'", 

329 stored.legacy_id, 

330 avatar.unique_id, 

331 ) 

332 stored.legacy_id = str(avatar.unique_id) 

333 stored.set_http_metadata(http_metadata) 

334 orm.add(stored) 

335 orm.commit() 

336 

337 return self.from_stored(stored) 

338 elif http_metadata is None or not http_metadata.url: 

339 raise OOBStoreError() 

340 

341 stored = Avatar( 

342 hash=hash_ if convert else None, 

343 height=img.height, 

344 width=img.width, 

345 url=avatar.url, 

346 legacy_id=avatar.unique_id, 

347 ) 

348 stored.set_http_metadata(http_metadata) 

349 

350 if response_headers: 

351 stored.etag = response_headers.get("etag") 

352 stored.last_modified = response_headers.get("last-modified") 

353 

354 with self.store.session(expire_on_commit=False) as orm: 

355 if avatar.url is not None: 

356 existing = orm.execute( 

357 select(Avatar).filter_by(url=avatar.url) 

358 ).scalar_one_or_none() 

359 if existing is not None: 

360 orm.delete(existing) 

361 orm.commit() 

362 orm.add(stored) 

363 orm.commit() 

364 return self.from_stored(stored) 

365 

366 async def __upload( 

367 self, 

368 session: AnySession | None, 

369 avatar: AvatarType, 

370 img: Image, 

371 hash_: str, 

372 size: int, 

373 ) -> AvatarMetadata | None: 

374 if not session: 

375 return None 

376 if config.USE_ATTACHMENT_ORIGINAL_URLS and avatar.url is not None: 

377 assert img.format 

378 return AvatarMetadata( 

379 url=avatar.url, 

380 width=img.width, 

381 height=img.height, 

382 bytes=size, 

383 type=img.format.lower(), 

384 id=hash_, 

385 ) 

386 

387 return await AttachmentUploader(session.xmpp, session).upload_avatar( 

388 avatar, img, hash_ 

389 ) 

390 

391 

392def _get_png_bytes(img: Image) -> bytes: 

393 with io.BytesIO() as f: 

394 img.save(f, format="PNG") 

395 return f.getvalue() 

396 

397 

398def _get_any_bytes(img_bytes: bytes | None, avatar: AvatarType) -> bytes: 

399 r = img_bytes or avatar.data 

400 if r is None: 

401 if avatar.path is None: 

402 raise RuntimeError("NEVER") 

403 r = avatar.path.read_bytes() 

404 return r 

405 

406 

407avatar_cache = AvatarCache() 

408log = logging.getLogger(__name__) 

409_download_lock = asyncio.Lock() 

410 

411__all__ = ("AvatarType", "CachedAvatar", "avatar_cache")