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 userbot.modules.avatars import repository
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
@@ -20,9 +20,12 @@ async def capture_avatar( # noqa: PLR0913
storage_key: str | None = None
file_size: int | None = None
downloaded = False
buffer = await client.download_media(file_id, in_memory=True)
if isinstance(buffer, BytesIO):
data = buffer.getvalue()
try:
data = await download_bytes(client, file_id)
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)
file_size = len(data)
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.errors import FileIdInvalid, FileReferenceExpired, FileReferenceInvalid
from pyrogram.types import Photo
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.registry import register
@@ -29,13 +28,12 @@ async def _download(
client: Client, owner_id: int, unique_id: str, file_id: str
) -> bytes | None:
try:
buffer = await client.download_media(file_id, in_memory=True)
return await download_bytes(client, file_id)
except STALE_FILE_ID:
fresh = await _fresh_file_id(client, owner_id, unique_id)
if fresh is None:
return None
buffer = await client.download_media(fresh, in_memory=True)
return buffer.getvalue() if isinstance(buffer, BytesIO) else None
return await download_bytes(client, fresh)
@register("fetch_avatar")
@@ -1,8 +1,7 @@
from io import BytesIO
from pyrogram.types import Sticker
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.registry import register
@@ -30,10 +29,9 @@ async def fetch_custom_emoji(ctx: JobContext) -> None:
if not stickers:
return
sticker = stickers[0]
buffer = await client.download_media(sticker.file_id, in_memory=True)
if not isinstance(buffer, BytesIO):
data = await download_bytes(client, sticker.file_id)
if data is None:
return
data = buffer.getvalue()
storage_key = capture.storage.put(data)
await upsert_downloaded(
ctx.pool,
@@ -1,4 +1,3 @@
from io import BytesIO
from typing import Any
from pyrogram import Client
@@ -6,6 +5,8 @@ from pyrogram.types import Message
from userbot.modules.capture import repository
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
_MEDIA_ATTRS = (
@@ -80,9 +81,14 @@ async def capture_media( # noqa: PLR0913
downloaded = True
else:
target = message if getattr(message, kind or "", None) is obj else obj
buffer = await client.download_media(target, in_memory=True)
if isinstance(buffer, BytesIO):
data = buffer.getvalue()
try:
data = await download_bytes(client, target)
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)
file_size = len(data)
downloaded = True
@@ -1,9 +1,8 @@
from io import BytesIO
from pyrogram import Client
from pyrogram.types import Story
from userbot.modules.capture.context import CaptureContext
from userbot.modules.download import download_bytes
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
)
if not (stored or story.deleted or story.media is None):
buffer = await client.download_media(story, in_memory=True)
if isinstance(buffer, BytesIO):
data = buffer.getvalue()
data = await download_bytes(client, story)
if data is not None:
storage_key = capture.storage.put(data)
file_size = len(data)
downloaded = True