functions.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436
  1. """
  2. MIT License
  3. Copyright (c) 2024 TheHamkerCat
  4. Permission is hereby granted, free of charge, to any person obtaining a copy
  5. of this software and associated documentation files (the "Software"), to deal
  6. in the Software without restriction, including without limitation the rights
  7. to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
  8. copies of the Software, and to permit persons to whom the Software is
  9. furnished to do so, subject to the following conditions:
  10. The above copyright notice and this permission notice shall be included in all
  11. copies or substantial portions of the Software.
  12. THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
  13. IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
  14. FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
  15. AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
  16. LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
  17. OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
  18. SOFTWARE.
  19. """
  20. import asyncio
  21. import ipaddress
  22. import socket
  23. from asyncio import gather
  24. from datetime import datetime, timedelta
  25. from io import BytesIO
  26. from math import atan2, cos, radians, sin, sqrt
  27. from os import execvp
  28. from random import randint
  29. from re import findall
  30. from re import sub as re_sub
  31. from sys import executable
  32. from typing import Dict
  33. from urllib.parse import urlparse
  34. import aiofiles
  35. import speedtest
  36. from PIL import Image, ImageDraw, ImageFilter, ImageFont
  37. from pyrogram import errors
  38. from pyrogram.enums import MessageEntityType
  39. from pyrogram.types import Message
  40. from wbb import aiohttpsession as aiosession
  41. from wbb.utils.dbfunctions import start_restart_stage
  42. from wbb.utils.http import get, post
  43. async def restart(m: Message):
  44. if m:
  45. await start_restart_stage(m.chat.id, m.id)
  46. execvp(executable, [executable, "-m", "wbb"])
  47. def generate_captcha():
  48. # Generate one letter
  49. def gen_letter():
  50. return chr(randint(65, 90))
  51. def rndColor():
  52. return (randint(64, 255), randint(64, 255), randint(64, 255))
  53. def rndColor2():
  54. return (randint(32, 127), randint(32, 127), randint(32, 127))
  55. wrong_answers = ["".join(gen_letter() for _ in range(4)) for _ in range(8)]
  56. width, height = 320, 100
  57. correct_answer = ""
  58. font = ImageFont.truetype("assets/arial.ttf", 55)
  59. # White background
  60. image = Image.new("RGB", (width, height), (255, 255, 255))
  61. draw = ImageDraw.Draw(image)
  62. # Noise (draw random lines instead of painting every pixel)
  63. for _ in range(120):
  64. x1, y1 = randint(0, width), randint(0, height)
  65. x2, y2 = randint(0, width), randint(0, height)
  66. draw.line((x1, y1, x2, y2), fill=rndColor(), width=1)
  67. # Text
  68. for t in range(4):
  69. letter = gen_letter()
  70. correct_answer += letter
  71. draw.text((60 * t + 50, 15), letter, font=font, fill=rndColor2())
  72. image = image.filter(ImageFilter.BLUR)
  73. buf = BytesIO() # on memory
  74. image.save(buf, "JPEG")
  75. buf.name = "captcha.jpg"
  76. buf.seek(0)
  77. return [buf, correct_answer, wrong_answers]
  78. def test_speedtest():
  79. def speed_convert(size):
  80. power = 2**10
  81. zero = 0
  82. units = {0: "", 1: "Kb/s", 2: "Mb/s", 3: "Gb/s", 4: "Tb/s"}
  83. while size > power:
  84. size /= power
  85. zero += 1
  86. return f"{round(size, 2)} {units[zero]}"
  87. speed = speedtest.Speedtest()
  88. info = speed.get_best_server()
  89. download = speed.download()
  90. upload = speed.upload()
  91. return [speed_convert(download), speed_convert(upload), info]
  92. async def get_http_status_code(url: str) -> int:
  93. async with aiosession.head(url) as resp:
  94. return resp.status
  95. async def make_carbon(code):
  96. url = "https://carbonara.solopov.dev/api/cook"
  97. async with aiosession.post(url, json={"code": code}) as resp:
  98. image = BytesIO(await resp.read())
  99. image.name = "carbon.png"
  100. return image
  101. async def transfer_sh(file_or_message):
  102. if isinstance(file_or_message, Message):
  103. file_or_message = await file_or_message.download()
  104. file = file_or_message
  105. async with aiofiles.open(file, "rb") as f:
  106. params = {file: await f.read()}
  107. resp = await post("https://transfer.sh/", data=params)
  108. url = resp.strip()
  109. return url
  110. async def calc_distance_from_ip(ip1: str, ip2: str) -> float:
  111. Radius_Earth = 6371.0088
  112. data1, data2 = await gather(
  113. get(f"http://ipinfo.io/{ip1}"),
  114. get(f"http://ipinfo.io/{ip2}"),
  115. )
  116. lat1, lon1 = data1["loc"].split(",")
  117. lat2, lon2 = data2["loc"].split(",")
  118. lat1, lon1 = radians(float(lat1)), radians(float(lon1))
  119. lat2, lon2 = radians(float(lat2)), radians(float(lon2))
  120. dlon = lon2 - lon1
  121. dlat = lat2 - lat1
  122. a = sin(dlat / 2) ** 2 + cos(lat1) * cos(lat2) * sin(dlon / 2) ** 2
  123. c = 2 * atan2(sqrt(a), sqrt(1 - a))
  124. distance = Radius_Earth * c
  125. return distance
  126. def is_safe_url(url: str) -> bool:
  127. """
  128. Check if a URL is safe to fetch (prevents SSRF).
  129. It ensures the URL uses http/https and that all resolved IPs are public.
  130. """
  131. try:
  132. parsed = urlparse(url.strip())
  133. if parsed.scheme not in {"http", "https"}:
  134. return False
  135. hostname = parsed.hostname
  136. if not hostname:
  137. return False
  138. port = parsed.port
  139. if port is None:
  140. port = 443 if parsed.scheme == "https" else 80
  141. # Resolve all addresses (IPv4 + IPv6). URL is safe only if all are public.
  142. address_info = socket.getaddrinfo(hostname, port, type=socket.SOCK_STREAM)
  143. resolved_ips = {item[4][0] for item in address_info}
  144. if not resolved_ips:
  145. return False
  146. for ip_addr in resolved_ips:
  147. ip = ipaddress.ip_address(ip_addr)
  148. if not ip.is_global:
  149. return False
  150. return True
  151. except (ValueError, OSError):
  152. return False
  153. def get_urls_from_text(text: str) -> bool:
  154. regex = r"""(?i)\b((?:https?://|www\d{0,3}[.]|[a-z0-9.\-]
  155. [.][a-z]{2,4}/)(?:[^\s()<>]+|\(([^\s()<>]+|(
  156. \([^\s()<>]+\)))*\))+(?:\(([^\s()<>]+|(\([^\
  157. ()<>]+\)))*\)|[^\s`!()\[\]{};:'".,<>?«»“”‘’]))""".strip()
  158. return [x[0] for x in findall(regex, str(text))]
  159. async def time_converter(message: Message, time_value: str) -> datetime:
  160. unit = ["m", "h", "d"] # m == minutes | h == hours | d == days
  161. check_unit = "".join(list(filter(time_value[-1].lower().endswith, unit)))
  162. currunt_time = datetime.now()
  163. time_digit = time_value[:-1]
  164. if not time_digit.isdigit():
  165. return await message.reply_text("时间格式不正确,请使用 10m、2h 或 1d。")
  166. if check_unit == "m":
  167. temp_time = currunt_time + timedelta(minutes=int(time_digit))
  168. elif check_unit == "h":
  169. temp_time = currunt_time + timedelta(hours=int(time_digit))
  170. elif check_unit == "d":
  171. temp_time = currunt_time + timedelta(days=int(time_digit))
  172. else:
  173. return await message.reply_text("时间格式不正确,请使用 10m、2h 或 1d。")
  174. return temp_time
  175. async def extract_userid(message, text: str):
  176. """
  177. NOT TO BE USED OUTSIDE THIS FILE
  178. """
  179. def is_int(text: str):
  180. try:
  181. int(text)
  182. except ValueError:
  183. return False
  184. return True
  185. text = text.strip()
  186. if is_int(text):
  187. return int(text)
  188. entities = message.entities
  189. app = message._client
  190. if len(entities) < 2:
  191. return (await app.get_users(text)).id
  192. entity = entities[1]
  193. if entity.type == MessageEntityType.MENTION:
  194. return (await app.get_users(text)).id
  195. if entity.type == MessageEntityType.TEXT_MENTION:
  196. return entity.user.id
  197. return None
  198. async def extract_user_and_reason(message, sender_chat=False):
  199. args = message.text.strip().split()
  200. text = message.text
  201. user = None
  202. reason = None
  203. try:
  204. if message.reply_to_message:
  205. reply = message.reply_to_message
  206. # if reply to a message and no reason is given
  207. if not reply.from_user:
  208. if (
  209. reply.sender_chat
  210. and reply.sender_chat != message.chat.id
  211. and sender_chat
  212. ):
  213. id_ = reply.sender_chat.id
  214. else:
  215. return None, None
  216. else:
  217. id_ = reply.from_user.id
  218. if len(args) < 2:
  219. reason = None
  220. else:
  221. reason = text.split(None, 1)[1]
  222. return id_, reason
  223. # if not reply to a message and no reason is given
  224. if len(args) == 2:
  225. user = text.split(None, 1)[1]
  226. return await extract_userid(message, user), None
  227. # if reason is given
  228. if len(args) > 2:
  229. user, reason = text.split(None, 2)[1:]
  230. return await extract_userid(message, user), reason
  231. return user, reason
  232. except errors.UsernameInvalid:
  233. return "", ""
  234. async def extract_user(message):
  235. return (await extract_user_and_reason(message))[0]
  236. def get_file_id_from_message(
  237. message,
  238. max_file_size=3145728,
  239. mime_types=["image/png", "image/jpeg"],
  240. ):
  241. file_id = None
  242. if message.document:
  243. if int(message.document.file_size) > max_file_size:
  244. return
  245. mime_type = message.document.mime_type
  246. if mime_types and mime_type not in mime_types:
  247. return
  248. file_id = message.document.file_id
  249. if message.sticker:
  250. if message.sticker.is_animated:
  251. if not message.sticker.thumbs:
  252. return
  253. file_id = message.sticker.thumbs[0].file_id
  254. else:
  255. file_id = message.sticker.file_id
  256. if message.photo:
  257. file_id = message.photo.file_id
  258. if message.animation:
  259. if not message.animation.thumbs:
  260. return
  261. file_id = message.animation.thumbs[0].file_id
  262. if message.video:
  263. if not message.video.thumbs:
  264. return
  265. file_id = message.video.thumbs[0].file_id
  266. return file_id
  267. def extract_text_and_keyb(ikb, text: str, row_width: int = 2):
  268. keyboard = {}
  269. try:
  270. text = text.strip()
  271. if text.startswith("`"):
  272. text = text[1:]
  273. if text.endswith("`"):
  274. text = text[:-1]
  275. if "~~" in text:
  276. text = text.replace("~~", "¤¤")
  277. text, keyb = text.split("~")
  278. if "¤¤" in text:
  279. text = text.replace("¤¤", "~~")
  280. keyb = findall(r"\[.+\,.+\]", keyb)
  281. for btn_str in keyb:
  282. btn_str = re_sub(r"[\[\]]", "", btn_str)
  283. btn_str = btn_str.split(",")
  284. btn_txt, btn_url = btn_str[0], btn_str[1].strip()
  285. if not get_urls_from_text(btn_url):
  286. continue
  287. keyboard[btn_txt] = btn_url
  288. keyboard = ikb(keyboard, row_width)
  289. except Exception:
  290. return
  291. return text, keyboard
  292. async def check_format(ikb, raw_text: str):
  293. keyb = findall(r"\[.+\,.+\]", raw_text)
  294. if keyb and not "~" in raw_text:
  295. raw_text = raw_text.replace("button=", "\n~\nbutton=")
  296. return raw_text
  297. if "~" in raw_text and keyb:
  298. if not extract_text_and_keyb(ikb, raw_text):
  299. return ""
  300. else:
  301. return raw_text
  302. else:
  303. return raw_text
  304. async def get_data_and_name(replied_message, message):
  305. text = message.text.markdown if message.text else message.caption.markdown
  306. name = text.split(None, 1)[1].strip()
  307. text = name.split(" ", 1)
  308. if len(text) > 1:
  309. name = text[0]
  310. data = text[1].strip()
  311. if replied_message and (replied_message.sticker or replied_message.video_note):
  312. data = None
  313. else:
  314. if replied_message and (replied_message.sticker or replied_message.video_note):
  315. data = None
  316. elif (
  317. replied_message and not replied_message.text and not replied_message.caption
  318. ):
  319. data = None
  320. else:
  321. data = (
  322. replied_message.text.markdown
  323. if replied_message.text
  324. else replied_message.caption.markdown
  325. )
  326. command = message.command[0]
  327. match = f"/{command} " + name
  328. if not message.reply_to_message and message.text:
  329. if match == data:
  330. data = "error"
  331. elif not message.reply_to_message and not message.text:
  332. if match == data:
  333. data = None
  334. return data, name
  335. async def get_specific_usernames(client, user_ids: list) -> Dict[int, str]:
  336. def _fetch_users():
  337. ids_str = ",".join(str(uid) for uid in user_ids)
  338. with client.storage.conn:
  339. query = f"""
  340. SELECT usernames.id, usernames.username
  341. FROM usernames
  342. WHERE usernames.id IN ({ids_str})
  343. AND username IS NOT NULL
  344. """
  345. result = client.storage.conn.execute(query).fetchall()
  346. users_ = {}
  347. for row in result:
  348. users_[row[0]] = row[1]
  349. return users_
  350. try:
  351. users = await asyncio.to_thread(_fetch_users)
  352. return users
  353. except Exception as e:
  354. print(f"Error fetching users: {e}")
  355. return {}