Ver código fonte

.upload and .download: [BETA]

TheHamkerCat 5 anos atrás
pai
commit
1ca8addfc4

+ 0 - 1
requirements.txt

@@ -19,7 +19,6 @@ TgCrypto
 uvloop
 youtube_dl
 bs4
-wget
 python-dotenv
 feedparser
 pyromod

+ 2 - 2
wbb/__main__.py

@@ -182,7 +182,7 @@ async def help_parser(name, keyboard=None):
             paginate_modules(0, HELPABLE, "help")
         )
     return (
-        """Hello {first_name}! My name is {bot_name}!
+        """Hello {first_name}, My name is {bot_name}.
 I'm a group management bot with some useful features.
 You can choose an option below, by clicking a button.
 Also you can ask anything in Support Group.
@@ -225,7 +225,7 @@ async def help_button(client, query):
     back_match = re.match(r"help_back", query.data)
     create_match = re.match(r"help_create", query.data)
     top_text = f"""
-Hello {query.from_user.first_name}! My name is {BOT_NAME}!
+Hello {query.from_user.first_name}, My name is {BOT_NAME}.
 I'm a group management bot with some usefule features.
 You can choose an option below, by clicking a button.
 Also you can ask anything in Support Group.

+ 114 - 0
wbb/modules/download_upload.py

@@ -0,0 +1,114 @@
+from os import remove
+from os.path import isfile
+from time import ctime, time
+
+from pyrogram import filters
+from pyrogram.types import Message
+
+from wbb import SUDOERS, USERBOT_PREFIX, app2
+from wbb.core.sections import section
+from wbb.modules.userbot import add_task, eor, rm_task, tasks
+from wbb.utils.downloader import download
+from wbb.utils.functions import progress
+
+
+@app2.on_message(
+    filters.user(SUDOERS)
+    & filters.command("download", prefixes=USERBOT_PREFIX)
+)
+async def download_func(_, message: Message):
+    reply = message.reply_to_message
+    start = task_id = int(time())
+
+    if reply:
+        m = await eor(message, text="Downloading...")
+
+        task = await add_task(
+            reply.download,
+            task_id=task_id,
+            progress=progress,
+            progress_args=(start, task_id, m),
+        )
+        await task
+        await rm_task(task_id)
+
+        elapsed = int(time() - start)
+        body = {
+            "Started": ctime(start),
+            "Time": f"{elapsed}s",
+        }
+        return await eor(m, text=section("Downloaded", body))
+
+    text = message.text
+    if len(text.split()) < 2:
+        return await eor(message, text="Invalid Arguments")
+
+    url = text.split(None, 1)[1]
+    task_id = int(time())
+
+    body = {
+        "Started": ctime(start),
+        "Task ID": task_id,
+        "URL": url,
+    }
+    m = await eor(
+        message,
+        text=section("Downloading", body, underline=False),
+        disable_web_page_preview=True,
+    )
+
+    try:
+        await download(
+            url,
+            progress_func=(progress, [start, task_id, m]),
+            task_id=task_id,
+        )
+    except Exception as e:
+        return await eor(m, text=f"**Error:** `{str(e)}`")
+
+    elapsed = int(time() - start)
+    body = {
+        "Started": ctime(start),
+        "Took": f"{elapsed}s",
+        "Task ID": task_id,
+    }
+    await eor(m, text=section("Downloaded", body, underline=False))
+
+
+@app2.on_message(
+    filters.user(SUDOERS)
+    & filters.command("upload", prefixes=USERBOT_PREFIX)
+)
+async def upload_func(_, message: Message):
+    if len(message.text.split()) != 2:
+        return await eor(message, text="Invalid Arguments")
+
+    url_or_path = message.text.split(None, 1)[1]
+
+    m = await eor(message, text="Uploading..")
+    start = task_id = int(time())
+
+    async def upload_file(path):
+        task = await add_task(
+            message.reply_document,
+            task_id,
+            path,
+            progress=progress,
+            progress_args=(start, task_id, m),
+        )
+        await task
+        await rm_task(task_id)
+        elapsed = int(time() - start)
+        return await eor(m, text=f"Uploaded in {elapsed}s")
+
+    try:
+        if isfile(url_or_path):
+            return await upload_file(url_or_path)
+        path = await download(
+            url_or_path,
+            task_id=task_id,
+            progress_func=(progress, [start, task_id, m]),
+        )
+        return await upload_file(path)
+    except Exception as e:
+        return await eor(m, text=f"**Error:** `{str(e)}`")

+ 5 - 0
wbb/modules/purge_me.py

@@ -77,6 +77,11 @@ __HELP__ = """
         To remove a user from sudoers.
 `.sudoers`
         To list sudo users.
+
+`.download [URL or reply to a file]`
+        Download a file from TG or URL
+`.upload [URL or File Path]`
+        Upload a file from local or URL
 """
 
 

+ 34 - 20
wbb/modules/userbot.py

@@ -10,7 +10,7 @@ import re
 import subprocess
 import sys
 import traceback
-from asyncio import get_event_loop
+from asyncio import Lock, create_task
 from html import escape
 from inspect import getfullargspec
 from io import StringIO
@@ -29,6 +29,7 @@ p = print
 r = None
 arq = arq
 arrow = lambda x: x.text + "\n`→`"
+TASKS_LOCK = Lock()
 tasks = {}
 
 
@@ -48,6 +49,27 @@ async def eor(msg: Message, **kwargs):
     )
 
 
+async def add_task(taskFunc, task_id, *args, **kwargs):
+    global tasks
+    task = create_task(taskFunc(*args, **kwargs))
+    tasks[task_id] = task
+    return task
+
+
+async def rm_task(task_id=None):
+    global tasks
+    async with TASKS_LOCK:
+        for key, value in list(tasks.items()):
+            if value.done() or value.cancelled():
+                del tasks[key]
+
+        if task_id:
+            if task_id in tasks:
+                if not tasks[task_id].done():
+                    tasks[task_id].cancel()
+                del tasks[task_id]
+
+
 @app2.on_message(
     filters.user(SUDOERS)
     & ~filters.forwarded
@@ -70,8 +92,7 @@ async def task_cancel(_, message: Message):
     if mid not in tasks:
         return await m.delete()
 
-    tasks[mid].cancel()
-    del tasks[mid]
+    await rm_task(mid)
     await eor(message, text=f"{arrow(m)} Task cancelled")
 
 
@@ -82,11 +103,7 @@ async def task_cancel(_, message: Message):
     & filters.command("lsTasks", prefixes=USERBOT_PREFIX)
 )
 async def task_list(_, message: Message):
-    global tasks
-    for key, value in tasks.items():
-        if value.done():
-            del tasks[key]
-
+    await rm_task()
     if not tasks:
         return await eor(
             message,
@@ -115,11 +132,6 @@ async def executor(client, message: Message):
 
     m = message
     p = print
-    loop = get_event_loop()
-
-    for key, value in tasks.items():
-        if value.done():
-            del tasks[key]
 
     # To prevent keyboard input attacks
     if m.reply_to_message:
@@ -133,16 +145,18 @@ async def executor(client, message: Message):
     redirected_error = sys.stderr = StringIO()
     stdout, stderr, exc = None, None, None
     try:
-        task = loop.create_task(aexec(cmd, client, message))
-
-        # Save it in tasks, so we can cancel it with .cancelTask
-        tasks[m.message_id] = task
+        task = await add_task(
+            aexec,
+            m.message_id,
+            cmd,
+            client,
+            message,
+        )
         await task
     except Exception as e:
         exc = str(e)
 
-    if m.message_id in tasks:
-        del tasks[m.message_id]
+    await rm_task()
 
     stdout = redirected_output.getvalue()
     stderr = redirected_error.getvalue()
@@ -215,7 +229,7 @@ async def shellrunner(client, message: Message):
         r = message.reply_to_message
         if r.reply_markup:
             if isinstance(r.reply_markup, ReplyKeyboardMarkup):
-                return await message.edit("INSECURE!")
+                return await eor(message, text="INSECURE!")
 
     text = message.text.split(None, 1)[1]
     if "\n" in text:

+ 0 - 3
wbb/utils/aiodownloader/__init__.py

@@ -1,3 +0,0 @@
-from .downloader import Handler
-
-Handler = Handler

+ 0 - 158
wbb/utils/aiodownloader/downloader.py

@@ -1,158 +0,0 @@
-import asyncio
-import os
-from typing import Optional
-
-import aiofiles
-import aiohttp
-
-
-class DownloadJob:
-    """
-    Download Job
-
-    :param file_url: url where the file is located
-    :param session: aiohttp session to be used on the job
-    :param file_name: name to be used on the file. Defaults to the last part of the url
-    :param save_path: dir where the file should be saved. Defaults to the current dir
-    """
-
-    def __init__(
-        self,
-        session: aiohttp.ClientSession,
-        file_url: str,
-        save_path: Optional[str] = None,
-        chunk_size: Optional[int] = 1024,
-    ):
-
-        self.file_url = file_url
-        self._session = session
-        self._chunk_size = chunk_size
-
-        self.file_name = file_url.split("/")[~0][0:230]
-        self.file_path = (
-            os.path.join(save_path, self.file_name)
-            if save_path
-            else self.file_name
-        )
-
-        self.completed = False
-        self.progress = 0
-        self.size = 0
-
-    async def get_size(self) -> int:
-        """
-        Gets the content-length of the file from the file url.
-        :return: the files size in bytes
-        """
-        if not self.size:
-            async with self._session.get(self.file_url) as resp:
-                if 200 <= resp.status < 300:
-                    self.size = int(resp.headers["Content-Length"])
-
-                else:
-                    raise aiohttp.errors.HttpProcessingError(
-                        message=f"There was a problem processing {self.file_url}",
-                        code=resp.status,
-                    )
-
-        return self.size
-
-    def _downloaded(self, chunk_size):
-        """
-        Method to be called when a chunk of the file is downloaded. It updates the
-        progress, adds the size to the _advance variable and checks if the download
-        is completed
-        """
-        self.progress += chunk_size
-
-    async def download(self):
-        """
-        Downloads the file from the given url to a file in the given path.
-        """
-
-        async with self._session.get(self.file_url) as resp:
-            # Checkning the response code
-            if 200 <= resp.status < 300:
-                # Saving the data to the file chunk by chunk.
-                async with aiofiles.open(
-                    self.file_path, "wb"
-                ) as file:
-
-                    # Downloading the file using the aiohttp.StreamReader
-                    async for data in resp.content.iter_chunked(
-                        self._chunk_size
-                    ):
-                        await file.write(data)
-                        self._downloaded(self._chunk_size)
-
-                self.completed = True
-                return self
-
-            else:
-                raise aiohttp.errors.HttpProcessingError(
-                    message=f"There was a problem processing {self.file_url}",
-                    code=resp.status,
-                )
-
-
-class Handler:
-    """
-    Top level interface with the downloader. It creates the download jobs and handles them.
-
-    :param loop: asyncio loop. if not provided asyncio.get_event_loop() will be used
-    :param session: aiohttp session to be used on all downloads. If not provided it will be
-    created
-    :param chunk_size: chunk bytes sizes to get from the file source. Defaults to 1024 bytes
-    """
-
-    def __init__(
-        self,
-        loop: Optional[asyncio.BaseEventLoop] = None,
-        session: Optional[aiohttp.ClientSession] = None,
-        chunk_size: Optional[int] = 1024,
-    ):
-
-        self._loop = loop or asyncio.get_event_loop()
-        self._session = session or aiohttp.ClientSession(
-            loop=self._loop
-        )
-        self._chunk_size = chunk_size
-
-    def _job_factory(
-        self, file_url: str, save_path: Optional[str] = None
-    ) -> DownloadJob:
-        """
-        Shortcut for creating a download job. It adds the session and the chunk size.
-        :param file_url: url where the file is located
-        :param save_path: save path for the download
-        :return:
-        """
-        return DownloadJob(
-            self._session, file_url, save_path, self._chunk_size
-        )
-
-    async def download(
-        self, url: str, save_path: Optional[str] = None
-    ) -> DownloadJob:
-        """
-        Downloads a bulk of files from the given list of urls to the given path.
-
-        :param files_url: list of urls where the files are located
-        :param save_path: path to be used for saving the files. Defaults to the current dir
-        """
-
-        job = self._job_factory(url, save_path=save_path)
-
-        task = asyncio.ensure_future(job.download())
-
-        await task
-        file_name = url.split("/")[-1]
-        file_name = (
-            file_name[0:230] if len(file_name) > 230 else file_name
-        )
-        path = (
-            os.getcwd() + "/" + file_name
-            if not save_path
-            else save_path
-        )
-        return path

+ 93 - 0
wbb/utils/downloader.py

@@ -0,0 +1,93 @@
+from inspect import iscoroutinefunction
+from os.path import abspath as absolute_path
+from time import time
+
+import aiofiles
+
+from wbb import aiohttpsession as session
+from wbb.modules.userbot import add_task, rm_task, tasks
+
+
+def ensure_status(status_code: int):
+    if status_code < 200 or status_code >= 300:
+        raise Exception(f"HttpProcessingError: {status_code}")
+
+
+async def download_url(
+    url,
+    file_path,
+    chunk_size,
+    progress_func,
+):
+    global tasks
+
+    file_path = file_path or url.split("/")[-1][:20]
+
+    progress = 0
+    total_size = 0
+
+    # Check if the provided progress function is
+    # async for sync in order to call it correctly.
+    if progress_func:
+        p_f_args = progress_func[1]
+        progress_func = progress_func[0]
+        is_async = iscoroutinefunction(progress_func)
+
+    # Check if the server responds
+    async with session.get(url) as resp:
+        ensure_status(resp.status)
+        total_size = int(resp.headers["Content-Length"])
+
+    async with session.get(url) as response:
+        ensure_status(response.status)
+
+        async with aiofiles.open(file_path, "wb") as f:
+
+            # Save content in file using aiohttp streamReader.
+            async for chunk in response.content.iter_chunked(
+                chunk_size
+            ):
+                await f.write(chunk)
+
+                # Call the progress func on each chunk downloaded
+                progress += chunk_size
+                if progress_func:
+
+                    if is_async:
+                        await progress_func(
+                            progress, total_size, *p_f_args
+                        )
+                    else:
+                        progress_func(progress, total_size, *p_f_args)
+
+    return absolute_path(file_path)
+
+
+async def download(
+    url: str,
+    file_path: str = None,
+    chunk_size: int = 1024,
+    progress_func=None,
+    task_id: int = int(time()),
+):
+    """
+    :url: url where the file is located
+    :file_path: path/to/file
+    :chunk_size: size of a single chunk
+    """
+    global tasks
+
+    # Create a task and add it to main tasks dict
+    # So we can cancel it using .cancelTask
+
+    task = await add_task(
+        download_url,
+        task_id,
+        url=url,
+        file_path=file_path,
+        chunk_size=chunk_size,
+        progress_func=progress_func,
+    )
+    await task
+    await rm_task(task_id)
+    return task.result()

+ 29 - 19
wbb/utils/functions.py

@@ -24,38 +24,54 @@ SOFTWARE.
 from asyncio import gather, get_running_loop
 from datetime import datetime, timedelta
 from io import BytesIO
-from math import atan2, cos, radians, sin, sqrt
+from math import atan2, cos, floor, radians, sin, sqrt
 from os import execvp
 from random import randint
 from re import findall
 from re import sub as re_sub
 from sys import executable
+from time import time
 
 import aiofiles
 import aiohttp
 import speedtest
 from PIL import Image, ImageDraw, ImageFilter, ImageFont
 from pyrogram.types import Message
-from wget import download
 
 from wbb import aiohttpsession as aiosession
-from wbb.utils import aiodownloader
+from wbb.modules.userbot import eor
 from wbb.utils.dbfunctions import start_restart_stage
 from wbb.utils.http import get
 
-"""
-Just import 'downloader' anywhere and do downloader.download() to
-download file from a given url
-"""
-downloader = aiodownloader.Handler()
 
-# Another downloader, but with wget
+async def progress(
+    current: int,
+    total: int,
+    start: int,
+    task_id: int,
+    message: Message,
+):
+    percentage = current / total * 100
+    elapsed = time() - start
+    speed = (current / 1000000) / elapsed  # In MB/s
+    eta = (100 / percentage * elapsed) / 60  # In Minutes
+
+    # Only edit at every 5 seconds
+    if round(elapsed % 5.00) != 0 and current != total:
+        return
+
+    round_pct = floor(percentage / 10)
 
+    bar = ("▰" * round_pct) + ("▱" * (10 - round_pct))
 
-async def download_url(url: str):
-    loop = get_running_loop()
-    file = await loop.run_in_executor(None, download, url)
-    return file
+    text = f"""
+{bar}
+**Progress:** {int(percentage)}%
+**Speed:** {round(speed, 2)}MB/s
+**ETA:** {round(eta, 2)}m
+**Task ID:** {task_id}
+"""
+    await eor(message, text=text)
 
 
 async def restart(m: Message):
@@ -128,12 +144,6 @@ def test_speedtest():
     return [speed_convert(download), speed_convert(upload), info]
 
 
-async def file_size_from_url(url: str) -> int:
-    async with aiosession.head(url) as resp:
-        size = int(resp.headers["content-length"])
-    return size
-
-
 async def get_http_status_code(url: str) -> int:
     async with aiosession.head(url) as resp:
         return resp.status

+ 11 - 5
wbb/utils/misc.py

@@ -80,15 +80,21 @@ def paginate_modules(page_n, module_dict, prefix, chat=None):
             )
         )
 
-    max_num_pages = ceil(len(pairs) / 7)
+    COLUMN_SIZE = 4
+
+    max_num_pages = ceil(len(pairs) / COLUMN_SIZE)
     modulo_page = page_n % max_num_pages
 
     # can only have a certain amount of buttons side by side
-    if len(pairs) > 7:
-        pairs = pairs[modulo_page * 7 : 7 * (modulo_page + 1)] + [
+    if len(pairs) > COLUMN_SIZE:
+        pairs = pairs[
+            modulo_page
+            * COLUMN_SIZE : COLUMN_SIZE
+            * (modulo_page + 1)
+        ] + [
             (
                 EqInlineKeyboardButton(
-                    "<",
+                    "❮",
                     callback_data="{}_prev({})".format(
                         prefix, modulo_page
                     ),
@@ -100,7 +106,7 @@ def paginate_modules(page_n, module_dict, prefix, chat=None):
                     ),
                 ),
                 EqInlineKeyboardButton(
-                    ">",
+                    "❯",
                     callback_data="{}_next({})".format(
                         prefix, modulo_page
                     ),