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 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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user