fix(userbot): drop stale media sessions and survive download timeouts

This commit is contained in:
hh
2026-09-01 21:05:48 +02:00
parent cf277c56cf
commit 9c12628304
6 changed files with 83 additions and 24 deletions
@@ -1,9 +1,9 @@
from io import BytesIO
from pyrogram import Client from pyrogram import Client
from userbot.modules.avatars import repository from userbot.modules.avatars import repository
from userbot.modules.capture.context import CaptureContext from userbot.modules.capture.context import CaptureContext
from userbot.modules.download import download_bytes
from utils.logging import logger
async def capture_avatar( # noqa: PLR0913 async def capture_avatar( # noqa: PLR0913
@@ -20,9 +20,12 @@ async def capture_avatar( # noqa: PLR0913
storage_key: str | None = None storage_key: str | None = None
file_size: int | None = None file_size: int | None = None
downloaded = False downloaded = False
buffer = await client.download_media(file_id, in_memory=True) try:
if isinstance(buffer, BytesIO): data = await download_bytes(client, file_id)
data = buffer.getvalue() except TimeoutError:
logger.warning(f"[yellow]Avatar download timed out for {owner_id}.[/]")
data = None
if data is not None:
storage_key = ctx.storage.put(data) storage_key = ctx.storage.put(data)
file_size = len(data) file_size = len(data)
downloaded = True downloaded = True
+56
View File
@@ -0,0 +1,56 @@
import asyncio
import contextlib
from io import BytesIO
from pyrogram import Client
from pyrogram.types import (
Animation,
Audio,
Document,
Message,
Photo,
Sticker,
Story,
Video,
VideoNote,
Voice,
)
from utils.logging import logger
type Downloadable = (
str
| Message
| Story
| Audio
| Document
| Photo
| Sticker
| Animation
| Video
| Voice
| VideoNote
)
STOP_TIMEOUT = 5.0
_reset_lock = asyncio.Lock()
async def _reset_media_sessions(client: Client) -> None:
async with _reset_lock:
sessions = list(client.media_sessions.values())
client.media_sessions.clear()
for session in sessions:
with contextlib.suppress(Exception):
await asyncio.wait_for(session.stop(), STOP_TIMEOUT)
logger.warning(f"[yellow]Dropped {len(sessions)} stale media session(s).[/]")
async def download_bytes(client: Client, target: Downloadable) -> bytes | None:
try:
buffer = await client.download_media(target, in_memory=True)
except TimeoutError:
await _reset_media_sessions(client)
buffer = await client.download_media(target, in_memory=True)
return buffer.getvalue() if isinstance(buffer, BytesIO) else None
@@ -1,10 +1,9 @@
from io import BytesIO
from pyrogram import Client from pyrogram import Client
from pyrogram.errors import FileIdInvalid, FileReferenceExpired, FileReferenceInvalid from pyrogram.errors import FileIdInvalid, FileReferenceExpired, FileReferenceInvalid
from pyrogram.types import Photo from pyrogram.types import Photo
from userbot.modules.avatars.repository import get_avatar_file, mark_avatar_downloaded from userbot.modules.avatars.repository import get_avatar_file, mark_avatar_downloaded
from userbot.modules.download import download_bytes
from userbot.modules.jobs.context import JobContext from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register from userbot.modules.jobs.registry import register
@@ -29,13 +28,12 @@ async def _download(
client: Client, owner_id: int, unique_id: str, file_id: str client: Client, owner_id: int, unique_id: str, file_id: str
) -> bytes | None: ) -> bytes | None:
try: try:
buffer = await client.download_media(file_id, in_memory=True) return await download_bytes(client, file_id)
except STALE_FILE_ID: except STALE_FILE_ID:
fresh = await _fresh_file_id(client, owner_id, unique_id) fresh = await _fresh_file_id(client, owner_id, unique_id)
if fresh is None: if fresh is None:
return None return None
buffer = await client.download_media(fresh, in_memory=True) return await download_bytes(client, fresh)
return buffer.getvalue() if isinstance(buffer, BytesIO) else None
@register("fetch_avatar") @register("fetch_avatar")
@@ -1,8 +1,7 @@
from io import BytesIO
from pyrogram.types import Sticker from pyrogram.types import Sticker
from userbot.modules.custom_emoji.repository import is_downloaded, upsert_downloaded from userbot.modules.custom_emoji.repository import is_downloaded, upsert_downloaded
from userbot.modules.download import download_bytes
from userbot.modules.jobs.context import JobContext from userbot.modules.jobs.context import JobContext
from userbot.modules.jobs.registry import register from userbot.modules.jobs.registry import register
@@ -30,10 +29,9 @@ async def fetch_custom_emoji(ctx: JobContext) -> None:
if not stickers: if not stickers:
return return
sticker = stickers[0] sticker = stickers[0]
buffer = await client.download_media(sticker.file_id, in_memory=True) data = await download_bytes(client, sticker.file_id)
if not isinstance(buffer, BytesIO): if data is None:
return return
data = buffer.getvalue()
storage_key = capture.storage.put(data) storage_key = capture.storage.put(data)
await upsert_downloaded( await upsert_downloaded(
ctx.pool, ctx.pool,
@@ -1,4 +1,3 @@
from io import BytesIO
from typing import Any from typing import Any
from pyrogram import Client from pyrogram import Client
@@ -6,6 +5,8 @@ from pyrogram.types import Message
from userbot.modules.capture import repository from userbot.modules.capture import repository
from userbot.modules.capture.context import CaptureContext from userbot.modules.capture.context import CaptureContext
from userbot.modules.download import download_bytes
from utils.logging import logger
from utils.policy.models import CaptureToggles from utils.policy.models import CaptureToggles
_MEDIA_ATTRS = ( _MEDIA_ATTRS = (
@@ -80,9 +81,14 @@ async def capture_media( # noqa: PLR0913
downloaded = True downloaded = True
else: else:
target = message if getattr(message, kind or "", None) is obj else obj target = message if getattr(message, kind or "", None) is obj else obj
buffer = await client.download_media(target, in_memory=True) try:
if isinstance(buffer, BytesIO): data = await download_bytes(client, target)
data = buffer.getvalue() except TimeoutError:
logger.warning(
f"[yellow]Media download timed out for {chat_id}/{message_id}.[/]"
)
data = None
if data is not None:
storage_key = ctx.storage.put(data) storage_key = ctx.storage.put(data)
file_size = len(data) file_size = len(data)
downloaded = True downloaded = True
@@ -1,9 +1,8 @@
from io import BytesIO
from pyrogram import Client from pyrogram import Client
from pyrogram.types import Story from pyrogram.types import Story
from userbot.modules.capture.context import CaptureContext from userbot.modules.capture.context import CaptureContext
from userbot.modules.download import download_bytes
from userbot.modules.stories import repository from userbot.modules.stories import repository
@@ -24,9 +23,8 @@ async def save_story(client: Client, capture: CaptureContext, story: Story) -> N
capture.pool, capture.account_id, peer_id, story.id capture.pool, capture.account_id, peer_id, story.id
) )
if not (stored or story.deleted or story.media is None): if not (stored or story.deleted or story.media is None):
buffer = await client.download_media(story, in_memory=True) data = await download_bytes(client, story)
if isinstance(buffer, BytesIO): if data is not None:
data = buffer.getvalue()
storage_key = capture.storage.put(data) storage_key = capture.storage.put(data)
file_size = len(data) file_size = len(data)
downloaded = True downloaded = True