Переглянути джерело

fix: multiple issues in urls + catastrophic backtracking in Regex, freezes the bot

TheHamkerCat 7 місяців тому
батько
коміт
2428bf104e
5 змінених файлів з 158 додано та 73 видалено
  1. 1 0
      Dockerfile
  2. 5 37
      wbb/modules/misc.py
  3. 76 18
      wbb/modules/regex.py
  4. 31 2
      wbb/modules/rss.py
  5. 45 16
      wbb/utils/functions.py

+ 1 - 0
Dockerfile

@@ -10,6 +10,7 @@ ENV PYTHONUNBUFFERED=1
 RUN apt-get update -y && apt-get install -y --no-install-recommends \
 RUN apt-get update -y && apt-get install -y --no-install-recommends \
     curl ca-certificates \
     curl ca-certificates \
     git gcc build-essential \
     git gcc build-essential \
+    iputils-ping \
     && rm -rf /var/lib/apt/lists/*
     && rm -rf /var/lib/apt/lists/*
 
 
 # install uv
 # install uv

+ 5 - 37
wbb/modules/misc.py

@@ -21,6 +21,7 @@ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
 OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
 OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
 SOFTWARE.
 SOFTWARE.
 """
 """
+
 import re
 import re
 import secrets
 import secrets
 import string
 import string
@@ -34,8 +35,6 @@ from wbb import SUDOERS, USERBOT_PREFIX, app, app2, arq, eor
 from wbb.core.decorators.errors import capture_err
 from wbb.core.decorators.errors import capture_err
 from wbb.utils import random_line
 from wbb.utils import random_line
 from wbb.utils.http import get
 from wbb.utils.http import get
-from wbb.utils.json_prettify import json_prettify
-from wbb.utils.pastebin import paste
 
 
 __MODULE__ = "Misc"
 __MODULE__ = "Misc"
 __HELP__ = """
 __HELP__ = """
@@ -61,9 +60,6 @@ __HELP__ = """
     Translate A Message
     Translate A Message
     Ex: /tr en
     Ex: /tr en
 
 
-/json [URL]
-    Get parsed JSON response from a rest API.
-
 /arq
 /arq
     Statistics Of ARQ API.
     Statistics Of ARQ API.
 
 
@@ -166,9 +162,7 @@ async def rtfm(_, message):
     await message.delete()
     await message.delete()
     if not message.reply_to_message:
     if not message.reply_to_message:
         return await message.reply_text("Reply To A Message lol")
         return await message.reply_text("Reply To A Message lol")
-    await message.reply_to_message.reply_text(
-        "Are You Lost? READ THE FUCKING DOCS!"
-    )
+    await message.reply_to_message.reply_text("Are You Lost? READ THE FUCKING DOCS!")
 
 
 
 
 @app.on_message(filters.command("runs"))
 @app.on_message(filters.command("runs"))
@@ -222,16 +216,12 @@ async def getid(client, message):
 @capture_err
 @capture_err
 async def random(_, message):
 async def random(_, message):
     if len(message.command) != 2:
     if len(message.command) != 2:
-        return await message.reply_text(
-            '"/random" Needs An Argurment.' " Ex: `/random 5`"
-        )
+        return await message.reply_text('"/random" Needs An Argurment. Ex: `/random 5`')
     length = message.text.split(None, 1)[1]
     length = message.text.split(None, 1)[1]
     try:
     try:
         if 1 < int(length) < 1000:
         if 1 < int(length) < 1000:
             alphabet = string.ascii_letters + string.digits
             alphabet = string.ascii_letters + string.digits
-            password = "".join(
-                secrets.choice(alphabet) for i in range(int(length))
-            )
+            password = "".join(secrets.choice(alphabet) for i in range(int(length)))
             await message.reply_text(f"`{password}`")
             await message.reply_text(f"`{password}`")
         else:
         else:
             await message.reply_text("Specify A Length Between 1-1000")
             await message.reply_text("Specify A Length Between 1-1000")
@@ -265,30 +255,8 @@ async def tr(_, message):
     await message.reply_text(result.result.translatedText)
     await message.reply_text(result.result.translatedText)
 
 
 
 
-@app.on_message(filters.command("json"))
-@capture_err
-async def json_fetch(_, message):
-    if len(message.command) != 2:
-        return await message.reply_text("/json [URL]")
-    url = message.text.split(None, 1)[1]
-    m = await message.reply_text("Fetching")
-    try:
-        data = await get(url)
-        data = await json_prettify(data)
-        if len(data) < 4090:
-            await m.edit(data)
-        else:
-            link = await paste(data)
-            await m.edit(
-                f"[OUTPUT_TOO_LONG]({link})",
-                disable_web_page_preview=True,
-            )
-    except Exception as e:
-        await m.edit(str(e))
-
-
 @app.on_message(filters.command(["kickme", "banme"]))
 @app.on_message(filters.command(["kickme", "banme"]))
 async def kickbanme(_, message):
 async def kickbanme(_, message):
     await message.reply_text(
     await message.reply_text(
-        "Haha, it doesn't work that way, You're stuck with everyone here."
+        "Haha, it doesn't work that way, You're stuck here with everyone."
     )
     )

+ 76 - 18
wbb/modules/regex.py

@@ -1,6 +1,7 @@
 # https://github.com/PaulSonOfLars/tgbot/blob/master/tg_bot/modules/sed.py
 # https://github.com/PaulSonOfLars/tgbot/blob/master/tg_bot/modules/sed.py
+import asyncio
+import multiprocessing as mp
 import re
 import re
-import sre_constants
 
 
 from pyrogram import filters
 from pyrogram import filters
 
 
@@ -11,6 +12,62 @@ __MODULE__ = "Sed"
 __HELP__ = "**Usage:**\ns/foo/bar"
 __HELP__ = "**Usage:**\ns/foo/bar"
 
 
 DELIMITERS = ("/", ":", "|", "_")
 DELIMITERS = ("/", ":", "|", "_")
+REGEX_TIMEOUT_SECONDS = 5
+
+
+def _regex_sub_worker(
+    pattern: str,
+    replacement: str,
+    source_text: str,
+    ignore_case: bool,
+    replace_all: bool,
+    result_queue,
+):
+    flags = re.I if ignore_case else 0
+    count = 0 if replace_all else 1
+
+    try:
+        result = re.sub(pattern, replacement, source_text, count=count, flags=flags)
+        result_queue.put(("ok", result))
+    except re.error:
+        result_queue.put(("regex_error", ""))
+
+
+def run_regex_with_timeout(
+    pattern: str,
+    replacement: str,
+    source_text: str,
+    ignore_case: bool,
+    replace_all: bool,
+) -> str:
+    result_queue = mp.Queue(maxsize=1)
+    process = mp.Process(
+        target=_regex_sub_worker,
+        args=(
+            pattern,
+            replacement,
+            source_text,
+            ignore_case,
+            replace_all,
+            result_queue,
+        ),
+    )
+    process.start()
+    process.join(REGEX_TIMEOUT_SECONDS)
+
+    if process.is_alive():
+        process.terminate()
+        process.join()
+        raise asyncio.TimeoutError
+
+    if result_queue.empty():
+        return ""
+
+    status, result = result_queue.get()
+    if status == "regex_error":
+        raise re.error("invalid regex")
+
+    return result
 
 
 
 
 @app.on_message(
 @app.on_message(
@@ -31,31 +88,31 @@ async def sed(_, message):
             to_fix = message.reply_to_message.caption
             to_fix = message.reply_to_message.caption
         else:
         else:
             return
             return
-        try:
-            repl, repl_with, flags = sed_result
-        except Exception:
+        if not sed_result:
             return
             return
+        repl, repl_with, flags = sed_result
 
 
         if not repl:
         if not repl:
             return await message.reply_text(
             return await message.reply_text(
-                "You're trying to replace... " "nothing with something?"
+                "You're trying to replace... nothing with something?"
             )
             )
 
 
         try:
         try:
             if infinite_checker(repl):
             if infinite_checker(repl):
                 return await message.reply_text("Nice try -_-")
                 return await message.reply_text("Nice try -_-")
 
 
-            if "i" in flags and "g" in flags:
-                text = re.sub(repl, repl_with, to_fix, flags=re.I).strip()
-            elif "i" in flags:
-                text = re.sub(
-                    repl, repl_with, to_fix, count=1, flags=re.I
-                ).strip()
-            elif "g" in flags:
-                text = re.sub(repl, repl_with, to_fix).strip()
-            else:
-                text = re.sub(repl, repl_with, to_fix, count=1).strip()
-        except sre_constants.error:
+            text = await asyncio.to_thread(
+                run_regex_with_timeout,
+                repl,
+                repl_with,
+                to_fix,
+                "i" in flags,
+                "g" in flags,
+            )
+            text = text.strip()
+        except asyncio.TimeoutError:
+            return await message.reply_text("Regex took too long to compute.")
+        except re.error:
             return
             return
 
 
         # empty string errors -_-
         # empty string errors -_-
@@ -75,8 +132,9 @@ def infinite_checker(repl):
         r"\(.{1,}\)\{.{1,}(,)?\}\(.*\)(\+|\* |\{.*\})",
         r"\(.{1,}\)\{.{1,}(,)?\}\(.*\)(\+|\* |\{.*\})",
     ]
     ]
     for match in regex:
     for match in regex:
-        status = re.search(match, repl)
-        return bool(status)
+        if re.search(match, repl):
+            return True
+    return False
 
 
 
 
 def separate_sed(sed_string):
 def separate_sed(sed_string):

+ 31 - 2
wbb/modules/rss.py

@@ -19,7 +19,11 @@ from wbb.utils.dbfunctions import (
     remove_rss_feed,
     remove_rss_feed,
     update_rss_feed,
     update_rss_feed,
 )
 )
-from wbb.utils.functions import get_http_status_code, get_urls_from_text
+from wbb.utils.functions import (
+    get_http_status_code,
+    get_urls_from_text,
+    is_safe_url,
+)
 from wbb.utils.rss import Feed
 from wbb.utils.rss import Feed
 
 
 __MODULE__ = "RSS"
 __MODULE__ = "RSS"
@@ -34,6 +38,13 @@ __HELP__ = f"""
 """
 """
 
 
 
 
+def get_parsed_feed_url(parsed, fallback: str) -> str:
+    href = parsed.get("href")
+    if isinstance(href, str) and href:
+        return href
+    return fallback
+
+
 async def rss_worker():
 async def rss_worker():
     log.info("RSS Worker started")
     log.info("RSS Worker started")
     while not await sleep(RSS_DELAY):
     while not await sleep(RSS_DELAY):
@@ -47,9 +58,20 @@ async def rss_worker():
             chat = _feed["chat_id"]
             chat = _feed["chat_id"]
             try:
             try:
                 url = _feed["url"]
                 url = _feed["url"]
+                if not is_safe_url(url):
+                    await remove_rss_feed(chat)
+                    log.info(f"Removed RSS Feed from {chat} (Unsafe URL)")
+                    continue
+
                 last_title = _feed.get("last_title")
                 last_title = _feed.get("last_title")
 
 
                 parsed = await loop.run_in_executor(None, parse, url)
                 parsed = await loop.run_in_executor(None, parse, url)
+                final_url = get_parsed_feed_url(parsed, url)
+                if not is_safe_url(final_url):
+                    await remove_rss_feed(chat)
+                    log.info(f"Removed RSS Feed from {chat} (Unsafe redirect)")
+                    continue
+
                 feed = Feed(parsed)
                 feed = Feed(parsed)
 
 
                 if feed.title == last_title:
                 if feed.title == last_title:
@@ -91,6 +113,9 @@ async def add_feed_func(_, m: Message):
         return await m.reply("[ERROR]: Invalid URL")
         return await m.reply("[ERROR]: Invalid URL")
 
 
     url = urls[0]
     url = urls[0]
+    if not is_safe_url(url):
+        return await m.reply("[ERROR]: URL is not allowed (SSRF protection).")
+
     status = await get_http_status_code(url)
     status = await get_http_status_code(url)
     if status != 200:
     if status != 200:
         return await m.reply("[ERROR]: Invalid Url")
         return await m.reply("[ERROR]: Invalid Url")
@@ -99,6 +124,10 @@ async def add_feed_func(_, m: Message):
     try:
     try:
         loop = get_event_loop()
         loop = get_event_loop()
         parsed = await loop.run_in_executor(None, parse, url)
         parsed = await loop.run_in_executor(None, parse, url)
+        final_url = get_parsed_feed_url(parsed, url)
+        if not is_safe_url(final_url):
+            return await m.reply("[ERROR]: URL is not allowed (SSRF protection).")
+
         feed = Feed(parsed)
         feed = Feed(parsed)
     except Exception:
     except Exception:
         return await m.reply(ns)
         return await m.reply(ns)
@@ -112,7 +141,7 @@ async def add_feed_func(_, m: Message):
         await m.reply(feed.parsed(), disable_web_page_preview=True)
         await m.reply(feed.parsed(), disable_web_page_preview=True)
     except Exception:
     except Exception:
         return await m.reply(ns)
         return await m.reply(ns)
-    await add_rss_feed(chat_id, parsed.url, feed.title)
+    await add_rss_feed(chat_id, final_url, feed.title)
 
 
 
 
 @app.on_message(filters.command("rm_feed"))
 @app.on_message(filters.command("rm_feed"))

+ 45 - 16
wbb/utils/functions.py

@@ -23,6 +23,8 @@ SOFTWARE.
 """
 """
 
 
 import asyncio
 import asyncio
+import ipaddress
+import socket
 
 
 from asyncio import gather
 from asyncio import gather
 from datetime import datetime, timedelta
 from datetime import datetime, timedelta
@@ -30,10 +32,11 @@ from io import BytesIO
 from math import atan2, cos, radians, sin, sqrt
 from math import atan2, cos, radians, sin, sqrt
 from os import execvp
 from os import execvp
 from random import randint
 from random import randint
-from re import findall, search
+from re import findall
 from re import sub as re_sub
 from re import sub as re_sub
 from sys import executable
 from sys import executable
 from typing import Dict
 from typing import Dict
+from urllib.parse import urlparse
 
 
 import aiofiles
 import aiofiles
 import speedtest
 import speedtest
@@ -64,9 +67,7 @@ def generate_captcha():
     def rndColor2():
     def rndColor2():
         return (randint(32, 127), randint(32, 127), randint(32, 127))
         return (randint(32, 127), randint(32, 127), randint(32, 127))
 
 
-    wrong_answers = [
-        "".join(gen_letter() for _ in range(4)) for _ in range(8)
-    ]
+    wrong_answers = ["".join(gen_letter() for _ in range(4)) for _ in range(8)]
 
 
     width, height = 320, 100
     width, height = 320, 100
     correct_answer = ""
     correct_answer = ""
@@ -89,8 +90,8 @@ def generate_captcha():
         draw.text((60 * t + 50, 15), letter, font=font, fill=rndColor2())
         draw.text((60 * t + 50, 15), letter, font=font, fill=rndColor2())
 
 
     image = image.filter(ImageFilter.BLUR)
     image = image.filter(ImageFilter.BLUR)
-    
-    buf = BytesIO() # on memory
+
+    buf = BytesIO()  # on memory
     image.save(buf, "JPEG")
     image.save(buf, "JPEG")
     buf.name = "captcha.jpg"
     buf.name = "captcha.jpg"
     buf.seek(0)
     buf.seek(0)
@@ -156,6 +157,40 @@ async def calc_distance_from_ip(ip1: str, ip2: str) -> float:
     return distance
     return distance
 
 
 
 
+def is_safe_url(url: str) -> bool:
+    """
+    Check if a URL is safe to fetch (prevents SSRF).
+    It ensures the URL uses http/https and that all resolved IPs are public.
+    """
+    try:
+        parsed = urlparse(url.strip())
+        if parsed.scheme not in {"http", "https"}:
+            return False
+
+        hostname = parsed.hostname
+        if not hostname:
+            return False
+
+        port = parsed.port
+        if port is None:
+            port = 443 if parsed.scheme == "https" else 80
+
+        # Resolve all addresses (IPv4 + IPv6). URL is safe only if all are public.
+        address_info = socket.getaddrinfo(hostname, port, type=socket.SOCK_STREAM)
+        resolved_ips = {item[4][0] for item in address_info}
+        if not resolved_ips:
+            return False
+
+        for ip_addr in resolved_ips:
+            ip = ipaddress.ip_address(ip_addr)
+            if not ip.is_global:
+                return False
+
+        return True
+    except (ValueError, OSError):
+        return False
+
+
 def get_urls_from_text(text: str) -> bool:
 def get_urls_from_text(text: str) -> bool:
     regex = r"""(?i)\b((?:https?://|www\d{0,3}[.]|[a-z0-9.\-]
     regex = r"""(?i)\b((?:https?://|www\d{0,3}[.]|[a-z0-9.\-]
                 [.][a-z]{2,4}/)(?:[^\s()<>]+|\(([^\s()<>]+|(
                 [.][a-z]{2,4}/)(?:[^\s()<>]+|\(([^\s()<>]+|(
@@ -349,19 +384,13 @@ async def get_data_and_name(replied_message, message):
     if len(text) > 1:
     if len(text) > 1:
         name = text[0]
         name = text[0]
         data = text[1].strip()
         data = text[1].strip()
-        if replied_message and (
-            replied_message.sticker or replied_message.video_note
-        ):
+        if replied_message and (replied_message.sticker or replied_message.video_note):
             data = None
             data = None
     else:
     else:
-        if replied_message and (
-            replied_message.sticker or replied_message.video_note
-        ):
+        if replied_message and (replied_message.sticker or replied_message.video_note):
             data = None
             data = None
         elif (
         elif (
-            replied_message
-            and not replied_message.text
-            and not replied_message.caption
+            replied_message and not replied_message.text and not replied_message.caption
         ):
         ):
             data = None
             data = None
         else:
         else:
@@ -383,7 +412,7 @@ async def get_data_and_name(replied_message, message):
 
 
 async def get_specific_usernames(client, user_ids: list) -> Dict[int, str]:
 async def get_specific_usernames(client, user_ids: list) -> Dict[int, str]:
     def _fetch_users():
     def _fetch_users():
-        ids_str = ','.join(str(uid) for uid in user_ids)
+        ids_str = ",".join(str(uid) for uid in user_ids)
 
 
         with client.storage.conn:
         with client.storage.conn:
             query = f"""
             query = f"""