fixing a lot of custom emojis
This commit is contained in:
@@ -21,9 +21,12 @@ def list_installed_modules() -> list[tuple[str, str | None]]:
|
|||||||
if not os.path.isdir(base):
|
if not os.path.isdir(base):
|
||||||
return []
|
return []
|
||||||
items: list[tuple[str, str | None]] = []
|
items: list[tuple[str, str | None]] = []
|
||||||
for name in sorted(os.listdir(base)):
|
for raw_name in sorted(os.listdir(base)):
|
||||||
path = os.path.join(base, name)
|
path = os.path.join(base, raw_name)
|
||||||
if os.path.isdir(path) and not name.startswith("."):
|
if os.path.isdir(path) and not raw_name.startswith("."):
|
||||||
|
name = (raw_name or "").strip()
|
||||||
|
if not name:
|
||||||
|
continue
|
||||||
ver = None
|
ver = None
|
||||||
vp = os.path.join(path, "VERSION")
|
vp = os.path.join(path, "VERSION")
|
||||||
if os.path.isfile(vp):
|
if os.path.isfile(vp):
|
||||||
|
|||||||
+1
-1
@@ -13,7 +13,7 @@ asyncpg==0.30.0
|
|||||||
attrs==24.2.0
|
attrs==24.2.0
|
||||||
babel==2.17.0
|
babel==2.17.0
|
||||||
cachetools==5.5.1
|
cachetools==5.5.1
|
||||||
certifi==2023.11.17 # aiocryptopay требует certifi<2024; обновить после выхода совместимой версии aiocryptopay
|
certifi==2023.11.17
|
||||||
cffi==1.17.1
|
cffi==1.17.1
|
||||||
charset-normalizer==3.4.0
|
charset-normalizer==3.4.0
|
||||||
click==8.2.2
|
click==8.2.2
|
||||||
|
|||||||
+33
-500
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from typing import Any, Iterable
|
from typing import Any
|
||||||
|
|
||||||
from aiogram import Bot
|
from aiogram import Bot
|
||||||
from aiogram.enums import MessageEntityType
|
from aiogram.enums import MessageEntityType
|
||||||
@@ -17,6 +17,10 @@ _CODE_BLOCK_RE = re.compile(r"<code[^>]*>(.*?)</code>", re.IGNORECASE | re.DOTAL
|
|||||||
_PRE_BLOCK_RE = re.compile(r"<pre[^>]*>(.*?)</pre>", re.IGNORECASE | re.DOTALL)
|
_PRE_BLOCK_RE = re.compile(r"<pre[^>]*>(.*?)</pre>", re.IGNORECASE | re.DOTALL)
|
||||||
|
|
||||||
|
|
||||||
|
def _set_parse_mode_none(kwargs: dict[str, Any]) -> None:
|
||||||
|
kwargs["parse_mode"] = None
|
||||||
|
|
||||||
|
|
||||||
def _get_protected_ranges(text: str) -> list[tuple[int, int]]:
|
def _get_protected_ranges(text: str) -> list[tuple[int, int]]:
|
||||||
"""Return ranges inside <code> and <pre> tags to skip replacements."""
|
"""Return ranges inside <code> and <pre> tags to skip replacements."""
|
||||||
ranges: list[tuple[int, int]] = []
|
ranges: list[tuple[int, int]] = []
|
||||||
@@ -67,35 +71,48 @@ async def _fetch_placeholder(emoji_id: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
async def _replace_markers(text: str) -> tuple[str, list[MessageEntity]]:
|
async def _replace_markers(text: str) -> tuple[str, list[MessageEntity]]:
|
||||||
"""Replace markers with placeholders and build custom emoji entities."""
|
"""Replace markers with placeholders and build custom emoji entities.
|
||||||
|
"""
|
||||||
if not text:
|
if not text:
|
||||||
return text, []
|
return text, []
|
||||||
|
|
||||||
entities: list[MessageEntity] = []
|
entities: list[MessageEntity] = []
|
||||||
result = text
|
|
||||||
|
|
||||||
protected_ranges = _get_protected_ranges(text)
|
protected_ranges = _get_protected_ranges(text)
|
||||||
matches = [m for m in _MARKER_RE.finditer(text) if not _is_in_ranges(m.start(), protected_ranges)]
|
matches = [m for m in _MARKER_RE.finditer(text) if not _is_in_ranges(m.start(), protected_ranges)]
|
||||||
for match in reversed(matches):
|
if not matches:
|
||||||
|
return text, []
|
||||||
|
|
||||||
|
replacements: list[tuple[int, int, str, str]] = []
|
||||||
|
for match in matches:
|
||||||
emoji_id = match.group(1) or match.group(2)
|
emoji_id = match.group(1) or match.group(2)
|
||||||
marker = match.group(0)
|
start, end = match.start(), match.end()
|
||||||
marker_pos = match.start()
|
|
||||||
|
|
||||||
placeholder = await _fetch_placeholder(emoji_id)
|
placeholder = await _fetch_placeholder(emoji_id)
|
||||||
result = result[:marker_pos] + placeholder + result[marker_pos + len(marker) :]
|
replacements.append((start, end, str(emoji_id), placeholder))
|
||||||
|
|
||||||
offset = _utf16_len(result[:marker_pos])
|
parts: list[str] = []
|
||||||
length = _utf16_len(placeholder)
|
pos = 0
|
||||||
|
for start, end, _emoji_id, placeholder in replacements:
|
||||||
|
parts.append(text[pos:start])
|
||||||
|
parts.append(placeholder)
|
||||||
|
pos = end
|
||||||
|
parts.append(text[pos:])
|
||||||
|
result = "".join(parts)
|
||||||
|
|
||||||
entities.insert(
|
offset_utf16 = 0
|
||||||
0,
|
pos = 0
|
||||||
|
for start, end, emoji_id, placeholder in replacements:
|
||||||
|
offset_utf16 += _utf16_len(text[pos:start])
|
||||||
|
length_utf16 = _utf16_len(placeholder)
|
||||||
|
entities.append(
|
||||||
MessageEntity(
|
MessageEntity(
|
||||||
type=MessageEntityType.CUSTOM_EMOJI,
|
type=MessageEntityType.CUSTOM_EMOJI,
|
||||||
offset=offset,
|
offset=offset_utf16,
|
||||||
length=length,
|
length=length_utf16,
|
||||||
custom_emoji_id=str(emoji_id),
|
custom_emoji_id=emoji_id,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
offset_utf16 += length_utf16
|
||||||
|
pos = end
|
||||||
|
|
||||||
return result, entities
|
return result, entities
|
||||||
|
|
||||||
@@ -253,11 +270,6 @@ async def _process_text(
|
|||||||
return processed, merged
|
return processed, merged
|
||||||
|
|
||||||
|
|
||||||
def _set_parse_mode_none(kwargs: dict[str, Any]) -> None:
|
|
||||||
if kwargs.get("parse_mode") is not None:
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
|
|
||||||
|
|
||||||
def patch_bot_methods() -> bool:
|
def patch_bot_methods() -> bool:
|
||||||
"""Patch Message methods to auto-handle custom emojis."""
|
"""Patch Message methods to auto-handle custom emojis."""
|
||||||
global _BOT
|
global _BOT
|
||||||
@@ -390,482 +402,3 @@ def initialize_custom_emojis() -> bool:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"[CustomEmojis] Init failed: {e}", exc_info=True)
|
logger.error(f"[CustomEmojis] Init failed: {e}", exc_info=True)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
"""
|
|
||||||
Support custom Telegram emojis in bot texts.
|
|
||||||
Allows writing custom emoji IDs in texts via a special syntax.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import re
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from aiogram import Bot
|
|
||||||
from aiogram.enums import MessageEntityType
|
|
||||||
from aiogram.types import MessageEntity
|
|
||||||
|
|
||||||
from logger import logger
|
|
||||||
|
|
||||||
_emoji_placeholder_cache: dict[str, str] = {}
|
|
||||||
_bot_instance: Bot | None = None
|
|
||||||
|
|
||||||
|
|
||||||
async def _get_emoji_placeholder(emoji_id: str) -> str:
|
|
||||||
"""Fetch placeholder emoji for custom_emoji_id via API, with cache."""
|
|
||||||
if emoji_id in _emoji_placeholder_cache:
|
|
||||||
return _emoji_placeholder_cache[emoji_id]
|
|
||||||
|
|
||||||
if _bot_instance is None:
|
|
||||||
return "😀"
|
|
||||||
|
|
||||||
try:
|
|
||||||
stickers = await _bot_instance.get_custom_emoji_stickers(custom_emoji_ids=[emoji_id])
|
|
||||||
if stickers:
|
|
||||||
sticker = stickers[0]
|
|
||||||
placeholder = None
|
|
||||||
if hasattr(sticker, "emoji") and sticker.emoji:
|
|
||||||
placeholder = sticker.emoji
|
|
||||||
elif hasattr(sticker, "alt") and sticker.alt:
|
|
||||||
placeholder = sticker.alt
|
|
||||||
|
|
||||||
if placeholder:
|
|
||||||
_emoji_placeholder_cache[emoji_id] = placeholder
|
|
||||||
return placeholder
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
return "😀"
|
|
||||||
|
|
||||||
|
|
||||||
def _get_utf16_length(text: str) -> int:
|
|
||||||
"""Calculate length in UTF-16 code units."""
|
|
||||||
return len(text.encode("utf-16-le")) // 2
|
|
||||||
|
|
||||||
|
|
||||||
async def parse_custom_emoji_markers(text: str) -> tuple[str, list[MessageEntity]]:
|
|
||||||
"""Parse text with custom emoji markers to placeholders and entities."""
|
|
||||||
if not text:
|
|
||||||
return text, []
|
|
||||||
|
|
||||||
entities: list[MessageEntity] = []
|
|
||||||
result_text = text
|
|
||||||
pattern = r"\{emoji:(\d+)\}|\[emoji:(\d+)\]"
|
|
||||||
matches = list(re.finditer(pattern, text))
|
|
||||||
|
|
||||||
for match in reversed(matches):
|
|
||||||
emoji_id = match.group(1) or match.group(2)
|
|
||||||
marker_text = match.group(0)
|
|
||||||
marker_pos_in_result = match.start()
|
|
||||||
|
|
||||||
emoji_placeholder = await _get_emoji_placeholder(emoji_id)
|
|
||||||
result_text = (
|
|
||||||
result_text[:marker_pos_in_result]
|
|
||||||
+ emoji_placeholder
|
|
||||||
+ result_text[marker_pos_in_result + len(marker_text) :]
|
|
||||||
)
|
|
||||||
|
|
||||||
text_before = result_text[:marker_pos_in_result]
|
|
||||||
offset_utf16 = _get_utf16_length(text_before)
|
|
||||||
placeholder_utf16_length = _get_utf16_length(emoji_placeholder)
|
|
||||||
|
|
||||||
entity = MessageEntity(
|
|
||||||
type=MessageEntityType.CUSTOM_EMOJI,
|
|
||||||
offset=offset_utf16,
|
|
||||||
length=placeholder_utf16_length,
|
|
||||||
custom_emoji_id=str(emoji_id),
|
|
||||||
)
|
|
||||||
entities.insert(0, entity)
|
|
||||||
|
|
||||||
return result_text, entities
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_html_entities(text: str) -> list[MessageEntity]:
|
|
||||||
"""Parse HTML markup and build entities manually."""
|
|
||||||
entities: list[MessageEntity] = []
|
|
||||||
|
|
||||||
link_patterns = [
|
|
||||||
(r'<a href="([^"]+)"\s*>', r"</a>"),
|
|
||||||
(r"<a href='([^']+)'\s*>", r"</a>"),
|
|
||||||
]
|
|
||||||
|
|
||||||
link_tags = []
|
|
||||||
for open_pattern, close_pattern in link_patterns:
|
|
||||||
for open_match in re.finditer(open_pattern, text):
|
|
||||||
url = open_match.group(1)
|
|
||||||
link_tags.append({
|
|
||||||
"type": "open",
|
|
||||||
"url": url,
|
|
||||||
"open_start": open_match.start(),
|
|
||||||
"open_end": open_match.end(),
|
|
||||||
"close_pattern": close_pattern,
|
|
||||||
})
|
|
||||||
|
|
||||||
for close_match in re.finditer(r"</a>", text):
|
|
||||||
link_tags.append({"type": "close", "pos": close_match.start()})
|
|
||||||
|
|
||||||
link_tags.sort(key=lambda x: x.get("open_start", x.get("pos", 0)))
|
|
||||||
|
|
||||||
link_stacks = []
|
|
||||||
tag_positions = []
|
|
||||||
|
|
||||||
for tag in link_tags:
|
|
||||||
if tag["type"] == "open":
|
|
||||||
link_stacks.append(tag)
|
|
||||||
else:
|
|
||||||
if link_stacks:
|
|
||||||
matching_open = link_stacks.pop()
|
|
||||||
open_pos = matching_open["open_end"]
|
|
||||||
close_pos = tag["pos"]
|
|
||||||
content = text[open_pos:close_pos]
|
|
||||||
|
|
||||||
tag_positions.append({
|
|
||||||
"open_pos": open_pos,
|
|
||||||
"close_pos": close_pos,
|
|
||||||
"content": content,
|
|
||||||
"entity_type": "text_link",
|
|
||||||
"url": matching_open["url"],
|
|
||||||
})
|
|
||||||
|
|
||||||
other_patterns = [
|
|
||||||
(r"<b\s*>", r"</b>", "bold"),
|
|
||||||
(r"<strong\s*>", r"</strong>", "bold"),
|
|
||||||
(r"<i\s*>", r"</i>", "italic"),
|
|
||||||
(r"<em\s*>", r"</em>", "italic"),
|
|
||||||
(r"<u\s*>", r"</u>", "underline"),
|
|
||||||
(r"<ins\s*>", r"</ins>", "underline"),
|
|
||||||
(r"<s\s*>", r"</s>", "strikethrough"),
|
|
||||||
(r"<strike\s*>", r"</strike>", "strikethrough"),
|
|
||||||
(r"<del\s*>", r"</del>", "strikethrough"),
|
|
||||||
(r"<code\s*>", r"</code>", "code"),
|
|
||||||
(r"<pre\s*>", r"</pre>", "pre"),
|
|
||||||
(r"<blockquote\s*>", r"</blockquote>", "blockquote"),
|
|
||||||
]
|
|
||||||
|
|
||||||
other_tags = []
|
|
||||||
for open_pattern, close_pattern, entity_type in other_patterns:
|
|
||||||
for open_match in re.finditer(open_pattern, text):
|
|
||||||
other_tags.append({
|
|
||||||
"type": "open",
|
|
||||||
"entity_type": entity_type,
|
|
||||||
"open_end": open_match.end(),
|
|
||||||
"close_pattern": close_pattern,
|
|
||||||
})
|
|
||||||
|
|
||||||
for close_match in re.finditer(close_pattern, text):
|
|
||||||
other_tags.append({"type": "close", "entity_type": entity_type, "pos": close_match.start()})
|
|
||||||
|
|
||||||
other_tags.sort(key=lambda x: x.get("open_end", x.get("pos", 0)))
|
|
||||||
|
|
||||||
other_stacks: dict[str, list[dict[str, Any]]] = {}
|
|
||||||
|
|
||||||
for tag in other_tags:
|
|
||||||
entity_type = tag["entity_type"]
|
|
||||||
|
|
||||||
if tag["type"] == "open":
|
|
||||||
if entity_type not in other_stacks:
|
|
||||||
other_stacks[entity_type] = []
|
|
||||||
other_stacks[entity_type].append(tag)
|
|
||||||
else:
|
|
||||||
if entity_type in other_stacks and other_stacks[entity_type]:
|
|
||||||
matching_open = other_stacks[entity_type].pop()
|
|
||||||
open_pos = matching_open["open_end"]
|
|
||||||
close_pos = tag["pos"]
|
|
||||||
content = text[open_pos:close_pos]
|
|
||||||
|
|
||||||
tag_positions.append({
|
|
||||||
"open_pos": open_pos,
|
|
||||||
"close_pos": close_pos,
|
|
||||||
"content": content,
|
|
||||||
"entity_type": entity_type,
|
|
||||||
"url": None,
|
|
||||||
})
|
|
||||||
|
|
||||||
tag_positions.sort(key=lambda x: x["open_pos"])
|
|
||||||
|
|
||||||
for tag_info in tag_positions:
|
|
||||||
content = tag_info["content"]
|
|
||||||
open_pos = tag_info["open_pos"]
|
|
||||||
entity_type = tag_info["entity_type"]
|
|
||||||
url = tag_info["url"]
|
|
||||||
|
|
||||||
text_before_content = text[:open_pos]
|
|
||||||
offset_utf16 = _get_utf16_length(text_before_content)
|
|
||||||
|
|
||||||
length_utf16 = _get_utf16_length(content)
|
|
||||||
|
|
||||||
entity_dict = {
|
|
||||||
"type": MessageEntityType(entity_type),
|
|
||||||
"offset": offset_utf16,
|
|
||||||
"length": length_utf16,
|
|
||||||
}
|
|
||||||
if url:
|
|
||||||
entity_dict["url"] = url
|
|
||||||
|
|
||||||
entity = MessageEntity(**entity_dict)
|
|
||||||
entities.append(entity)
|
|
||||||
|
|
||||||
return entities
|
|
||||||
|
|
||||||
|
|
||||||
async def process_text_with_custom_emojis(text: str) -> tuple[str, list[MessageEntity] | None]:
|
|
||||||
"""Process text with custom emoji markers."""
|
|
||||||
if not text or not isinstance(text, str):
|
|
||||||
return text, None
|
|
||||||
|
|
||||||
processed_text, entities = await parse_custom_emoji_markers(text)
|
|
||||||
return (processed_text, entities) if entities else (text, None)
|
|
||||||
|
|
||||||
|
|
||||||
def patch_bot_methods() -> bool:
|
|
||||||
"""Patch Message methods to handle custom emojis."""
|
|
||||||
global _bot_instance
|
|
||||||
try:
|
|
||||||
from aiogram.types import Message
|
|
||||||
from bot import bot
|
|
||||||
|
|
||||||
_bot_instance = bot
|
|
||||||
|
|
||||||
if not hasattr(Message, "_original_answer"):
|
|
||||||
Message._original_answer = Message.answer
|
|
||||||
Message._original_edit_text = Message.edit_text
|
|
||||||
Message._original_edit_caption = Message.edit_caption
|
|
||||||
Message._original_answer_photo = Message.answer_photo
|
|
||||||
Message._original_answer_video = Message.answer_video
|
|
||||||
Message._original_answer_animation = Message.answer_animation
|
|
||||||
Message._original_edit_media = Message.edit_media
|
|
||||||
|
|
||||||
async def process_text_and_entities(
|
|
||||||
text: str,
|
|
||||||
entities: list[MessageEntity] | None = None,
|
|
||||||
parse_mode: str | None = None,
|
|
||||||
):
|
|
||||||
"""Process text and entities for custom emojis."""
|
|
||||||
if not text:
|
|
||||||
return text, entities
|
|
||||||
|
|
||||||
processed_text_with_html, custom_entities = await process_text_with_custom_emojis(text)
|
|
||||||
|
|
||||||
html_entities_with_html = []
|
|
||||||
if custom_entities and "<" in processed_text_with_html and ">" in processed_text_with_html:
|
|
||||||
html_entities_with_html = _parse_html_entities(processed_text_with_html)
|
|
||||||
|
|
||||||
if html_entities_with_html or custom_entities:
|
|
||||||
text_without_html = ""
|
|
||||||
utf16_offset_map: dict[int, int] = {}
|
|
||||||
html_pos = 0
|
|
||||||
plain_utf16 = 0
|
|
||||||
html_utf16 = 0
|
|
||||||
|
|
||||||
while html_pos < len(processed_text_with_html):
|
|
||||||
if processed_text_with_html[html_pos] == "<":
|
|
||||||
while html_pos < len(processed_text_with_html) and processed_text_with_html[html_pos] != ">":
|
|
||||||
char = processed_text_with_html[html_pos]
|
|
||||||
char_utf16_len = _get_utf16_length(char)
|
|
||||||
for i in range(char_utf16_len):
|
|
||||||
utf16_offset_map[html_utf16 + i] = plain_utf16
|
|
||||||
html_utf16 += char_utf16_len
|
|
||||||
html_pos += 1
|
|
||||||
if html_pos < len(processed_text_with_html):
|
|
||||||
char = processed_text_with_html[html_pos]
|
|
||||||
char_utf16_len = _get_utf16_length(char)
|
|
||||||
for i in range(char_utf16_len):
|
|
||||||
utf16_offset_map[html_utf16 + i] = plain_utf16
|
|
||||||
html_utf16 += char_utf16_len
|
|
||||||
html_pos += 1
|
|
||||||
else:
|
|
||||||
char = processed_text_with_html[html_pos]
|
|
||||||
text_without_html += char
|
|
||||||
char_utf16_len = _get_utf16_length(char)
|
|
||||||
for i in range(char_utf16_len):
|
|
||||||
utf16_offset_map[html_utf16 + i] = plain_utf16 + i
|
|
||||||
plain_utf16 += char_utf16_len
|
|
||||||
html_utf16 += char_utf16_len
|
|
||||||
html_pos += 1
|
|
||||||
|
|
||||||
def recalculate_offset(entity_offset_utf16: int) -> int:
|
|
||||||
"""Recalculate UTF-16 offset from HTML to plain text."""
|
|
||||||
if entity_offset_utf16 in utf16_offset_map:
|
|
||||||
return utf16_offset_map[entity_offset_utf16]
|
|
||||||
|
|
||||||
sorted_keys = sorted(utf16_offset_map.keys())
|
|
||||||
best_match = None
|
|
||||||
for key in sorted_keys:
|
|
||||||
if key <= entity_offset_utf16:
|
|
||||||
best_match = key
|
|
||||||
else:
|
|
||||||
break
|
|
||||||
|
|
||||||
return utf16_offset_map[best_match] if best_match is not None else entity_offset_utf16
|
|
||||||
|
|
||||||
html_entities = []
|
|
||||||
for entity in html_entities_with_html:
|
|
||||||
new_offset = recalculate_offset(entity.offset)
|
|
||||||
entity_end_offset = entity.offset + entity.length
|
|
||||||
new_end_offset = recalculate_offset(entity_end_offset)
|
|
||||||
new_length = new_end_offset - new_offset
|
|
||||||
entity_dict = entity.model_dump()
|
|
||||||
entity_dict["offset"] = new_offset
|
|
||||||
entity_dict["length"] = new_length
|
|
||||||
html_entities.append(MessageEntity(**entity_dict))
|
|
||||||
|
|
||||||
corrected_custom_entities = []
|
|
||||||
for entity in custom_entities:
|
|
||||||
new_offset = recalculate_offset(entity.offset)
|
|
||||||
entity_dict = entity.model_dump()
|
|
||||||
entity_dict["offset"] = new_offset
|
|
||||||
corrected_custom_entities.append(MessageEntity(**entity_dict))
|
|
||||||
|
|
||||||
final_entities: list[MessageEntity] = []
|
|
||||||
if html_entities:
|
|
||||||
final_entities.extend(html_entities)
|
|
||||||
if corrected_custom_entities:
|
|
||||||
final_entities.extend(corrected_custom_entities)
|
|
||||||
if entities:
|
|
||||||
final_entities.extend(entities)
|
|
||||||
|
|
||||||
if final_entities:
|
|
||||||
final_entities = sorted(final_entities, key=lambda e: e.offset)
|
|
||||||
|
|
||||||
return text_without_html, final_entities if final_entities else None
|
|
||||||
|
|
||||||
return text, entities
|
|
||||||
|
|
||||||
async def patched_message_answer(self, text: str, entities: list[MessageEntity] | None = None, **kwargs):
|
|
||||||
"""Patched Message.answer."""
|
|
||||||
processed_text, final_entities = await process_text_and_entities(text, entities, kwargs.get("parse_mode"))
|
|
||||||
if final_entities:
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
return await self._original_answer(
|
|
||||||
text=processed_text, entities=final_entities if final_entities else None, **kwargs
|
|
||||||
)
|
|
||||||
|
|
||||||
async def patched_message_edit_text(self, text: str, entities: list[MessageEntity] | None = None, **kwargs):
|
|
||||||
"""Patched Message.edit_text."""
|
|
||||||
processed_text, final_entities = await process_text_and_entities(text, entities, kwargs.get("parse_mode"))
|
|
||||||
if final_entities:
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
return await self._original_edit_text(text=processed_text, entities=final_entities, **kwargs)
|
|
||||||
|
|
||||||
async def patched_message_edit_caption(
|
|
||||||
self,
|
|
||||||
caption: str | None = None,
|
|
||||||
caption_entities: list[MessageEntity] | None = None,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
"""Patched Message.edit_caption."""
|
|
||||||
if not caption:
|
|
||||||
return await self._original_edit_caption(caption=caption, caption_entities=caption_entities, **kwargs)
|
|
||||||
processed_caption, final_entities = await process_text_and_entities(
|
|
||||||
caption, caption_entities, kwargs.get("parse_mode")
|
|
||||||
)
|
|
||||||
if final_entities:
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
return await self._original_edit_caption(
|
|
||||||
caption=processed_caption, caption_entities=final_entities, **kwargs
|
|
||||||
)
|
|
||||||
|
|
||||||
async def patched_message_answer_photo(
|
|
||||||
self,
|
|
||||||
photo: Any,
|
|
||||||
caption: str | None = None,
|
|
||||||
caption_entities: list[MessageEntity] | None = None,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
"""Patched Message.answer_photo."""
|
|
||||||
if not caption:
|
|
||||||
return await self._original_answer_photo(
|
|
||||||
photo=photo, caption=caption, caption_entities=caption_entities, **kwargs
|
|
||||||
)
|
|
||||||
processed_caption, final_entities = await process_text_and_entities(
|
|
||||||
caption, caption_entities, kwargs.get("parse_mode")
|
|
||||||
)
|
|
||||||
if final_entities:
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
return await self._original_answer_photo(
|
|
||||||
photo=photo, caption=processed_caption, caption_entities=final_entities, **kwargs
|
|
||||||
)
|
|
||||||
|
|
||||||
async def patched_message_answer_video(
|
|
||||||
self,
|
|
||||||
video: Any,
|
|
||||||
caption: str | None = None,
|
|
||||||
caption_entities: list[MessageEntity] | None = None,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
"""Patched Message.answer_video."""
|
|
||||||
if not caption:
|
|
||||||
return await self._original_answer_video(
|
|
||||||
video=video, caption=caption, caption_entities=caption_entities, **kwargs
|
|
||||||
)
|
|
||||||
processed_caption, final_entities = await process_text_and_entities(
|
|
||||||
caption, caption_entities, kwargs.get("parse_mode")
|
|
||||||
)
|
|
||||||
if final_entities:
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
return await self._original_answer_video(
|
|
||||||
video=video, caption=processed_caption, caption_entities=final_entities, **kwargs
|
|
||||||
)
|
|
||||||
|
|
||||||
async def patched_message_answer_animation(
|
|
||||||
self,
|
|
||||||
animation: Any,
|
|
||||||
caption: str | None = None,
|
|
||||||
caption_entities: list[MessageEntity] | None = None,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
"""Patched Message.answer_animation."""
|
|
||||||
if not caption:
|
|
||||||
return await self._original_answer_animation(
|
|
||||||
animation=animation,
|
|
||||||
caption=caption,
|
|
||||||
caption_entities=caption_entities,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
processed_caption, final_entities = await process_text_and_entities(
|
|
||||||
caption, caption_entities, kwargs.get("parse_mode")
|
|
||||||
)
|
|
||||||
if final_entities:
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
return await self._original_answer_animation(
|
|
||||||
animation=animation,
|
|
||||||
caption=processed_caption,
|
|
||||||
caption_entities=final_entities,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
async def patched_message_edit_media(self, media: Any, **kwargs):
|
|
||||||
"""Patched Message.edit_media."""
|
|
||||||
if hasattr(media, "caption") and media.caption:
|
|
||||||
processed_caption, final_entities = await process_text_and_entities(
|
|
||||||
media.caption, getattr(media, "caption_entities", None), kwargs.get("parse_mode")
|
|
||||||
)
|
|
||||||
media.caption = processed_caption
|
|
||||||
if final_entities:
|
|
||||||
if hasattr(media, "parse_mode"):
|
|
||||||
media.parse_mode = None
|
|
||||||
kwargs["parse_mode"] = None
|
|
||||||
media.caption_entities = final_entities
|
|
||||||
return await self._original_edit_media(media=media, **kwargs)
|
|
||||||
|
|
||||||
Message.answer = patched_message_answer
|
|
||||||
Message.edit_text = patched_message_edit_text
|
|
||||||
Message.edit_caption = patched_message_edit_caption
|
|
||||||
Message.answer_photo = patched_message_answer_photo
|
|
||||||
Message.answer_video = patched_message_answer_video
|
|
||||||
Message.answer_animation = patched_message_answer_animation
|
|
||||||
Message.edit_media = patched_message_edit_media
|
|
||||||
|
|
||||||
return True
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"[CustomEmojis] Error while patching bot methods: {e}", exc_info=True)
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def initialize_custom_emojis() -> bool:
|
|
||||||
"""Initialize custom emoji support."""
|
|
||||||
try:
|
|
||||||
return patch_bot_methods()
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"[CustomEmojis] Error during initialization: {e}", exc_info=True)
|
|
||||||
return False
|
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ def load_modules_from_folder(folder: str = "modules") -> list[Router]:
|
|||||||
base_path = Path(folder)
|
base_path = Path(folder)
|
||||||
|
|
||||||
for _finder, name, _ispkg in pkgutil.iter_modules([str(base_path)]):
|
for _finder, name, _ispkg in pkgutil.iter_modules([str(base_path)]):
|
||||||
|
name = (name or "").strip()
|
||||||
if not _is_safe_module_name(name):
|
if not _is_safe_module_name(name):
|
||||||
logger.warning(f"[Modules] Пропуск недопустимого имени модуля: {name!r}")
|
logger.warning(f"[Modules] Пропуск недопустимого имени модуля: {name!r}")
|
||||||
continue
|
continue
|
||||||
@@ -50,6 +51,7 @@ def load_module_webhooks(folder: str = "modules") -> list[dict]:
|
|||||||
base_path = Path(folder)
|
base_path = Path(folder)
|
||||||
|
|
||||||
for _finder, name, _ispkg in pkgutil.iter_modules([str(base_path)]):
|
for _finder, name, _ispkg in pkgutil.iter_modules([str(base_path)]):
|
||||||
|
name = (name or "").strip()
|
||||||
if not _is_safe_module_name(name):
|
if not _is_safe_module_name(name):
|
||||||
continue
|
continue
|
||||||
if not manager.should_autostart(name):
|
if not manager.should_autostart(name):
|
||||||
@@ -74,6 +76,7 @@ def load_module_fast_flow_handlers(folder: str = "modules") -> dict:
|
|||||||
base_path = Path(folder)
|
base_path = Path(folder)
|
||||||
|
|
||||||
for _finder, name, _ispkg in pkgutil.iter_modules([str(base_path)]):
|
for _finder, name, _ispkg in pkgutil.iter_modules([str(base_path)]):
|
||||||
|
name = (name or "").strip()
|
||||||
if not _is_safe_module_name(name):
|
if not _is_safe_module_name(name):
|
||||||
continue
|
continue
|
||||||
if not manager.should_autostart(name):
|
if not manager.should_autostart(name):
|
||||||
|
|||||||
@@ -13,6 +13,10 @@ IGNORE_SUBMODULES = {"models", "schemas", "db"}
|
|||||||
STATE_FILE = os.getenv("MODULES_STATE_FILE", "storage/modules_state.json")
|
STATE_FILE = os.getenv("MODULES_STATE_FILE", "storage/modules_state.json")
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_module_name(name: str | None) -> str:
|
||||||
|
return (name or "").strip()
|
||||||
|
|
||||||
|
|
||||||
class ModuleRecord:
|
class ModuleRecord:
|
||||||
def __init__(self, name: str, pkg: str) -> None:
|
def __init__(self, name: str, pkg: str) -> None:
|
||||||
self.name = name
|
self.name = name
|
||||||
@@ -36,7 +40,8 @@ class ModulesManager:
|
|||||||
if os.path.isfile(STATE_FILE):
|
if os.path.isfile(STATE_FILE):
|
||||||
with open(STATE_FILE, encoding="utf-8") as f:
|
with open(STATE_FILE, encoding="utf-8") as f:
|
||||||
data = json.load(f)
|
data = json.load(f)
|
||||||
self.disabled = set(data.get("disabled", []))
|
raw = data.get("disabled", [])
|
||||||
|
self.disabled = {_normalize_module_name(n) for n in raw if _normalize_module_name(n)}
|
||||||
else:
|
else:
|
||||||
os.makedirs(os.path.dirname(STATE_FILE), exist_ok=True)
|
os.makedirs(os.path.dirname(STATE_FILE), exist_ok=True)
|
||||||
self._save_state()
|
self._save_state()
|
||||||
@@ -52,6 +57,7 @@ class ModulesManager:
|
|||||||
logger.warning(f"[Modules] Не удалось сохранить состояние: {e}")
|
logger.warning(f"[Modules] Не удалось сохранить состояние: {e}")
|
||||||
|
|
||||||
def adopt(self, name: str, router: Router):
|
def adopt(self, name: str, router: Router):
|
||||||
|
name = _normalize_module_name(name)
|
||||||
rec = self.registry.get(name) or ModuleRecord(name, self.pkg(name))
|
rec = self.registry.get(name) or ModuleRecord(name, self.pkg(name))
|
||||||
rec.router = router
|
rec.router = router
|
||||||
rec.enabled = True
|
rec.enabled = True
|
||||||
@@ -61,6 +67,7 @@ class ModulesManager:
|
|||||||
return bool(name and name.isidentifier() and "." not in name and "/" not in name and "\\" not in name)
|
return bool(name and name.isidentifier() and "." not in name and "/" not in name and "\\" not in name)
|
||||||
|
|
||||||
async def start(self, name: str) -> None:
|
async def start(self, name: str) -> None:
|
||||||
|
name = _normalize_module_name(name)
|
||||||
if not self._is_safe_module_name(name):
|
if not self._is_safe_module_name(name):
|
||||||
raise ValueError(f"[Modules] Недопустимое имя модуля: {name!r}")
|
raise ValueError(f"[Modules] Недопустимое имя модуля: {name!r}")
|
||||||
rec = self.registry.get(name) or ModuleRecord(name, self.pkg(name))
|
rec = self.registry.get(name) or ModuleRecord(name, self.pkg(name))
|
||||||
@@ -95,6 +102,7 @@ class ModulesManager:
|
|||||||
logger.info(f"[Modules] {name} запущен.")
|
logger.info(f"[Modules] {name} запущен.")
|
||||||
|
|
||||||
async def stop(self, name: str) -> None:
|
async def stop(self, name: str) -> None:
|
||||||
|
name = _normalize_module_name(name)
|
||||||
rec = self.registry.get(name)
|
rec = self.registry.get(name)
|
||||||
if not rec or not rec.enabled:
|
if not rec or not rec.enabled:
|
||||||
logger.info(f"[Modules] {name} уже остановлен или не найден.")
|
logger.info(f"[Modules] {name} уже остановлен или не найден.")
|
||||||
@@ -124,6 +132,7 @@ class ModulesManager:
|
|||||||
logger.info(f"[Modules] {name} остановлен.")
|
logger.info(f"[Modules] {name} остановлен.")
|
||||||
|
|
||||||
async def restart(self, name: str) -> None:
|
async def restart(self, name: str) -> None:
|
||||||
|
name = _normalize_module_name(name)
|
||||||
logger.info(f"[Modules] Перезапуск {name}...")
|
logger.info(f"[Modules] Перезапуск {name}...")
|
||||||
await self.stop(name)
|
await self.stop(name)
|
||||||
await self.start(name)
|
await self.start(name)
|
||||||
@@ -142,6 +151,7 @@ class ModulesManager:
|
|||||||
importlib.invalidate_caches()
|
importlib.invalidate_caches()
|
||||||
|
|
||||||
def is_enabled(self, name: str) -> bool:
|
def is_enabled(self, name: str) -> bool:
|
||||||
|
name = _normalize_module_name(name)
|
||||||
rec = self.registry.get(name)
|
rec = self.registry.get(name)
|
||||||
if not rec or not rec.router:
|
if not rec or not rec.router:
|
||||||
return False
|
return False
|
||||||
@@ -153,10 +163,10 @@ class ModulesManager:
|
|||||||
return bool(sub and rec.router in sub)
|
return bool(sub and rec.router in sub)
|
||||||
|
|
||||||
def is_disabled(self, name: str) -> bool:
|
def is_disabled(self, name: str) -> bool:
|
||||||
return name in self.disabled
|
return _normalize_module_name(name) in self.disabled
|
||||||
|
|
||||||
def should_autostart(self, name: str) -> bool:
|
def should_autostart(self, name: str) -> bool:
|
||||||
return name not in self.disabled
|
return _normalize_module_name(name) not in self.disabled
|
||||||
|
|
||||||
|
|
||||||
manager = ModulesManager()
|
manager = ModulesManager()
|
||||||
|
|||||||
Reference in New Issue
Block a user