Compare commits
22 Commits
75388d8693
...
0.2
| Author | SHA1 | Date | |
|---|---|---|---|
| b28b7c583d | |||
|
|
f8e9a528e2 | ||
|
|
9266d8f32b | ||
|
|
d5e30f8c36 | ||
|
|
a803b80d59 | ||
|
|
8094aaab10 | ||
|
|
28ef360d5f | ||
|
|
af7c9c7467 | ||
|
|
ddc1c36c72 | ||
|
|
ff6656fdb0 | ||
| 3b4e9d6a1e | |||
| 0ae0f8ca53 | |||
| 422534ea45 | |||
| e7460c0a9e | |||
| 4c1507555a | |||
| ce9c157b72 | |||
| cee1275e3c | |||
| 9b702df01f | |||
| 4a12bdbd40 | |||
| a2a1b6c2ed | |||
| c519ae379d | |||
| 8c1399d781 |
9
.dockerignore
Normal file
9
.dockerignore
Normal file
@@ -0,0 +1,9 @@
|
|||||||
|
.venv/
|
||||||
|
session_storage/
|
||||||
|
__pycache__/
|
||||||
|
runtime/
|
||||||
|
*.swp
|
||||||
|
*.swo
|
||||||
|
*.vscode
|
||||||
|
*.json
|
||||||
|
*.sqlite
|
||||||
4
.gitignore
vendored
4
.gitignore
vendored
@@ -1,7 +1,9 @@
|
|||||||
.venv/
|
.venv/
|
||||||
session_storage/
|
session_storage/
|
||||||
__pycache__/
|
__pycache__/
|
||||||
|
runtime/
|
||||||
*.swp
|
*.swp
|
||||||
*.swo
|
*.swo
|
||||||
*.vscode
|
*.vscode
|
||||||
*.json
|
*.json
|
||||||
|
*.sqlite
|
||||||
16
Dockerfile
Normal file
16
Dockerfile
Normal file
@@ -0,0 +1,16 @@
|
|||||||
|
FROM python:3.12-slim
|
||||||
|
|
||||||
|
RUN apt-get update && apt-get install --no-install-recommends -y libolm3 libolm-dev git && rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
|
WORKDIR /usr/src/app
|
||||||
|
|
||||||
|
COPY requirements.txt ./
|
||||||
|
RUN pip install --no-cache-dir -r requirements.txt
|
||||||
|
|
||||||
|
COPY . .
|
||||||
|
|
||||||
|
WORKDIR /runtime
|
||||||
|
RUN chown -R 1000:1000 /usr/src /runtime
|
||||||
|
|
||||||
|
USER 1000:1000
|
||||||
|
CMD [ "python", "/usr/src/app/main.py" ]
|
||||||
112
README.md
112
README.md
@@ -4,7 +4,30 @@
|
|||||||
Раньше был бот с таким же функционалом, но для Telegram. Telegram больше не в
|
Раньше был бот с таким же функционалом, но для Telegram. Telegram больше не в
|
||||||
почёте, и теперь у меня всё в локальном Matrix, поэтому бот тоже перенесён сюда.
|
почёте, и теперь у меня всё в локальном Matrix, поэтому бот тоже перенесён сюда.
|
||||||
|
|
||||||
## Подготовка окружения
|
## Для запуска рекомендуется `docker`
|
||||||
|
|
||||||
|
1. [Скачайте образ](https://git.tyukalov.su/nikita/-/packages/container/2026-matrix-csonac) (измените версию на нужную вам):
|
||||||
|
```bash
|
||||||
|
sudo docker pull git.tyukalov.su/nikita/2026-matrix-csonac:0.2
|
||||||
|
```
|
||||||
|
2. Выполните первый запуск в интерактивном режиме
|
||||||
|
```bash
|
||||||
|
sudo mkdir runtime
|
||||||
|
sudo chown 1000:1000 runtime
|
||||||
|
sudo docker run -v ./runtime:/runtime -ti git.tyukalov.su/nikita/2026-matrix-csonac:0.2
|
||||||
|
```
|
||||||
|
3. В папке `runtime` появится файл `config.json`. Внесите туда нужные настройки
|
||||||
|
4. Запустите ещё раз в интерактивном режиме, чтобы авторизоваться
|
||||||
|
```bash
|
||||||
|
sudo docker run -v ./runtime:/runtime -ti git.tyukalov.su/nikita/2026-matrix-csonac:0.2
|
||||||
|
```
|
||||||
|
5. Дальше можно запускать контейнер в фоне
|
||||||
|
```bash
|
||||||
|
sudo docker run -p 127.0.0.1:4980:4980 -v ./runtime:/runtime --detach git.tyukalov.su/nikita/2026-matrix-csonac:0.2
|
||||||
|
```
|
||||||
|
> Рекомендуется запускать контейнер при помощи `docker compose`.
|
||||||
|
|
||||||
|
## Запуск без `docker`
|
||||||
|
|
||||||
1. Клонируйте репозиторий и перейдите в его директорию
|
1. Клонируйте репозиторий и перейдите в его директорию
|
||||||
```bash
|
```bash
|
||||||
@@ -26,18 +49,25 @@ python main.py
|
|||||||
python main.py
|
python main.py
|
||||||
```
|
```
|
||||||
|
|
||||||
## Запуск
|
## Как работает бот
|
||||||
|
|
||||||
Для запуска приложения вы можете либо активировать `venv`, либо просто использовать полный путь до интерпретатора Python:
|
Для каждой службы (канала уведомлений) создаётся своя собственная комната и в
|
||||||
```bash
|
такие комнаты добавляется бот. Затем вся работа с ботом производится при помощи
|
||||||
~/2026-matrix-downloader/.venv/bin/python main.py
|
команд - текстовых сообщений, начинающихся с *восклицательного знака*.
|
||||||
```
|
Поддерживаются следующие команды:
|
||||||
```bash
|
- `!help` - получить справку
|
||||||
. .venv/bin/activate
|
- `!info` - получить сведения о комнате
|
||||||
python main.py
|
- `!tokens` - получить список токенов
|
||||||
```
|
- `!auth [SERVICE_NAME]` - создать новый токен (если указать имя сервиса, то
|
||||||
|
приложение, использующее токен, не сможет самостоятельно указывать имя
|
||||||
|
сервиса - всегда будет использовано имя, указанное вами)
|
||||||
|
- `!deauth <TOKEN>` - удалить токен
|
||||||
|
- `!name <TOKEN> [SERVICE_NAME]` - указать (или удалить) имя сервиса для токена
|
||||||
|
- `!ban <IP> <SECONDS> [REASON]` - забанить указанный IP на указанное число секунд (можно указать причину)
|
||||||
|
- `!unban <IP>` - разбанить указанный IP адрес
|
||||||
|
- `!bans` - получить список забаненных IP адресов
|
||||||
|
|
||||||
## Как работает этот бот
|
## Как работает веб-сервер
|
||||||
|
|
||||||
Этот бот запускает веб-сервер, используя `web_host` и `web_port` из файла
|
Этот бот запускает веб-сервер, используя `web_host` и `web_port` из файла
|
||||||
конфигурации.
|
конфигурации.
|
||||||
@@ -46,35 +76,45 @@ python main.py
|
|||||||
> `nginx`. Это также позволит вам использовать защищённое соединение, что
|
> `nginx`. Это также позволит вам использовать защищённое соединение, что
|
||||||
> исключит возможность применения атаки Man-in-the-Middle для перехвата токена.
|
> исключит возможность применения атаки Man-in-the-Middle для перехвата токена.
|
||||||
|
|
||||||
Все запросы к веб-серверу требуют авторизации, используя HTTP заголовок
|
Все запросы к веб-серверу являются GET-запросами и требуют авторизации.
|
||||||
`Authorization` и схему `Bearer`. Токен авторизации генерируется посредством
|
Авторизоваться можно двумя путями:
|
||||||
взаимодействия с ботом в Matrix.
|
1. Использовать HTTP-заголовок `Authorization` и схему `Bearer`. Например:
|
||||||
|
```plain
|
||||||
|
Authorization: Bearer 1234567890abcdef
|
||||||
|
```
|
||||||
|
2. Использовать URL параметр `token`. Например:
|
||||||
|
```plain
|
||||||
|
https://csonac.su/api/notify?token=f1829e94d...
|
||||||
|
```
|
||||||
|
Рекомендуется использовать первый способ (HTTP-заголовок), так как это позволяет
|
||||||
|
избежать раскрытия токена в логах сервера и других местах, где можно посмотреть
|
||||||
|
URL прошлых запросов.
|
||||||
|
|
||||||
> Для каждого сервиса, использующего службу, рекомендуется генерировать свой
|
> Для каждого сервиса, использующего бота, рекомендуется генерировать свой
|
||||||
> собственный токен. Это позволит отозвать токен только для одной службы, если
|
> собственный токен. Это позволит отозвать токен только для одного серсива, если
|
||||||
> он будет украден.
|
> токен будет украден.
|
||||||
|
|
||||||
Доступные эндпоинты:
|
Доступные эндпоинты:
|
||||||
- `POST /<channel>/notify`
|
- `GET /api/notify`
|
||||||
- **Описание.** Используется, чтобы отправить уведомление в указанный канал.
|
- **Описание.** Используется, чтобы отправить уведомление в канал.
|
||||||
Вместо `<channel>` указывается код канала, получаемый при помощи команды
|
- **Параметры запроса**
|
||||||
`!channel_code`, выполненной в комнате Matrix.
|
- `channel` - код канала, получаемый при помощи команды `!info`,
|
||||||
- **Тело запроса.** Тело запроса представляет собой `json` объект:
|
выполненной в комнате Matrix
|
||||||
```json
|
- `service` - имя сервиса (учитывается, только если для токена не было
|
||||||
|
настроено имя сервиса через бота)
|
||||||
|
- `text` - текст уведомления (форматирование не поддерживается)
|
||||||
|
- **Ответ**
|
||||||
|
- В случае успеха сервер вернёт `200` и JSON следующего формата:
|
||||||
|
```json
|
||||||
{
|
{
|
||||||
"service": "<название сервиса, передаётся только если разрешено>",
|
"notification_id": "Notification ID will be here"
|
||||||
"subservice": "<название подсервиса>",
|
|
||||||
"text": "<текст уведомления>",
|
|
||||||
"urgent": false
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
- **Тело ответа.** Тело ответа представляет собой `json` объект. В случае
|
> В текущей версии `notification_id` не имеет практической пользы и его
|
||||||
успеха в нём будут все поля, перечисляемые ниже. В случае провала - только
|
> формат будет меняться.
|
||||||
поле `error`, содержащее текстовое описание ошибки.
|
- В случае ошибки сервер вернёт JSON следующего формата:
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"error": null,
|
"detail": "Error description in English"
|
||||||
"notification_id": "<здесь будет Notification ID>"
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
- ``
|
|
||||||
202
bot.py
202
bot.py
@@ -1,202 +0,0 @@
|
|||||||
"""
|
|
||||||
matrix-nio basics wrapper
|
|
||||||
"""
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import time
|
|
||||||
import traceback
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from nio import LoginError, LoginResponse, SyncResponse
|
|
||||||
from nio import WhoamiResponse
|
|
||||||
from nio import AsyncClient, AsyncClientConfig
|
|
||||||
|
|
||||||
from datatypes import AppConfig
|
|
||||||
import util
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# DATA
|
|
||||||
#
|
|
||||||
_app_config: AppConfig
|
|
||||||
_client: AsyncClient
|
|
||||||
_task: asyncio.Task | None = None
|
|
||||||
_stop: asyncio.Event | None = None
|
|
||||||
_since: str | None = None
|
|
||||||
_last_since_save_time: float = 0
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# CALLBACKS
|
|
||||||
#
|
|
||||||
async def _sync_callback(response: SyncResponse):
|
|
||||||
global _since, _last_since_save_time
|
|
||||||
t = time.time()
|
|
||||||
_since = response.next_batch
|
|
||||||
if t - _last_since_save_time >= 120.0:
|
|
||||||
await util.set_next_batch(_app_config, _since)
|
|
||||||
_last_since_save_time = t
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# PRIVATE
|
|
||||||
#
|
|
||||||
async def _bot_login_using_token(token: str, device_id: str) -> bool:
|
|
||||||
"""Tries to login using access_token. Returns True on success."""
|
|
||||||
util.log_info("Authorizing using access_token...")
|
|
||||||
_client.restore_login(
|
|
||||||
user_id=f"@{_app_config.matrix_user}:{util.get_hostname_from_url(_app_config.matrix_homeserver)}",
|
|
||||||
device_id=device_id,
|
|
||||||
access_token=token
|
|
||||||
)
|
|
||||||
result = await _client.whoami()
|
|
||||||
if type(result) is not WhoamiResponse:
|
|
||||||
return False
|
|
||||||
util.log_info(f"Logged in as {result.user_id} using access_token")
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def _bot_login_using_password(password: str) -> tuple[str, str] | tuple[None, None]:
|
|
||||||
"""Tries to login using password. Returns (access_token, device_id) on success."""
|
|
||||||
util.log_info("Authorizing using password...")
|
|
||||||
result = await _client.login(password=password)
|
|
||||||
if type(result) is LoginResponse:
|
|
||||||
util.log_info(f"Authorized using password")
|
|
||||||
return result.access_token, result.device_id
|
|
||||||
elif type(result) is LoginError:
|
|
||||||
util.log_error(f"Failed to authorize using password: {result.message}")
|
|
||||||
return None, None
|
|
||||||
else:
|
|
||||||
raise RuntimeError(f"Invalid login result: {result}")
|
|
||||||
|
|
||||||
async def _bot_login(config: AppConfig) -> bool:
|
|
||||||
"""Tries to login"""
|
|
||||||
try:
|
|
||||||
# get the session token and try to use it
|
|
||||||
session_token, device_id = await util.get_session_data(config)
|
|
||||||
if session_token is not None and device_id is not None:
|
|
||||||
if await _bot_login_using_token(session_token, device_id):
|
|
||||||
return True
|
|
||||||
await util.set_session_data(config, None)
|
|
||||||
util.log_warning("Existing access_token is deleted")
|
|
||||||
# get the password and try to use it
|
|
||||||
password = await util.get_password()
|
|
||||||
if password is None:
|
|
||||||
util.log_error("No password provided (consider using MATRIX_PASSWORD environment variable)")
|
|
||||||
return False
|
|
||||||
access_token, device_id = await _bot_login_using_password(password)
|
|
||||||
if access_token is not None and device_id is not None:
|
|
||||||
await util.set_session_data(config, (access_token, device_id))
|
|
||||||
util.log_warning("Saved new access_token and device_id")
|
|
||||||
return True
|
|
||||||
# can't login
|
|
||||||
return False
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
return False
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def _bot_loop(config: AppConfig) -> None:
|
|
||||||
"""Bot loop"""
|
|
||||||
global _client, _since
|
|
||||||
# app stop task
|
|
||||||
if _stop is None:
|
|
||||||
raise RuntimeError("_stop can't be None")
|
|
||||||
stop_task = asyncio.create_task(_stop.wait())
|
|
||||||
# setup the callback for syncing
|
|
||||||
_client.add_response_callback(_sync_callback, SyncResponse) # type: ignore
|
|
||||||
# login
|
|
||||||
login_task = asyncio.create_task(_bot_login(config))
|
|
||||||
done, _ = await asyncio.wait(
|
|
||||||
[login_task, stop_task],
|
|
||||||
return_when=asyncio.FIRST_COMPLETED
|
|
||||||
)
|
|
||||||
# stopped
|
|
||||||
if stop_task in done:
|
|
||||||
login_task.cancel()
|
|
||||||
return
|
|
||||||
# failed to login
|
|
||||||
if login_task.exception() or not login_task.result():
|
|
||||||
util.request_app_stop("Can't authorize into matrix")
|
|
||||||
return
|
|
||||||
# load initial `next_batch`
|
|
||||||
_since = await util.get_next_batch(config)
|
|
||||||
# sync forever
|
|
||||||
while True:
|
|
||||||
# sync
|
|
||||||
sync_task = asyncio.create_task(_client.sync_forever(timeout=5000, since=_since))
|
|
||||||
done, _ = await asyncio.wait(
|
|
||||||
[sync_task, stop_task],
|
|
||||||
return_when=asyncio.FIRST_COMPLETED
|
|
||||||
)
|
|
||||||
# stopped
|
|
||||||
if stop_task in done:
|
|
||||||
util.log_info("Stopping sync_forever...")
|
|
||||||
_client.stop_sync_forever()
|
|
||||||
util.log_info("Waiting for sync_forever to quit...")
|
|
||||||
await sync_task
|
|
||||||
sync_task.cancel()
|
|
||||||
break
|
|
||||||
# something happened
|
|
||||||
try:
|
|
||||||
sync_task.result()
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
await asyncio.sleep(1)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# PUBLIC
|
|
||||||
#
|
|
||||||
async def start(config: AppConfig) -> bool:
|
|
||||||
"""Starts the bot"""
|
|
||||||
global _client
|
|
||||||
global _task, _stop, _app_config
|
|
||||||
if _task is not None:
|
|
||||||
return False
|
|
||||||
_app_config = config
|
|
||||||
# create the bot
|
|
||||||
store_dir = Path.cwd() / config.store_dir
|
|
||||||
store_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
client_config = AsyncClientConfig(
|
|
||||||
store_name="storefile",
|
|
||||||
encryption_enabled=True,
|
|
||||||
store_sync_tokens=False
|
|
||||||
)
|
|
||||||
_client = AsyncClient(
|
|
||||||
homeserver=config.matrix_homeserver,
|
|
||||||
user=config.matrix_user,
|
|
||||||
store_path=str(store_dir),
|
|
||||||
config=client_config
|
|
||||||
)
|
|
||||||
_stop = asyncio.Event()
|
|
||||||
_task = asyncio.create_task(_bot_loop(config))
|
|
||||||
return True
|
|
||||||
|
|
||||||
def get_client() -> AsyncClient:
|
|
||||||
return _client
|
|
||||||
|
|
||||||
async def stop() -> None:
|
|
||||||
"""Stop the bot"""
|
|
||||||
global _task, _stop
|
|
||||||
if _task is None or _stop is None:
|
|
||||||
return
|
|
||||||
_stop.set()
|
|
||||||
try:
|
|
||||||
await _task
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
pass
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
_task = None
|
|
||||||
_stop = None
|
|
||||||
try:
|
|
||||||
await _client.close()
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
await util.set_next_batch(_app_config, _since)
|
|
||||||
|
|
||||||
251
bot_callbacks.py
Normal file
251
bot_callbacks.py
Normal file
@@ -0,0 +1,251 @@
|
|||||||
|
"""Callbacks for matrix messages"""
|
||||||
|
|
||||||
|
import html
|
||||||
|
import time
|
||||||
|
import traceback
|
||||||
|
|
||||||
|
from mab import MatrixBot
|
||||||
|
from nio import MatrixRoom, RoomMessageText
|
||||||
|
|
||||||
|
import util
|
||||||
|
from database import Database
|
||||||
|
from datatypes import *
|
||||||
|
|
||||||
|
#
|
||||||
|
# PRIVATE
|
||||||
|
#
|
||||||
|
_bot: MatrixBot = None # type: ignore
|
||||||
|
_db: Database = None # type: ignore
|
||||||
|
|
||||||
|
def _generate_help_message() -> str:
|
||||||
|
"""Generates help message in HTML markup"""
|
||||||
|
result = "<strong><i>Как использовать</i></strong><br><ol>"
|
||||||
|
result += "<li>Используя <code>!auth</code>, создайте токен</li>"
|
||||||
|
result += "<li>Добавьте бота в комнату для уведомлений</li>"
|
||||||
|
result += "<li>Выполните <code>!info</code> в комнате, чтобы узнать код канала уведомений</li>"
|
||||||
|
result += "<li>Используя полученные токен и код канала, отправьте уведомление через веб-запрос</li>"
|
||||||
|
result += "</ol>"
|
||||||
|
|
||||||
|
result += "<br><br><strong><i>Команды</i></strong>"
|
||||||
|
for aliases in _COMMANDS:
|
||||||
|
result += f"<br><code>!{aliases[0]}</code> - <i>{html.escape(_COMMANDS[aliases][1])}</i>"
|
||||||
|
return result
|
||||||
|
|
||||||
|
#
|
||||||
|
# GENERIC EVENT HANDLERS
|
||||||
|
#
|
||||||
|
async def _on_text(room: MatrixRoom, event: RoomMessageText) -> None:
|
||||||
|
"""This callback is called when text message is received"""
|
||||||
|
if event.sender == _bot.get_client().user_id:
|
||||||
|
return
|
||||||
|
text = event.body
|
||||||
|
parts = [p for p in text.split() if p.strip()]
|
||||||
|
if len(parts) < 1:
|
||||||
|
return
|
||||||
|
cmd = parts[0].lower()
|
||||||
|
if not cmd.startswith("!"):
|
||||||
|
return
|
||||||
|
cb = None
|
||||||
|
for aliases in _COMMANDS:
|
||||||
|
for alias in aliases:
|
||||||
|
if alias == cmd.lstrip("!"):
|
||||||
|
cb = _COMMANDS[aliases][0]
|
||||||
|
break
|
||||||
|
if cb:
|
||||||
|
break
|
||||||
|
if cb is None:
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>Используйте <code>!help</code></strong>")
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
await cb(room, parts[1:])
|
||||||
|
except:
|
||||||
|
traceback.print_exc()
|
||||||
|
|
||||||
|
#
|
||||||
|
# COMMAND HANDLERS
|
||||||
|
#
|
||||||
|
async def _on_cmd_help(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!help"""
|
||||||
|
await _bot.send_text_to_room(room.room_id, _generate_help_message(), is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_info(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!info"""
|
||||||
|
# get room info (or add it)
|
||||||
|
room_info = await _db.room_get(room.room_id)
|
||||||
|
if room_info is None:
|
||||||
|
room_info = await _db.room_create(room.room_id)
|
||||||
|
# failure
|
||||||
|
if room_info is None:
|
||||||
|
await _bot.send_text_to_room(
|
||||||
|
room.room_id,
|
||||||
|
"<strong>Нет информации о комнате</strong>",
|
||||||
|
is_html=True
|
||||||
|
)
|
||||||
|
return
|
||||||
|
# respond
|
||||||
|
response = "<strong><i>Сведения о комнате</i></strong><br>"
|
||||||
|
response += f"<strong>Канал:</strong> <code>{html.escape(room_info.code)}</code>"
|
||||||
|
await _bot.send_text_to_room(room.room_id, response, is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_tokens(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!info"""
|
||||||
|
# get all tokens and check if there are none
|
||||||
|
tokens = await _db.token_get_all()
|
||||||
|
if not tokens:
|
||||||
|
await _bot.send_text_to_room(
|
||||||
|
room.room_id,
|
||||||
|
"<strong>Нет токенов, используйте <code>!auth</code></strong>",
|
||||||
|
is_html=True
|
||||||
|
)
|
||||||
|
return
|
||||||
|
# create the list of tokens
|
||||||
|
response = "<strong><i>Список токенов</i></strong>"
|
||||||
|
for token in tokens:
|
||||||
|
response += f"<br><strong><code>{html.escape(token.code)}</code></strong><br><ul>"
|
||||||
|
if token.name is not None:
|
||||||
|
response += f"<li><strong>Имя службы:</strong> <code>{token.name}</code></li>"
|
||||||
|
else:
|
||||||
|
response += f"<li><strong>Имя службы:</strong> <i>указывается в запросе</i></li>"
|
||||||
|
response += f"<li><strong>Создан:</strong> <i>{util.date_to_text(token.created_at)}</i></li>"
|
||||||
|
response += f"<li><strong>Последнее использование:</strong> <i>{util.date_to_text(token.last_access_at)}</i></li>"
|
||||||
|
response += "</ul>"
|
||||||
|
# respond
|
||||||
|
await _bot.send_text_to_room(room.room_id, response, is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_auth(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!auth"""
|
||||||
|
# create new token
|
||||||
|
token = await _db.token_create()
|
||||||
|
# set the name if it is provided
|
||||||
|
if args:
|
||||||
|
token.name = " ".join(args)
|
||||||
|
await _db.token_set_name(token.code, token.name)
|
||||||
|
# create the response
|
||||||
|
response = "<strong>Создан новый токен</strong>"
|
||||||
|
response += f"<br><strong>Код:</strong> <code>{token.code}</code>"
|
||||||
|
response += f"<br><strong>Имя сервиса:</strong> "
|
||||||
|
if token.name:
|
||||||
|
response += f"<code>{html.escape(token.name)}</code>"
|
||||||
|
else:
|
||||||
|
response += "указывается в запросе"
|
||||||
|
# respond
|
||||||
|
await _bot.send_text_to_room(room.room_id, response, is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_deauth(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!deauth"""
|
||||||
|
# no token provided
|
||||||
|
if len(args) != 1:
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>Укажите токен, который надо удалить (должен быть ровно один аргумент)</strong>", is_html=True)
|
||||||
|
return
|
||||||
|
token = args[0]
|
||||||
|
# check if token does not exist
|
||||||
|
if await _db.token_get(token) is None:
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>Токен не найден</strong>", is_html=True)
|
||||||
|
return
|
||||||
|
# remove the token
|
||||||
|
await _db.token_delete(token)
|
||||||
|
# respond
|
||||||
|
response = f"<strong>Удалён токен <code>{token}</code></strong>"
|
||||||
|
await _bot.send_text_to_room(room.room_id, response, is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_name(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!name"""
|
||||||
|
# check arguments
|
||||||
|
if len(args) < 1:
|
||||||
|
error = "<strong>Требуется указать как минимум токен. Если хотите убрать имя, то имя указывать не надо. Если имя нужно назначить или изменить, то после токена укажите новое имя.</strong>"
|
||||||
|
await _bot.send_text_to_room(room.room_id, error, is_html=True)
|
||||||
|
return
|
||||||
|
token = args[0]
|
||||||
|
new_name = " ".join(args[1:])
|
||||||
|
if not new_name.strip():
|
||||||
|
new_name = None
|
||||||
|
# check if token exists
|
||||||
|
if await _db.token_get(token) is None:
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>Токен не существует</strong>", is_html=True)
|
||||||
|
return
|
||||||
|
# set new name
|
||||||
|
await _db.token_set_name(token, new_name)
|
||||||
|
# prepare the response
|
||||||
|
if new_name:
|
||||||
|
response = f"<strong>Новое имя для токена <code>{token}</code>: <code>{new_name}</code></strong>"
|
||||||
|
else:
|
||||||
|
response = f"<strong>Удалено имя для токена <code>{token}</code></strong>"
|
||||||
|
# respond
|
||||||
|
await _bot.send_text_to_room(room.room_id, response, is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_ban(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!ban"""
|
||||||
|
# check arguments
|
||||||
|
if len(args) < 2:
|
||||||
|
error = "<strong>Формат: <code>!ban <IP> <SECONDS> [REASON]</code></strong>"
|
||||||
|
await _bot.send_text_to_room(room.room_id, error, is_html=True)
|
||||||
|
return
|
||||||
|
# get args
|
||||||
|
try:
|
||||||
|
ip = args[0]
|
||||||
|
duration = float(args[1])
|
||||||
|
reason = " ".join(args[2:]) if args[2:] else "Manual ban"
|
||||||
|
except:
|
||||||
|
error = "<strong>Возникла ошибка. Наверняка неправильно указаны секунды.</strong>"
|
||||||
|
await _bot.send_text_to_room(room.room_id, error, is_html=True)
|
||||||
|
return
|
||||||
|
# ban
|
||||||
|
try:
|
||||||
|
await _db.ban_create(ip, time.time() + duration, reason)
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>IP адрес заблокирован</strong>", is_html=True)
|
||||||
|
except:
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>Возникла ошибка</strong>", is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_unban(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!unban"""
|
||||||
|
# check arguments
|
||||||
|
if len(args) != 1:
|
||||||
|
error = "<strong>Формат: <code>!ban <IP></code></strong>"
|
||||||
|
await _bot.send_text_to_room(room.room_id, error, is_html=True)
|
||||||
|
return
|
||||||
|
# unban
|
||||||
|
try:
|
||||||
|
await _db.ban_delete(args[0])
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>IP адрес разблокирован (если он был заблокирован)</strong>", is_html=True)
|
||||||
|
except:
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>Возникла ошибка</strong>", is_html=True)
|
||||||
|
|
||||||
|
async def _on_cmd_bans(room: MatrixRoom, args: list[str]) -> None:
|
||||||
|
"""!bans"""
|
||||||
|
# list
|
||||||
|
bans = _db.ban_get_all()
|
||||||
|
if not bans:
|
||||||
|
await _bot.send_text_to_room(room.room_id, "<strong>Нет заблокированных IP адресов</strong>", is_html=True)
|
||||||
|
return
|
||||||
|
# create the response
|
||||||
|
result = "<strong>Список заблокированных IP</strong><br><ul>"
|
||||||
|
for ban in bans:
|
||||||
|
result += f"<li><strong>{html.escape(ban["ip"])}</strong><ul>"
|
||||||
|
result += f"<li><strong>Действует до:</strong> <code>{util.date_to_text(ban["expires_at"])}</code></li>"
|
||||||
|
result += f"<li><strong>Причина:</strong> <code>{html.escape(ban["reason"])}</code></li>"
|
||||||
|
result += "</ul></li>"
|
||||||
|
result += "</ul>"
|
||||||
|
await _bot.send_text_to_room(room.room_id, result, is_html=True)
|
||||||
|
|
||||||
|
|
||||||
|
_COMMANDS = {
|
||||||
|
("help", "h", "?"): (_on_cmd_help, "Получить справку"),
|
||||||
|
("info", "room", "i", "r"): (_on_cmd_info, "Получить сведения о комнате"),
|
||||||
|
("tokens", "t"): (_on_cmd_tokens, "Получить список токенов"),
|
||||||
|
("auth", "create", "a", "c"): (_on_cmd_auth, "Создать новый токен"),
|
||||||
|
("deauth", "delete", "d"): (_on_cmd_deauth, "Удалить существующий токен"),
|
||||||
|
("name", "n"): (_on_cmd_name, "Задать (или удалить) имя для токена"),
|
||||||
|
("ban", "b"): (_on_cmd_ban, "Забанить IP адрес"),
|
||||||
|
("unban", "u"): (_on_cmd_unban, "Разбанить IP адрес"),
|
||||||
|
("bans", "l"): (_on_cmd_bans, "Получить список забаненных IP адресов"),
|
||||||
|
}
|
||||||
|
|
||||||
|
#
|
||||||
|
# PUBLIC
|
||||||
|
#
|
||||||
|
def setup(bot: MatrixBot, db: Database) -> None:
|
||||||
|
"""Setup the callbacks"""
|
||||||
|
global _bot, _db
|
||||||
|
_db = db
|
||||||
|
_bot = bot
|
||||||
|
_bot.add_event_callback(_on_text, RoomMessageText) # type: ignore
|
||||||
17
bot_types.py
17
bot_types.py
@@ -1,17 +0,0 @@
|
|||||||
"""Data types required for the bot"""
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from dataclasses import dataclass
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class MatrixBotConfig:
|
|
||||||
"""Configuration for MatrixBot"""
|
|
||||||
|
|
||||||
matrix_homeserver_url: str
|
|
||||||
"""Homeserver, for example: `https://matrix.domain.su`"""
|
|
||||||
|
|
||||||
matrix_username_localpart: str
|
|
||||||
"""Localpart of matrix username (without homeserver), for example: `valid-username`"""
|
|
||||||
|
|
||||||
storage_directory: Path
|
|
||||||
"""Path to the storage directory (will be created if needed)"""
|
|
||||||
@@ -11,7 +11,11 @@ import util
|
|||||||
DEFAULT_CONFIG = {
|
DEFAULT_CONFIG = {
|
||||||
"matrix_homeserver": "https://matrix.domain.net",
|
"matrix_homeserver": "https://matrix.domain.net",
|
||||||
"matrix_user": "short_username",
|
"matrix_user": "short_username",
|
||||||
"store_dir": "session_storage"
|
"store_dir": "session_storage",
|
||||||
|
"web_ip": "0.0.0.0",
|
||||||
|
"web_port": 4980,
|
||||||
|
"fails_to_ban": 50,
|
||||||
|
"ban_duration": 600
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -50,4 +54,5 @@ async def load_config(path: str = "config.json") -> AppConfig | None:
|
|||||||
cfg = AppConfig(**j)
|
cfg = AppConfig(**j)
|
||||||
return cfg
|
return cfg
|
||||||
except:
|
except:
|
||||||
|
traceback.print_exc()
|
||||||
return None
|
return None
|
||||||
387
database.py
387
database.py
@@ -0,0 +1,387 @@
|
|||||||
|
"""This module implements database operations."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import traceback
|
||||||
|
import time
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
import aiosqlite
|
||||||
|
from aiosqlite import Connection, Row
|
||||||
|
|
||||||
|
from datatypes import *
|
||||||
|
from util import get_hash
|
||||||
|
|
||||||
|
|
||||||
|
class Database:
|
||||||
|
#
|
||||||
|
# PRIVATE
|
||||||
|
#
|
||||||
|
BACKGROUND_ROUTINE_PERIOD = 60
|
||||||
|
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _setup_tables(conn: Connection) -> None:
|
||||||
|
SETUP_SQL_SCRIPT = """
|
||||||
|
CREATE TABLE IF NOT EXISTS rooms (
|
||||||
|
code TEXT PRIMARY KEY,
|
||||||
|
matrix_id TEXT NOT NULL UNIQUE
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS tokens (
|
||||||
|
code TEXT PRIMARY KEY,
|
||||||
|
name TEXT,
|
||||||
|
created_at REAL,
|
||||||
|
last_access_at REAL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS bans (
|
||||||
|
ip TEXT PRIMARY KEY,
|
||||||
|
expires_at INTEGER,
|
||||||
|
reason TEXT
|
||||||
|
);
|
||||||
|
"""
|
||||||
|
await conn.executescript(SETUP_SQL_SCRIPT)
|
||||||
|
await conn.commit()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _select_by_anded_kwargs(conn: Connection, table_name: str, **kwargs) -> list[dict[str, Any]]:
|
||||||
|
# prepare keys and values
|
||||||
|
keys = tuple(k for k in kwargs)
|
||||||
|
values = tuple(kwargs[k] for k in keys)
|
||||||
|
# prepare the statement
|
||||||
|
statement = f"SELECT * FROM {table_name}"
|
||||||
|
# add `WHERE` part if there kwargs
|
||||||
|
if keys:
|
||||||
|
statement += " WHERE "
|
||||||
|
statement += " AND ".join([f"{k}=?" for k in keys])
|
||||||
|
else:
|
||||||
|
values = None
|
||||||
|
# execute and enumerate
|
||||||
|
result = []
|
||||||
|
async with conn.execute(statement, values) as cursor:
|
||||||
|
async for row in cursor:
|
||||||
|
entry = {}
|
||||||
|
for k in row.keys():
|
||||||
|
entry[k] = row[k]
|
||||||
|
result.append(entry)
|
||||||
|
return result
|
||||||
|
|
||||||
|
async def _background_routine(self, stop_event: asyncio.Event) -> None:
|
||||||
|
"""This routine perform routine tasks."""
|
||||||
|
stop_task = asyncio.create_task(stop_event.wait())
|
||||||
|
wait_task = asyncio.create_task(asyncio.sleep(0))
|
||||||
|
self._logger.debug("Started background worker")
|
||||||
|
while True:
|
||||||
|
done, _ = await asyncio.wait(
|
||||||
|
[stop_task, wait_task],
|
||||||
|
return_when=asyncio.FIRST_COMPLETED
|
||||||
|
)
|
||||||
|
if stop_task in done:
|
||||||
|
wait_task.cancel()
|
||||||
|
break
|
||||||
|
wait_task = asyncio.create_task(
|
||||||
|
asyncio.sleep(self.BACKGROUND_ROUTINE_PERIOD)
|
||||||
|
)
|
||||||
|
if self._connection is None:
|
||||||
|
continue
|
||||||
|
# get current time
|
||||||
|
current_time = time.time()
|
||||||
|
# delete old bans
|
||||||
|
try:
|
||||||
|
statement = "DELETE FROM bans WHERE expires_at <= ? RETURNING ip"
|
||||||
|
async with self._connection.execute(statement, (current_time,)) as cursor:
|
||||||
|
async for row in cursor:
|
||||||
|
ip = row["ip"]
|
||||||
|
if ip in self._bans:
|
||||||
|
del self._bans[ip]
|
||||||
|
self._logger.debug(f"IP {ip} is not banned anymore")
|
||||||
|
await self._connection.commit()
|
||||||
|
self._logger.debug("Performed banned IPs cleanup")
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
# delete expired fails
|
||||||
|
try:
|
||||||
|
for ip in dict(self._fails):
|
||||||
|
self._fails[ip] = list(
|
||||||
|
filter(lambda x: x > current_time, self._fails[ip])
|
||||||
|
)
|
||||||
|
if not self._fails[ip]:
|
||||||
|
del self._fails[ip]
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
self._logger.debug("Stopped background worker")
|
||||||
|
|
||||||
|
async def _load_bans_from_database(self) -> None:
|
||||||
|
"""Loads bans information from database. Must be called when database connection is established."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
data = await self._select_by_anded_kwargs(self._connection, "bans")
|
||||||
|
for d in data:
|
||||||
|
self._bans[d["ip"]] = {
|
||||||
|
"expires_at": d["expires_at"],
|
||||||
|
"reason": d["reason"]
|
||||||
|
}
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
return
|
||||||
|
|
||||||
|
def _get_fails_for_ip(self, ip) -> int:
|
||||||
|
if ip not in self._fails:
|
||||||
|
return 0
|
||||||
|
t = time.time()
|
||||||
|
return sum([(1 if expires_at > t else 0) for expires_at in self._fails[ip]])
|
||||||
|
|
||||||
|
#
|
||||||
|
# PUBLIC
|
||||||
|
#
|
||||||
|
def __init__(self, path: Path):
|
||||||
|
self._path = path
|
||||||
|
self._logger = logging.getLogger("database")
|
||||||
|
self._logger.setLevel(logging.DEBUG)
|
||||||
|
|
||||||
|
self._connection: Connection | None = None
|
||||||
|
self._background_stop_event: asyncio.Event | None = None
|
||||||
|
self._background_task: asyncio.Task | None = None
|
||||||
|
self._bans: dict = {}
|
||||||
|
self._fails: dict[str, list[float]] = {}
|
||||||
|
|
||||||
|
async def connect(self) -> bool:
|
||||||
|
"""Connect to the database. Returns False on failure."""
|
||||||
|
if self._connection is not None:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
self._connection = await aiosqlite.connect(self._path)
|
||||||
|
self._connection.row_factory = Row
|
||||||
|
await self._setup_tables(self._connection)
|
||||||
|
self._bans = {}
|
||||||
|
self._fails = {}
|
||||||
|
await self._load_bans_from_database()
|
||||||
|
self._background_stop_event = asyncio.Event()
|
||||||
|
self._background_task = asyncio.create_task(
|
||||||
|
self._background_routine(self._background_stop_event)
|
||||||
|
)
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def disconnect(self) -> None:
|
||||||
|
"""Disconnect from the database."""
|
||||||
|
if self._connection is None:
|
||||||
|
return
|
||||||
|
self._background_stop_event.set() # type: ignore
|
||||||
|
await self._background_task # type: ignore
|
||||||
|
|
||||||
|
await self._connection.close()
|
||||||
|
self._connection = None
|
||||||
|
self._background_task = None
|
||||||
|
self._background_stop_event = None
|
||||||
|
|
||||||
|
|
||||||
|
async def room_create(self, matrix_id: str) -> ObjectRoom | None:
|
||||||
|
"""Create a room with specified matrix_id. Returns None if the room exists."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
# room exists
|
||||||
|
if await self.room_get(matrix_id) is not None:
|
||||||
|
return None
|
||||||
|
# create the code
|
||||||
|
code = None
|
||||||
|
string_to_hash = matrix_id
|
||||||
|
while code is None or (await self.room_get(code)) is not None:
|
||||||
|
code = get_hash(string_to_hash)[-10:]
|
||||||
|
string_to_hash += "A"
|
||||||
|
# code must always start with q
|
||||||
|
code = f"q{code}"
|
||||||
|
# add
|
||||||
|
try:
|
||||||
|
statement = "INSERT INTO rooms (code, matrix_id) VALUES (?, ?)"
|
||||||
|
await self._connection.execute(statement, (code, matrix_id))
|
||||||
|
await self._connection.commit()
|
||||||
|
self._logger.info(f"Added new room with code {code}")
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
return None
|
||||||
|
# return the added object
|
||||||
|
return ObjectRoom(
|
||||||
|
code=code,
|
||||||
|
matrix_id=matrix_id
|
||||||
|
)
|
||||||
|
|
||||||
|
async def room_get_all(self) -> list[ObjectRoom]:
|
||||||
|
"""Get room by code."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
rooms = await self._select_by_anded_kwargs(
|
||||||
|
self._connection,
|
||||||
|
"rooms"
|
||||||
|
)
|
||||||
|
result = []
|
||||||
|
for r in rooms:
|
||||||
|
result.append(ObjectRoom(
|
||||||
|
code=r["code"],
|
||||||
|
matrix_id=r["matrix_id"]
|
||||||
|
))
|
||||||
|
return result
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
return []
|
||||||
|
|
||||||
|
async def room_get(self, identifier: str) -> ObjectRoom | None:
|
||||||
|
"""Get room by identifier (either `code` or `matrix_id`)."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
kwargs = {}
|
||||||
|
kwargs["code" if identifier[0] == "q" else "matrix_id"] = identifier
|
||||||
|
rooms = await self._select_by_anded_kwargs(
|
||||||
|
self._connection,
|
||||||
|
"rooms",
|
||||||
|
**kwargs
|
||||||
|
)
|
||||||
|
if not rooms:
|
||||||
|
self._logger.debug(f"Room `{identifier}` is not found")
|
||||||
|
return None
|
||||||
|
return ObjectRoom(**rooms[0])
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def token_create(self) -> ObjectToken:
|
||||||
|
"""Create a token."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
code = get_hash(str(time.time()).encode() + os.urandom(64))
|
||||||
|
create_time = time.time()
|
||||||
|
statement = """
|
||||||
|
INSERT INTO tokens (code, name, created_at, last_access_at)
|
||||||
|
VALUES (?, ?, ?, ?)
|
||||||
|
"""
|
||||||
|
result = ObjectToken(
|
||||||
|
code=code,
|
||||||
|
name=None,
|
||||||
|
created_at=create_time,
|
||||||
|
last_access_at=create_time
|
||||||
|
)
|
||||||
|
await self._connection.execute(
|
||||||
|
statement,
|
||||||
|
(result.code, result.name, result.created_at, result.last_access_at)
|
||||||
|
)
|
||||||
|
await self._connection.commit()
|
||||||
|
return result
|
||||||
|
|
||||||
|
async def token_delete(self, code: str) -> None:
|
||||||
|
"""Delete a token."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
await self._connection.execute("DELETE FROM tokens WHERE code=?", (code,))
|
||||||
|
await self._connection.commit()
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
|
||||||
|
async def token_get(self, code: str) -> ObjectToken | None:
|
||||||
|
"""Get token by its code."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
result = await self._select_by_anded_kwargs(
|
||||||
|
self._connection,
|
||||||
|
"tokens",
|
||||||
|
code=code
|
||||||
|
)
|
||||||
|
if not result:
|
||||||
|
return None
|
||||||
|
result = result[0]
|
||||||
|
return ObjectToken(**result)
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def token_get_all(self) -> list[ObjectToken]:
|
||||||
|
"""Get list of all tokens."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
result = await self._select_by_anded_kwargs(
|
||||||
|
self._connection,
|
||||||
|
"tokens"
|
||||||
|
)
|
||||||
|
result = [ObjectToken(**r) for r in result]
|
||||||
|
return result
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
return []
|
||||||
|
|
||||||
|
async def token_set_name(self, code: str, name: str | None) -> None:
|
||||||
|
"""Set name for the token."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
statement = "UPDATE tokens SET name=? WHERE code=?"
|
||||||
|
await self._connection.execute(statement, (name, code))
|
||||||
|
await self._connection.commit()
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
|
||||||
|
|
||||||
|
async def ban_create(self, ip: str, expires_at: float, reason: str) -> None:
|
||||||
|
"""Save information about banned IP address. Replaces existing IPs."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
statement = "INSERT OR REPLACE INTO bans (ip, expires_at, reason) VALUES (?, ?, ?)"
|
||||||
|
await self._connection.execute(statement, (ip, expires_at, reason))
|
||||||
|
await self._connection.commit()
|
||||||
|
self._bans[ip] = {
|
||||||
|
"expires_at": expires_at,
|
||||||
|
"reason": reason
|
||||||
|
}
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
|
||||||
|
def ban_get_all(self) -> list[dict]:
|
||||||
|
"""Get information about banned IP addresses."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
return [{"ip": k, **self._bans[k]} for k in self._bans]
|
||||||
|
|
||||||
|
def ban_check(self, ip: str) -> bool:
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
return ip in self._bans
|
||||||
|
|
||||||
|
async def ban_delete(self, ip: str) -> None:
|
||||||
|
"""Unbans specified IP address."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
try:
|
||||||
|
statement = "DELETE FROM bans WHERE ip = ?"
|
||||||
|
await self._connection.execute(statement, (ip,))
|
||||||
|
await self._connection.commit()
|
||||||
|
if ip in self._bans:
|
||||||
|
del self._bans[ip]
|
||||||
|
except:
|
||||||
|
self._logger.error(traceback.format_exc())
|
||||||
|
|
||||||
|
|
||||||
|
def fail_create(self, ip: str, expires_at: float) -> int:
|
||||||
|
"""Adds failed access attempt. Returns count of failed attempts for IP."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
if ip not in self._fails:
|
||||||
|
self._fails[ip] = []
|
||||||
|
self._fails[ip].append(expires_at)
|
||||||
|
return self._get_fails_for_ip(ip)
|
||||||
|
|
||||||
|
def fail_clear(self, ip: str) -> None:
|
||||||
|
"""Clears failed attempts counter."""
|
||||||
|
if self._connection is None:
|
||||||
|
raise RuntimeError("Not connected to the database")
|
||||||
|
if ip in self._fails:
|
||||||
|
del self._fails[ip]
|
||||||
17
datatypes.py
17
datatypes.py
@@ -10,6 +10,19 @@ class AppConfig:
|
|||||||
matrix_homeserver: str
|
matrix_homeserver: str
|
||||||
matrix_user: str
|
matrix_user: str
|
||||||
store_dir: str
|
store_dir: str
|
||||||
|
web_ip: str
|
||||||
|
web_port: int
|
||||||
|
fails_to_ban: int
|
||||||
|
ban_duration: float
|
||||||
|
|
||||||
class MessageType(Enum):
|
@dataclass
|
||||||
TEXT = "m.text"
|
class ObjectRoom:
|
||||||
|
code: str
|
||||||
|
matrix_id: str
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ObjectToken:
|
||||||
|
code: str
|
||||||
|
name: str | None
|
||||||
|
created_at: float
|
||||||
|
last_access_at: float
|
||||||
148
logic.py
148
logic.py
@@ -1,148 +0,0 @@
|
|||||||
"""Main logic implementation"""
|
|
||||||
|
|
||||||
import traceback
|
|
||||||
from typing import Any
|
|
||||||
import html
|
|
||||||
|
|
||||||
from nio import AsyncClient
|
|
||||||
from nio import JoinResponse, RoomSendResponse, RoomSendError
|
|
||||||
from nio import MatrixInvitedRoom, InviteMemberEvent
|
|
||||||
from nio import MatrixRoom, RoomMessageText
|
|
||||||
|
|
||||||
from nio import OlmUnverifiedDeviceError
|
|
||||||
|
|
||||||
from datatypes import *
|
|
||||||
import util
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# DATA
|
|
||||||
#
|
|
||||||
_client: AsyncClient
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# PRIVATE
|
|
||||||
#
|
|
||||||
async def _verify_all_devices() -> None:
|
|
||||||
"""Verifies all known devices"""
|
|
||||||
for user_id in _client.device_store.users:
|
|
||||||
for device_id, olm_device in _client.device_store[user_id].items():
|
|
||||||
# can't trust ourselves
|
|
||||||
if device_id == _client.device_id and user_id == _client.user_id:
|
|
||||||
continue
|
|
||||||
# they are already verified
|
|
||||||
if olm_device.verified:
|
|
||||||
continue
|
|
||||||
# verify them
|
|
||||||
_client.verify_device(olm_device)
|
|
||||||
|
|
||||||
def _handle_html_in_kwargs(kwargs: dict[str, Any]) -> None:
|
|
||||||
"""Modified `kwargs` in-place so that `formatted_body` appears if needed"""
|
|
||||||
if "formatted_body" in kwargs or "body" not in kwargs:
|
|
||||||
return
|
|
||||||
is_html, text_without_html = util.check_and_remove_html(kwargs["body"])
|
|
||||||
if not is_html:
|
|
||||||
return
|
|
||||||
kwargs["format"] = "org.matrix.custom.html"
|
|
||||||
kwargs["formatted_body"] = kwargs["body"]
|
|
||||||
kwargs["body"] = text_without_html
|
|
||||||
|
|
||||||
async def _send_message_to(room_id: str, message_type: MessageType, **kwargs) -> str:
|
|
||||||
"""Sends a message to the room and returns event_id.
|
|
||||||
|
|
||||||
This function automatically detects `body` key in `kwargs` and checks
|
|
||||||
if it is a valid HTML. If it is a valid HTML, it will send it as such.
|
|
||||||
Moreover, `body` attribute will be cleaned from any HTML tags, so that
|
|
||||||
the text will be looking well. `formatted_body` attribute is added
|
|
||||||
automatically and you should not add it manually.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
# handle HTML
|
|
||||||
_handle_html_in_kwargs(kwargs)
|
|
||||||
# try to send the message
|
|
||||||
result = await _client.room_send(
|
|
||||||
room_id=room_id,
|
|
||||||
message_type="m.room.message",
|
|
||||||
content={
|
|
||||||
"msgtype": message_type.value,
|
|
||||||
**kwargs
|
|
||||||
}
|
|
||||||
)
|
|
||||||
# success
|
|
||||||
if type(result) is RoomSendResponse:
|
|
||||||
return result.event_id
|
|
||||||
# error
|
|
||||||
elif type(result) is RoomSendError:
|
|
||||||
raise Exception(result)
|
|
||||||
# unknown error
|
|
||||||
else:
|
|
||||||
raise RuntimeError()
|
|
||||||
except OlmUnverifiedDeviceError as e:
|
|
||||||
# verify everyone and retry
|
|
||||||
await _verify_all_devices()
|
|
||||||
return await _send_message_to(room_id, message_type, **kwargs)
|
|
||||||
except:
|
|
||||||
raise
|
|
||||||
|
|
||||||
async def _send_text_to(room_id: str, text: str) -> str:
|
|
||||||
"""Sends a text message to the room. `text` may be HTML"""
|
|
||||||
return await _send_message_to(
|
|
||||||
room_id=room_id,
|
|
||||||
message_type=MessageType.TEXT,
|
|
||||||
body=text
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# CALLBACKS
|
|
||||||
#
|
|
||||||
async def _message_callback(room: MatrixRoom, event: RoomMessageText) -> None:
|
|
||||||
"""Handle commands received from Matrix"""
|
|
||||||
try:
|
|
||||||
# do not process messages sent by ourselves
|
|
||||||
if event.sender == _client.user_id:
|
|
||||||
return
|
|
||||||
# prepare response
|
|
||||||
response = "<b>Получено сообщение</b>"
|
|
||||||
response += f"<br><br><b>Room ID:</b> {html.escape(room.room_id)}"
|
|
||||||
response += f"<br><b>Sender:</b> {html.escape(event.sender)}"
|
|
||||||
print(await _send_text_to(room.room_id, response))
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
async def _invite_callback(room: MatrixInvitedRoom, event: InviteMemberEvent) -> None:
|
|
||||||
"""Happens when the bot is invited to somewhere"""
|
|
||||||
try:
|
|
||||||
result = await _client.join(room.room_id)
|
|
||||||
if type(result) is JoinResponse:
|
|
||||||
util.log_info(f"Joined the room {room.room_id}")
|
|
||||||
else:
|
|
||||||
util.log_error(f"Can't join room {room.room_id}")
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
async def _generic_test_callback(*args, **kwargs) -> None:
|
|
||||||
"""Use this callback to check argument types"""
|
|
||||||
print("GENERIC TEST CALLBACK")
|
|
||||||
for a in args:
|
|
||||||
print(f" - {type(a)}")
|
|
||||||
for k in kwargs:
|
|
||||||
print(f" * {k} = {kwargs[k]}")
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# PUBLIC
|
|
||||||
#
|
|
||||||
async def setup(client: AsyncClient) -> None:
|
|
||||||
global _client
|
|
||||||
_client = client
|
|
||||||
client.add_event_callback(_message_callback, RoomMessageText) # type: ignore
|
|
||||||
client.add_event_callback(_invite_callback, InviteMemberEvent) # type: ignore
|
|
||||||
|
|
||||||
async def stop() -> None:
|
|
||||||
"""Stop all ongoing processes"""
|
|
||||||
pass
|
|
||||||
35
main.py
35
main.py
@@ -7,18 +7,20 @@ import traceback
|
|||||||
import signal
|
import signal
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import config
|
from mab import MatrixBot, MatrixBotConfig
|
||||||
import util
|
|
||||||
#import bot
|
|
||||||
#import logic
|
|
||||||
|
|
||||||
from new_bot import MatrixBot
|
|
||||||
from bot_types import MatrixBotConfig
|
|
||||||
|
|
||||||
from datatypes import AppConfig
|
from datatypes import AppConfig
|
||||||
|
import database
|
||||||
|
import config
|
||||||
|
import util
|
||||||
|
import bot_callbacks
|
||||||
|
import web
|
||||||
|
|
||||||
|
bot: MatrixBot = None # type: ignore
|
||||||
|
|
||||||
async def main() -> None:
|
async def main() -> None:
|
||||||
"""Entry point"""
|
"""Entry point"""
|
||||||
|
global bot
|
||||||
# setup signal handler
|
# setup signal handler
|
||||||
util.setup_app_stop_event()
|
util.setup_app_stop_event()
|
||||||
loop = asyncio.get_running_loop()
|
loop = asyncio.get_running_loop()
|
||||||
@@ -35,17 +37,34 @@ async def main() -> None:
|
|||||||
util.log_error("Could't load config")
|
util.log_error("Could't load config")
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# setup the bot
|
||||||
matrix_bot_config = MatrixBotConfig(
|
matrix_bot_config = MatrixBotConfig(
|
||||||
matrix_homeserver_url=cfg.matrix_homeserver,
|
matrix_homeserver_url=cfg.matrix_homeserver,
|
||||||
matrix_username_localpart=cfg.matrix_user,
|
matrix_username_localpart=cfg.matrix_user,
|
||||||
storage_directory=Path(cfg.store_dir)
|
storage_directory=Path(cfg.store_dir)
|
||||||
)
|
)
|
||||||
bot = MatrixBot(matrix_bot_config)
|
bot = MatrixBot(matrix_bot_config)
|
||||||
bot.start()
|
# setup the database
|
||||||
|
db = database.Database(Path("database.sqlite"))
|
||||||
|
# setup the callbacks
|
||||||
|
bot_callbacks.setup(bot, db)
|
||||||
|
# setup the web server
|
||||||
|
web_server = web.Web(cfg, bot, db)
|
||||||
|
|
||||||
|
# start the app
|
||||||
|
if not await db.connect():
|
||||||
|
util.log_error("Can't connect to the database!")
|
||||||
|
return
|
||||||
|
bot.start()
|
||||||
|
await web_server.start()
|
||||||
|
|
||||||
|
# wait for Ctrl+C
|
||||||
await util.get_app_stop_event().wait()
|
await util.get_app_stop_event().wait()
|
||||||
|
|
||||||
|
# stop the app
|
||||||
|
await web_server.stop()
|
||||||
await bot.stop()
|
await bot.stop()
|
||||||
|
await db.disconnect()
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
try:
|
try:
|
||||||
|
|||||||
340
new_bot.py
340
new_bot.py
@@ -1,340 +0,0 @@
|
|||||||
import asyncio
|
|
||||||
import aiofiles
|
|
||||||
import aioconsole
|
|
||||||
import traceback
|
|
||||||
import logging
|
|
||||||
import time
|
|
||||||
import json
|
|
||||||
import os
|
|
||||||
import re
|
|
||||||
from urllib.parse import urlparse
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from nio import AsyncClient, AsyncClientConfig, SyncResponse
|
|
||||||
from nio import LoginResponse, LoginError, WhoamiResponse, WhoamiError
|
|
||||||
|
|
||||||
from bot_types import *
|
|
||||||
|
|
||||||
|
|
||||||
class MatrixBot:
|
|
||||||
"""Asynchronous Matrix Bot Implementation.
|
|
||||||
|
|
||||||
Use objects of this class to build your bots. Manage the event loop
|
|
||||||
by yourself.
|
|
||||||
"""
|
|
||||||
NEXT_BATCH_DUMP_PERIOD = 120.0
|
|
||||||
|
|
||||||
#
|
|
||||||
# PRIVATE
|
|
||||||
#
|
|
||||||
@staticmethod
|
|
||||||
def _validate_matrix_homeserver_url(url: str) -> None:
|
|
||||||
"""Checks if `url` is a valid matrix homeserver URL.
|
|
||||||
Raises an Exception if it is not.
|
|
||||||
"""
|
|
||||||
parsed = urlparse(url)
|
|
||||||
if parsed.scheme not in ("http", "https"):
|
|
||||||
raise RuntimeError(
|
|
||||||
f"Scheme {parsed.scheme} is not a valid scheme for matrix homeserver URL"
|
|
||||||
)
|
|
||||||
if not parsed.netloc:
|
|
||||||
raise RuntimeError(f"{url} is not a valid matrix homeserver URL")
|
|
||||||
if parsed.path != "":
|
|
||||||
raise RuntimeError(f"{url} must have empty path (remove `{parsed.path}` after the hostname)")
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _validate_matrix_username_localpart(localpart: str) -> None:
|
|
||||||
"""Checks if `username` is a valid localpart of matrix username.
|
|
||||||
Raises an Exception if it is not.
|
|
||||||
"""
|
|
||||||
pattern = r"^[a-z0-9._=\-]+$"
|
|
||||||
if not bool(re.fullmatch(pattern, localpart)) or len(localpart) > 255:
|
|
||||||
raise RuntimeError(f"{localpart} is not a valid matrix username localpart")
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _validate_storage_directory(path: Path) -> None:
|
|
||||||
"""Checks if `path` is a valid storage directory and creates it.
|
|
||||||
Raises an Exception if it is not.
|
|
||||||
"""
|
|
||||||
path.mkdir(parents=True, exist_ok=True)
|
|
||||||
if not path.is_dir():
|
|
||||||
raise RuntimeError(f"Could not create directory {path}")
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _validate_bot_config(config: MatrixBotConfig) -> None:
|
|
||||||
"""Checks if `config` has errors.
|
|
||||||
Raises an Exception if it does.
|
|
||||||
"""
|
|
||||||
MatrixBot._validate_matrix_homeserver_url(config.matrix_homeserver_url)
|
|
||||||
MatrixBot._validate_matrix_username_localpart(config.matrix_username_localpart)
|
|
||||||
MatrixBot._validate_storage_directory(config.storage_directory)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _build_client(config: MatrixBotConfig) -> AsyncClient:
|
|
||||||
"""Builds `nio.AsyncClient` from `MatrixBotConfig`"""
|
|
||||||
# create the config for the client
|
|
||||||
client_config = AsyncClientConfig(
|
|
||||||
store_name="nio_store_file",
|
|
||||||
encryption_enabled=True,
|
|
||||||
store_sync_tokens=False
|
|
||||||
)
|
|
||||||
# create the client
|
|
||||||
client = AsyncClient(
|
|
||||||
homeserver=config.matrix_homeserver_url,
|
|
||||||
user=config.matrix_username_localpart,
|
|
||||||
store_path=str(config.storage_directory),
|
|
||||||
config=client_config
|
|
||||||
)
|
|
||||||
return client
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _build_matrix_username(config: MatrixBotConfig) -> str:
|
|
||||||
"""Builds full matrix username."""
|
|
||||||
homeserver_name = urlparse(config.matrix_homeserver_url).hostname
|
|
||||||
localpart = config.matrix_username_localpart
|
|
||||||
return f"@{localpart}:{homeserver_name}"
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
async def _wait_for_task_and_stop_event(payload_task: asyncio.Task, stop_wait_task: asyncio.Task, print_exc: bool = False) -> Any:
|
|
||||||
"""Waits for `payload_task` to finish. Uses `stop_wait_task` to check if payload task
|
|
||||||
should be cancelled. Returns value returned by `payload_task` task. Raises exception
|
|
||||||
raised by `payload_task` task.
|
|
||||||
|
|
||||||
If `stop_wait_task` finishes, then `asyncio.CancelledError` is raised.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
done, _ = await asyncio.wait(
|
|
||||||
[payload_task, stop_wait_task],
|
|
||||||
return_when=asyncio.FIRST_COMPLETED
|
|
||||||
)
|
|
||||||
if stop_wait_task in done:
|
|
||||||
payload_task.cancel()
|
|
||||||
await payload_task
|
|
||||||
raise asyncio.CancelledError()
|
|
||||||
return payload_task.result()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
raise
|
|
||||||
except:
|
|
||||||
if print_exc:
|
|
||||||
traceback.print_exc()
|
|
||||||
raise
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
async def _default_password_callback() -> str:
|
|
||||||
"""Gets password from `MATRIX_PASSWORD` envvar if it is set. Asks
|
|
||||||
the user for the password otherwise."""
|
|
||||||
if "MATRIX_PASSWORD" in os.environ:
|
|
||||||
return os.environ["MATRIX_PASSWORD"]
|
|
||||||
print("--- A password is required (btw you might use MATRIX_PASSWORD envvar) ---")
|
|
||||||
return await aioconsole.ainput("Password: ")
|
|
||||||
|
|
||||||
|
|
||||||
async def _write_next_batch(self, next_batch: str) -> None:
|
|
||||||
"""Writes `next_batch` value to disk."""
|
|
||||||
path = self._config.storage_directory / "next_batch"
|
|
||||||
async with aiofiles.open(path, "w") as f:
|
|
||||||
await f.write(next_batch)
|
|
||||||
self._logger.debug("next_batch value is written to the disk")
|
|
||||||
|
|
||||||
async def _read_next_batch(self) -> str | None:
|
|
||||||
"""Reads `next_batch` value from disk. Returns None if file does not exist."""
|
|
||||||
path = self._config.storage_directory / "next_batch"
|
|
||||||
if not path.is_file():
|
|
||||||
return None
|
|
||||||
async with aiofiles.open(path, "r") as f:
|
|
||||||
return (await f.read()).strip()
|
|
||||||
|
|
||||||
async def _write_session_data(self, *, access_token: str, device_id: str) -> None:
|
|
||||||
"""Write session data to disk."""
|
|
||||||
path = self._config.storage_directory / "session_data.json"
|
|
||||||
data = {
|
|
||||||
"access_token": access_token,
|
|
||||||
"device_id": device_id
|
|
||||||
}
|
|
||||||
async with aiofiles.open(path, "w") as f:
|
|
||||||
await f.write(json.dumps(data, indent=4))
|
|
||||||
self._logger.debug("Session data is writter to the disk")
|
|
||||||
|
|
||||||
async def _read_session_data(self) -> dict[str, Any] | None:
|
|
||||||
"""Read session data from disk."""
|
|
||||||
path = self._config.storage_directory / "session_data.json"
|
|
||||||
if not path.is_file():
|
|
||||||
return None
|
|
||||||
async with aiofiles.open(path, "r") as f:
|
|
||||||
j = json.loads(await f.read())
|
|
||||||
return j
|
|
||||||
|
|
||||||
#
|
|
||||||
# CALLBACKS
|
|
||||||
#
|
|
||||||
async def _callback_sync(self, response: SyncResponse) -> None:
|
|
||||||
"""This callback is called when AsyncClient syncs with the server"""
|
|
||||||
current_time = time.time()
|
|
||||||
delta_time = current_time - self._last_next_batch_dump
|
|
||||||
self._last_next_batch = response.next_batch
|
|
||||||
if delta_time >= self.NEXT_BATCH_DUMP_PERIOD:
|
|
||||||
self._last_next_batch_dump = current_time
|
|
||||||
try:
|
|
||||||
await self._write_next_batch(self._last_next_batch)
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
#
|
|
||||||
# LIFECYCLE
|
|
||||||
#
|
|
||||||
def _setup_client_callbacks(self) -> None:
|
|
||||||
# setup the callbacks
|
|
||||||
self._client.add_response_callback(self._callback_sync, SyncResponse) # type: ignore
|
|
||||||
|
|
||||||
async def _client_login_session_data(self, session_data: dict[str, Any]) -> None:
|
|
||||||
"""Login using session data. Raises and exception on failure."""
|
|
||||||
self._logger.debug("Using stored session data to log in")
|
|
||||||
# build user id
|
|
||||||
username = self._build_matrix_username(self._config)
|
|
||||||
self._client.restore_login(
|
|
||||||
user_id=username,
|
|
||||||
**session_data
|
|
||||||
)
|
|
||||||
result = await self._client.whoami()
|
|
||||||
if type(result) is WhoamiError:
|
|
||||||
self._logger.error(f"Can't log in using stored session data: '{result.message}'")
|
|
||||||
raise RuntimeError(result.message)
|
|
||||||
elif type(result) is not WhoamiResponse:
|
|
||||||
self._logger.error("Can't log in using stored session data, unknown error")
|
|
||||||
raise RuntimeError("Unknown response for whoami request")
|
|
||||||
self._logger.debug("Logged in using stored session data")
|
|
||||||
|
|
||||||
async def _client_login_password(self) -> None:
|
|
||||||
"""Login using password and save result to disk on success.
|
|
||||||
Raises an exception on failure.
|
|
||||||
"""
|
|
||||||
self._logger.debug("Using password to log in")
|
|
||||||
# get the password
|
|
||||||
password = await self._cb_password()
|
|
||||||
result = await self._client.login(password=password)
|
|
||||||
if type(result) is LoginResponse:
|
|
||||||
self._logger.debug("Logged in using password")
|
|
||||||
await self._write_session_data(
|
|
||||||
access_token=result.access_token,
|
|
||||||
device_id=result.device_id
|
|
||||||
)
|
|
||||||
elif type(result) is LoginError:
|
|
||||||
self._logger.error(f"Can't log in using password: '{result.message}'")
|
|
||||||
raise RuntimeError(result.message)
|
|
||||||
else:
|
|
||||||
self._logger.error(f"Can't log in using password, unknown error")
|
|
||||||
raise RuntimeError("Unknown login result")
|
|
||||||
|
|
||||||
async def _client_login(self) -> None:
|
|
||||||
"""This function logs in."""
|
|
||||||
# check if we have session data stored on the disk
|
|
||||||
session_data = await self._read_session_data()
|
|
||||||
# session data is present, try to log in
|
|
||||||
if session_data is not None:
|
|
||||||
self._logger.debug("Some session data found on the disk")
|
|
||||||
try:
|
|
||||||
await self._client_login_session_data(session_data)
|
|
||||||
return
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
raise
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
# no session data - login using password
|
|
||||||
try:
|
|
||||||
self._logger.debug("No session data found on the disk OR invalid data")
|
|
||||||
await self._client_login_password()
|
|
||||||
return
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
raise
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
# can't login
|
|
||||||
self._logger.error("Can't log in using available methods")
|
|
||||||
raise RuntimeError("All login methods have failed, can't continue")
|
|
||||||
|
|
||||||
async def _client_destroy(self) -> None:
|
|
||||||
"""Gracefully destroys the client."""
|
|
||||||
try:
|
|
||||||
self._logger.debug("Closing the client")
|
|
||||||
await self._client.close()
|
|
||||||
if self._last_next_batch is not None:
|
|
||||||
self._logger.debug("Saving next_batch")
|
|
||||||
await self._write_next_batch(self._last_next_batch)
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
async def _background_coroutine(self) -> None:
|
|
||||||
"""This function implements bot lifecycle."""
|
|
||||||
# we should stop when this task stops
|
|
||||||
self._logger.debug("_background_coroutine is started")
|
|
||||||
stop_wait_task = asyncio.create_task(self._stop_event.wait())
|
|
||||||
self._client = self._build_client(self._config)
|
|
||||||
self._setup_client_callbacks()
|
|
||||||
# perform login
|
|
||||||
login_task = asyncio.create_task(self._client_login())
|
|
||||||
try:
|
|
||||||
await self._wait_for_task_and_stop_event(login_task, stop_wait_task)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
self._logger.debug("Background task is cancelled during login")
|
|
||||||
await self._client_destroy()
|
|
||||||
return
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
self._logger.info("Succesfully logged in")
|
|
||||||
# test
|
|
||||||
wait_task = asyncio.create_task(asyncio.sleep(100000))
|
|
||||||
try:
|
|
||||||
self._logger.info("Bot is not implemented yet, sleeping forever")
|
|
||||||
await self._wait_for_task_and_stop_event(wait_task, stop_wait_task)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
self._logger.info("Background task is cancelled during eternal sleep")
|
|
||||||
await self._client_destroy()
|
|
||||||
return
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
|
|
||||||
#
|
|
||||||
# PUBLIC
|
|
||||||
#
|
|
||||||
def __init__(self, config: MatrixBotConfig) -> None:
|
|
||||||
# check if config is valid
|
|
||||||
self._validate_bot_config(config) # may raise an Exception
|
|
||||||
# save the config
|
|
||||||
self._config = config
|
|
||||||
# create the logger
|
|
||||||
self._logger = logging.getLogger(self._build_matrix_username(config))
|
|
||||||
self._logger.setLevel(logging.DEBUG)
|
|
||||||
|
|
||||||
# prepare some private data
|
|
||||||
self._background_task: asyncio.Task | None = None
|
|
||||||
self._client: AsyncClient = None # type: ignore
|
|
||||||
self._last_next_batch_dump: float = 0.0
|
|
||||||
self._last_next_batch: str | None = None
|
|
||||||
self._cb_password = self._default_password_callback
|
|
||||||
|
|
||||||
def start(self) -> None:
|
|
||||||
"""Start the bot.
|
|
||||||
Starts the bot in background task. Raises an exception if there are
|
|
||||||
problems (for example, the bot is already started). The bot will
|
|
||||||
do everything to keep itself running, including restarts. Use
|
|
||||||
`stop()` to stop the bot.
|
|
||||||
"""
|
|
||||||
if self._background_task is not None:
|
|
||||||
raise RuntimeError("The bot is already started!")
|
|
||||||
self._stop_event = asyncio.Event()
|
|
||||||
self._background_task = asyncio.create_task(
|
|
||||||
self._background_coroutine()
|
|
||||||
)
|
|
||||||
|
|
||||||
async def stop(self) -> None:
|
|
||||||
"""Stop the bot and wait for the bot stop."""
|
|
||||||
if self._background_task is None:
|
|
||||||
return
|
|
||||||
self._stop_event.set()
|
|
||||||
try:
|
|
||||||
await self._background_task
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
self._stop_event = None
|
|
||||||
self._background_task = None
|
|
||||||
@@ -1,3 +1,75 @@
|
|||||||
aioconsole
|
agent-detector==1.1.0
|
||||||
aiofiles
|
aioconsole==0.8.2
|
||||||
matrix-nio[e2e]
|
aiofiles==25.1.0
|
||||||
|
aiohappyeyeballs==2.7.1
|
||||||
|
aiohttp==3.14.3
|
||||||
|
aiohttp_socks==0.12.0
|
||||||
|
aiosignal==1.4.0
|
||||||
|
aiosqlite==0.22.1
|
||||||
|
annotated-doc==0.0.5
|
||||||
|
annotated-types==0.8.0
|
||||||
|
anyio==4.14.2
|
||||||
|
atomicwrites==1.4.1
|
||||||
|
attrs==26.1.0
|
||||||
|
build==1.5.0
|
||||||
|
cachetools==7.1.7
|
||||||
|
certifi==2026.7.22
|
||||||
|
click==8.5.0
|
||||||
|
detect-installer==0.1.0
|
||||||
|
dnspython==2.8.0
|
||||||
|
email-validator==2.3.0
|
||||||
|
fastapi==0.141.1
|
||||||
|
fastapi-cli==0.0.32
|
||||||
|
fastapi-cloud-cli==0.24.0
|
||||||
|
fastar==0.12.0
|
||||||
|
frozenlist==1.8.0
|
||||||
|
h11==0.16.0
|
||||||
|
h2==4.4.1
|
||||||
|
hpack==4.2.0
|
||||||
|
httpcore==1.0.9
|
||||||
|
httptools==0.8.0
|
||||||
|
httpx==0.28.1
|
||||||
|
hyperframe==6.1.0
|
||||||
|
idna==3.19
|
||||||
|
Jinja2==3.1.6
|
||||||
|
jsonschema==4.26.0
|
||||||
|
jsonschema-specifications==2025.9.1
|
||||||
|
mab @ git+https://git.tyukalov.su/nikita/mab@v0.0.2
|
||||||
|
markdown-it-py==4.2.0
|
||||||
|
MarkupSafe==3.0.3
|
||||||
|
matrix-nio==0.26.0
|
||||||
|
mdurl==0.1.2
|
||||||
|
multidict==6.7.1
|
||||||
|
packaging==26.3
|
||||||
|
peewee==3.19.0
|
||||||
|
propcache==0.5.2
|
||||||
|
pycryptodome==3.23.0
|
||||||
|
pydantic==2.13.5
|
||||||
|
pydantic-extra-types==2.11.1
|
||||||
|
pydantic-settings==2.15.0
|
||||||
|
pydantic_core==2.46.5
|
||||||
|
Pygments==2.21.0
|
||||||
|
pyproject_hooks==1.2.0
|
||||||
|
python-dotenv==1.2.3
|
||||||
|
python-multipart==0.0.32
|
||||||
|
python-socks==3.0.0
|
||||||
|
PyYAML==6.0.3
|
||||||
|
referencing==0.37.0
|
||||||
|
rich==15.0.0
|
||||||
|
rich-toolkit==0.20.3
|
||||||
|
rignore==0.8.1
|
||||||
|
rpds-py==2026.6.3
|
||||||
|
sentry-sdk==2.68.1
|
||||||
|
shellingham==1.5.4
|
||||||
|
starlette==1.6.0
|
||||||
|
typer==0.27.2
|
||||||
|
typing-inspection==0.4.4
|
||||||
|
typing_extensions==4.16.0
|
||||||
|
unpaddedbase64==2.1.0
|
||||||
|
urllib3==2.7.0
|
||||||
|
uvicorn==0.52.4
|
||||||
|
uvloop==0.22.1
|
||||||
|
vodozemac==0.10.0
|
||||||
|
watchfiles==1.2.0
|
||||||
|
websockets==17.1
|
||||||
|
yarl==1.24.5
|
||||||
|
|||||||
124
util.py
124
util.py
@@ -5,14 +5,11 @@ Utilities
|
|||||||
import asyncio
|
import asyncio
|
||||||
import sys
|
import sys
|
||||||
import os
|
import os
|
||||||
import json
|
|
||||||
import logging
|
import logging
|
||||||
import traceback
|
import traceback
|
||||||
from html.parser import HTMLParser
|
import datetime
|
||||||
from pathlib import Path
|
import hashlib
|
||||||
from urllib.parse import urlparse
|
|
||||||
|
|
||||||
import aiofiles
|
|
||||||
import aioconsole
|
import aioconsole
|
||||||
|
|
||||||
from datatypes import AppConfig
|
from datatypes import AppConfig
|
||||||
@@ -72,61 +69,6 @@ async def ainput(text: str = "") -> str:
|
|||||||
return ""
|
return ""
|
||||||
return await aioconsole.ainput(text)
|
return await aioconsole.ainput(text)
|
||||||
|
|
||||||
async def set_next_batch(config: AppConfig, next_batch: str | None) -> bool:
|
|
||||||
"""Save `next_batch` to session directory."""
|
|
||||||
try:
|
|
||||||
next_batch_file_path = Path.cwd() / config.store_dir / "next_batch.txt"
|
|
||||||
if next_batch is None:
|
|
||||||
next_batch_file_path.unlink(True)
|
|
||||||
return True
|
|
||||||
async with aiofiles.open(next_batch_file_path, "w") as f:
|
|
||||||
await f.write(next_batch)
|
|
||||||
return True
|
|
||||||
except:
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def get_next_batch(config: AppConfig) -> str | None:
|
|
||||||
"""Get `next_batch`"""
|
|
||||||
try:
|
|
||||||
next_batch_file_path = Path.cwd() / config.store_dir / "next_batch.txt"
|
|
||||||
async with aiofiles.open(next_batch_file_path, "r") as f:
|
|
||||||
next_batch = (await f.read()).strip()
|
|
||||||
if not next_batch:
|
|
||||||
next_batch = None
|
|
||||||
return next_batch
|
|
||||||
except:
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def get_session_data(config: AppConfig) -> tuple[str, str] | tuple[None, None]:
|
|
||||||
"""
|
|
||||||
Get (access_token, device_id) or (None, None)
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
token_file_path = Path.cwd() / config.store_dir / "auth.json"
|
|
||||||
async with aiofiles.open(token_file_path, "r") as f:
|
|
||||||
j = json.loads((await f.read()).strip())
|
|
||||||
return j["access_token"], j["device_id"]
|
|
||||||
except:
|
|
||||||
return None, None
|
|
||||||
|
|
||||||
async def set_session_data(config: AppConfig, token_device_pair: tuple[str, str] | None) -> bool:
|
|
||||||
"""
|
|
||||||
Set new (access_token, device_id) pair; use None to remove it.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True on success
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
token_file_path = Path.cwd() / config.store_dir / "auth.json"
|
|
||||||
if token_device_pair is None:
|
|
||||||
token_file_path.unlink(True)
|
|
||||||
return True
|
|
||||||
async with aiofiles.open(token_file_path, "w") as f:
|
|
||||||
await f.write(json.dumps({"access_token": token_device_pair[0], "device_id": token_device_pair[1]}))
|
|
||||||
return True
|
|
||||||
except:
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def get_password() -> str | None:
|
async def get_password() -> str | None:
|
||||||
if "MATRIX_PASSWORD" in os.environ:
|
if "MATRIX_PASSWORD" in os.environ:
|
||||||
return os.environ["MATRIX_PASSWORD"]
|
return os.environ["MATRIX_PASSWORD"]
|
||||||
@@ -134,31 +76,41 @@ async def get_password() -> str | None:
|
|||||||
return None
|
return None
|
||||||
return await ainput("Matrix password: ")
|
return await ainput("Matrix password: ")
|
||||||
|
|
||||||
def get_hostname_from_url(url: str) -> str | None:
|
def date_to_text(date: datetime.datetime | float, dow: bool = True, seconds: bool = True) -> str:
|
||||||
"""Returns `matrix.domain.net` for `https://matrix.domain.net/bla/bla/bla`"""
|
''' Returns date as formatted string.
|
||||||
try:
|
Day of week can be added.
|
||||||
return urlparse(url).hostname
|
Seconds can be added.
|
||||||
except:
|
'''
|
||||||
return None
|
if type(date) is float:
|
||||||
|
date = int(date)
|
||||||
|
if type(date) is int:
|
||||||
|
date = datetime.datetime.utcfromtimestamp(date)
|
||||||
|
|
||||||
def check_and_remove_html(possible_html: str) -> tuple[bool, str]:
|
# format string
|
||||||
"""Checks if `possible_html` is a valid HTML text and returns (is_html, text_without_tags)"""
|
format_string = ''
|
||||||
has_tags = False
|
if dow:
|
||||||
text_fragments = []
|
format_string += '%a, '
|
||||||
class Extractor(HTMLParser):
|
format_string += '%d.%m.%Y, %H:%M'
|
||||||
def handle_starttag(self, tag, attrs):
|
if seconds:
|
||||||
nonlocal has_tags
|
format_string += ':%S'
|
||||||
has_tags = True
|
# day of week to Russian
|
||||||
def handle_data(self, data):
|
translate_map = [
|
||||||
text_fragments.append(data)
|
('Mon', 'Пн'),
|
||||||
|
('Tue', 'Вт'),
|
||||||
|
('Wed', 'Ср'),
|
||||||
|
('Thu', 'Чт'),
|
||||||
|
('Fri', 'Пт'),
|
||||||
|
('Sat', 'Сб'),
|
||||||
|
('Sun', 'Вс')
|
||||||
|
]
|
||||||
|
result = date.strftime(format_string) # type: ignore
|
||||||
|
for en, ru in translate_map:
|
||||||
|
if en in result:
|
||||||
|
result = result.replace(en, ru)
|
||||||
|
break
|
||||||
|
return result
|
||||||
|
|
||||||
parser = Extractor(convert_charrefs=True)
|
def get_hash(data: bytes | str) -> str:
|
||||||
parser.feed(possible_html)
|
if type(data) is str:
|
||||||
|
data = data.encode(errors="ignore")
|
||||||
try:
|
return hashlib.sha256(data).hexdigest() # type: ignore
|
||||||
if has_tags:
|
|
||||||
return (True, " ".join("".join(text_fragments).split()))
|
|
||||||
except:
|
|
||||||
traceback.print_exc()
|
|
||||||
|
|
||||||
return (False, possible_html)
|
|
||||||
73
web.py
Normal file
73
web.py
Normal file
@@ -0,0 +1,73 @@
|
|||||||
|
"""This module implements API"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import traceback
|
||||||
|
import time
|
||||||
|
import html
|
||||||
|
|
||||||
|
from fastapi import FastAPI, Request
|
||||||
|
from fastapi.responses import JSONResponse
|
||||||
|
|
||||||
|
from mab import MatrixBot
|
||||||
|
import uvicorn
|
||||||
|
|
||||||
|
from datatypes import AppConfig
|
||||||
|
from database import Database
|
||||||
|
|
||||||
|
import web_routes
|
||||||
|
import web_middleware
|
||||||
|
|
||||||
|
class Web:
|
||||||
|
#
|
||||||
|
# PRIVATE
|
||||||
|
#
|
||||||
|
async def _cb_exception(self, request: Request, exc: Exception):
|
||||||
|
traceback.print_exception(exc)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=500,
|
||||||
|
content={"detail": "Internal Server Error"}
|
||||||
|
)
|
||||||
|
|
||||||
|
#
|
||||||
|
# PUBLIC
|
||||||
|
#
|
||||||
|
def __init__(self, app_config: AppConfig, bot: MatrixBot, db: Database):
|
||||||
|
self._api = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
self._api.state.database = db
|
||||||
|
self._api.state.matrix = bot
|
||||||
|
self._api.state.app_config = app_config
|
||||||
|
|
||||||
|
self._api.include_router(web_routes.router)
|
||||||
|
self._api.add_middleware(web_middleware.AuthMiddleware, ["/api"])
|
||||||
|
self._api.add_middleware(web_middleware.BanCheckMiddleware, ["/api"])
|
||||||
|
self._api.add_middleware(web_middleware.RealIpResolver)
|
||||||
|
self._api.add_exception_handler(Exception, self._cb_exception)
|
||||||
|
|
||||||
|
self._server_config = uvicorn.Config(
|
||||||
|
self._api,
|
||||||
|
host=app_config.web_ip,
|
||||||
|
port=app_config.web_port
|
||||||
|
)
|
||||||
|
self._server: uvicorn.Server | None = None
|
||||||
|
self._server_task: asyncio.Task | None = None
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
"""Start the API."""
|
||||||
|
if self._server is not None or self._server_task is not None:
|
||||||
|
raise RuntimeError("The server is already started")
|
||||||
|
self._server = uvicorn.Server(self._server_config)
|
||||||
|
self._server_task = asyncio.create_task(self._server.serve())
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
"""Stop the API server."""
|
||||||
|
if self._server is None or self._server_task is None:
|
||||||
|
raise RuntimeError("The server is not started yet")
|
||||||
|
self._server_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._server_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
except:
|
||||||
|
traceback.print_exc()
|
||||||
|
self._server_task = None
|
||||||
|
self._server = None
|
||||||
135
web_middleware.py
Normal file
135
web_middleware.py
Normal file
@@ -0,0 +1,135 @@
|
|||||||
|
import time
|
||||||
|
|
||||||
|
from fastapi import Request
|
||||||
|
from fastapi import status
|
||||||
|
from fastapi.responses import JSONResponse
|
||||||
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
|
from database import Database
|
||||||
|
from datatypes import ObjectToken, AppConfig
|
||||||
|
|
||||||
|
|
||||||
|
class AuthMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""Middleware for authorization"""
|
||||||
|
def _is_request_protected(self, request: Request) -> bool:
|
||||||
|
"""Returns True if request URL starts with one of private prefixes"""
|
||||||
|
for pp in self.private_prefixes:
|
||||||
|
if request.url.path.startswith(pp):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _fail_and_ban_if_needed(self, request: Request) -> None:
|
||||||
|
"""Adds a fail and bans the IP if too much failures are recorded."""
|
||||||
|
app_config: AppConfig = request.app.state.app_config
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
current_time: float = time.time()
|
||||||
|
ip: str = request.state.ip
|
||||||
|
total_fails = database.fail_create(ip, current_time + 10.0)
|
||||||
|
if total_fails >= app_config.fails_to_ban:
|
||||||
|
await database.ban_create(
|
||||||
|
ip,
|
||||||
|
current_time + app_config.ban_duration,
|
||||||
|
"Banned by AuthMiddleware"
|
||||||
|
)
|
||||||
|
database.fail_clear(ip)
|
||||||
|
|
||||||
|
def __init__(self, app, private_prefixes: list[str]):
|
||||||
|
super().__init__(app)
|
||||||
|
self.private_prefixes = private_prefixes
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
# do not check authorization for public endpoints
|
||||||
|
if not self._is_request_protected(request):
|
||||||
|
return await call_next(request)
|
||||||
|
token = None
|
||||||
|
# look for the token in headers
|
||||||
|
if "Authorization" in request.headers:
|
||||||
|
v = request.headers["Authorization"]
|
||||||
|
if not v.lower().startswith("bearer"):
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "Invalid authorization scheme"}
|
||||||
|
)
|
||||||
|
parts = [p for p in v.split(" ") if p.strip()]
|
||||||
|
if len(parts) != 2:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "No credentials provided in HTTP header"}
|
||||||
|
)
|
||||||
|
token = parts[1]
|
||||||
|
# look for the token in query parameters
|
||||||
|
elif "token" in request.query_params:
|
||||||
|
token = request.query_params.get("token")
|
||||||
|
if not token:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "The token is empry"}
|
||||||
|
)
|
||||||
|
# no token found
|
||||||
|
else:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "No token provided"}
|
||||||
|
)
|
||||||
|
# try to authorize
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
token_data: ObjectToken | None = await database.token_get(token)
|
||||||
|
if not token_data:
|
||||||
|
await self._fail_and_ban_if_needed(request)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "No such token"}
|
||||||
|
)
|
||||||
|
request.state.token = token_data
|
||||||
|
return await call_next(request)
|
||||||
|
|
||||||
|
|
||||||
|
class BanCheckMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""Middleware for checking if the IP is banned."""
|
||||||
|
def _is_request_protected(self, request: Request) -> bool:
|
||||||
|
"""Returns True if request URL starts with one of private prefixes"""
|
||||||
|
for pp in self.private_prefixes:
|
||||||
|
if request.url.path.startswith(pp):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def __init__(self, app, private_prefixes: list[str]):
|
||||||
|
super().__init__(app)
|
||||||
|
self.private_prefixes = private_prefixes
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
# do not check authorization for public endpoints
|
||||||
|
if not self._is_request_protected(request):
|
||||||
|
return await call_next(request)
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
if database.ban_check(request.state.ip):
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN,
|
||||||
|
content={"detail": "You are temporarily banned"}
|
||||||
|
)
|
||||||
|
return await call_next(request)
|
||||||
|
|
||||||
|
|
||||||
|
class RealIpResolver(BaseHTTPMiddleware):
|
||||||
|
"""Middleware that gets real IP address of the client."""
|
||||||
|
def __init__(self, app):
|
||||||
|
super().__init__(app)
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
# assign None by default
|
||||||
|
request.state.ip = None
|
||||||
|
# look for IP address
|
||||||
|
x_forwarded_for = request.headers.get("x-forwarded-for")
|
||||||
|
if x_forwarded_for:
|
||||||
|
request.state.ip = x_forwarded_for.split(",")[0].strip()
|
||||||
|
else:
|
||||||
|
x_real_ip = request.headers.get("x-real-ip")
|
||||||
|
if x_real_ip:
|
||||||
|
request.state.ip = x_real_ip
|
||||||
|
else:
|
||||||
|
request.state.ip = request.client.host if request.client else None
|
||||||
|
return await call_next(request)
|
||||||
70
web_routes.py
Normal file
70
web_routes.py
Normal file
@@ -0,0 +1,70 @@
|
|||||||
|
"""API Routes"""
|
||||||
|
|
||||||
|
import html
|
||||||
|
import traceback
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Request, Depends
|
||||||
|
from fastapi import HTTPException, status
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from mab import MatrixBot
|
||||||
|
from database import Database
|
||||||
|
from datatypes import *
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/api")
|
||||||
|
|
||||||
|
#
|
||||||
|
# MODELS
|
||||||
|
#
|
||||||
|
class ModelGetNotify(BaseModel):
|
||||||
|
channel: str = Field(
|
||||||
|
...,
|
||||||
|
min_length=11,
|
||||||
|
max_length=11,
|
||||||
|
description="Notification channel (use `!info` bot command)"
|
||||||
|
)
|
||||||
|
text: str = Field(
|
||||||
|
...,
|
||||||
|
min_length=1,
|
||||||
|
max_length=1024,
|
||||||
|
description="Text to use as notification body"
|
||||||
|
)
|
||||||
|
service: str | None = Field(
|
||||||
|
None,
|
||||||
|
min_length=1,
|
||||||
|
max_length=1024,
|
||||||
|
description="Service name to use (ignored if token has name)"
|
||||||
|
)
|
||||||
|
|
||||||
|
#
|
||||||
|
# ENDPOINTS
|
||||||
|
#
|
||||||
|
@router.get("/notify")
|
||||||
|
async def _get_notify(request: Request, params: ModelGetNotify = Depends()):
|
||||||
|
# improve readability
|
||||||
|
database: Database = request.app.state.database
|
||||||
|
matrix: MatrixBot = request.app.state.matrix
|
||||||
|
token: ObjectToken = request.state.token
|
||||||
|
# check if room exists
|
||||||
|
room = await database.room_get(params.channel)
|
||||||
|
if room is None:
|
||||||
|
raise HTTPException(status.HTTP_200_OK, "No such channel")
|
||||||
|
# prepare service name
|
||||||
|
service = token.name or params.service or "Unnamed"
|
||||||
|
# prepare text of the notification
|
||||||
|
text = f"<strong>Уведомление от службы <code>{html.escape(service)}</code></strong>"
|
||||||
|
text += "<br>" * 2
|
||||||
|
text += html.escape(params.text)
|
||||||
|
# try to send the notification
|
||||||
|
try:
|
||||||
|
notification_id = await matrix.send_text_to_room(
|
||||||
|
room.matrix_id,
|
||||||
|
text,
|
||||||
|
is_html=True
|
||||||
|
)
|
||||||
|
except:
|
||||||
|
traceback.print_exc()
|
||||||
|
raise HTTPException(status.HTTP_200_OK, "Failed to send matrix message")
|
||||||
|
# success
|
||||||
|
return {"notification_id": notification_id}
|
||||||
Reference in New Issue
Block a user