fix(userbot): drop stale media sessions and survive download timeouts
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user