Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions tests/api/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import asyncio
import json
from pathlib import Path

from alembic.command import upgrade
from alembic.config import Config
Expand Down Expand Up @@ -102,6 +103,16 @@ async def create_tables():


if TEST_FROM == "local":
if IS_SQLITE:
try:
database_file = Path(DATABASE_URL.partition("///")[2].split("?", 1)[0])
for suffix in ("", "-shm", "-wal"):
try:
database_file.with_name(f"{database_file.name}{suffix}").unlink()
except FileNotFoundError:
pass
except Exception:
pass
run_migrations()


Expand Down
123 changes: 123 additions & 0 deletions tests/telegram/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,123 @@
from __future__ import annotations

from dataclasses import dataclass
from types import SimpleNamespace
from unittest.mock import AsyncMock

import pytest

from app.db.models import UserStatus
from app.models.admin import AdminDetails


@dataclass
class FakeChat:
id: int = 1


@dataclass
class FakeFromUser:
full_name: str = "Test User"


class FakeState:
def __init__(self):
self._data: dict = {}
self._state = None

async def set_state(self, state):
self._state = state

async def get_state(self):
return self._state

async def clear(self):
self._state = None
self._data = {}

async def update_data(self, **kwargs):
self._data.update(kwargs)

async def get_data(self):
return dict(self._data)

async def get_value(self, key, default=None):
return self._data.get(key, default)


class FakeMessage:
def __init__(self, text: str = "text", message_id: int = 100, chat_id: int = 1):
self.text = text
self.message_id = message_id
self.chat = FakeChat(chat_id)
self.from_user = FakeFromUser()
self.bot = SimpleNamespace(delete_messages=AsyncMock())
self.answer = AsyncMock(side_effect=self._new_message)
self.reply = AsyncMock(side_effect=self._new_message)
self.edit_text = AsyncMock()
self.edit_reply_markup = AsyncMock()
self.answer_document = AsyncMock()
self.answer_photo = AsyncMock()
self.delete = AsyncMock()

async def _new_message(self, *args, **kwargs):
return FakeMessage(message_id=self.message_id + 1, chat_id=self.chat.id)


class FakeCallbackQuery:
def __init__(self, message: FakeMessage | None = None, data: str = "callback"):
self.message = message or FakeMessage()
self.data = data
self.from_user = self.message.from_user
self.bot = self.message.bot
self.answer = AsyncMock()


class FakeInlineQuery:
def __init__(self, query: str = ""):
self.query = query
self.answer = AsyncMock()


@pytest.fixture
def fake_state() -> FakeState:
return FakeState()


@pytest.fixture
def fake_message() -> FakeMessage:
return FakeMessage()


@pytest.fixture
def fake_callback() -> FakeCallbackQuery:
return FakeCallbackQuery()


@pytest.fixture
def fake_inline_query() -> FakeInlineQuery:
return FakeInlineQuery()


@pytest.fixture
def admin() -> AdminDetails:
return AdminDetails(username="testadmin", is_sudo=True)


@pytest.fixture
def regular_admin() -> AdminDetails:
return AdminDetails(username="owner", is_sudo=False)


@pytest.fixture
def fake_user():
user = SimpleNamespace(
id=11,
username="alice",
status=UserStatus.active,
subscription_url="https://example.com/sub/alice",
next_plan=None,
admin=SimpleNamespace(username="owner"),
)
user.model_dump = lambda: {}
return user
226 changes: 226 additions & 0 deletions tests/telegram/handlers/admin/test_bulk_actions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,226 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock

import pytest

from app.models.user import UsernameGenerationStrategy
from app.telegram.handlers.admin import bulk_actions as bulk_h
from app.telegram.keyboards.bulk_actions import (
BulkAction,
BulkActionPanel,
BulkTemplateSelector,
UsernameStrategySelector,
)


class _TextValue(str):
def __call__(self, *args, **kwargs):
return str(self)


class _DummyTexts:
def __getattr__(self, name):
return _TextValue(name)


@pytest.fixture(autouse=True)
def patch_event_types(monkeypatch, fake_message, fake_callback):
monkeypatch.setattr(bulk_h, "Texts", _DummyTexts())
monkeypatch.setattr(bulk_h, "Message", type(fake_message))
monkeypatch.setattr(bulk_h, "CallbackQuery", type(fake_callback))
monkeypatch.setattr(bulk_h, "delete_messages", AsyncMock())
monkeypatch.setattr(bulk_h, "add_to_messages_to_delete", AsyncMock())


@pytest.mark.asyncio
async def test_helpers_message_target_and_chunking(fake_message, fake_callback):
msg = type(fake_message)()
callback = type(fake_callback)(message=msg)

assert bulk_h._message_target(msg) is msg
assert bulk_h._message_target(callback) is msg
chunks = bulk_h._chunk_subscription_urls(["a" * 3000, "b" * 3000], limit=3800)
assert len(chunks) == 2


@pytest.mark.asyncio
async def test_bulk_actions_menu(admin, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())

await bulk_h.bulk_actions(event, admin=admin)

event.message.edit_text.assert_awaited_once()


@pytest.mark.asyncio
async def test_bulk_create_from_template_no_templates(monkeypatch, fake_state, admin, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())
monkeypatch.setattr(bulk_h.user_templates, "get_user_templates", AsyncMock(return_value=[]))

await bulk_h.bulk_create_from_template(event, db=object(), state=fake_state, admin=admin)

event.answer.assert_awaited_once()


@pytest.mark.asyncio
async def test_bulk_template_chosen(fake_state, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())
callback_data = BulkTemplateSelector.Callback(template_id=1)

await bulk_h.bulk_template_chosen(event, state=fake_state, callback_data=callback_data)

assert await fake_state.get_state() == bulk_h.forms.BulkCreateFromTemplate.count


@pytest.mark.asyncio
async def test_bulk_template_count_invalid(fake_state, fake_message):
event = type(fake_message)(text="x")

await bulk_h.bulk_template_count(event, state=fake_state)

event.reply.assert_awaited_once()


@pytest.mark.asyncio
async def test_bulk_template_strategy_random(monkeypatch, fake_state, admin, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())
callback_data = UsernameStrategySelector.Callback(strategy=UsernameGenerationStrategy.random)
fake_state._data = {"template_id": 1, "count": 2}
monkeypatch.setattr(bulk_h, "delete_messages", AsyncMock())
monkeypatch.setattr(bulk_h, "_perform_bulk_creation", AsyncMock())

await bulk_h.bulk_template_strategy(event, db=object(), state=fake_state, admin=admin, callback_data=callback_data)

bulk_h._perform_bulk_creation.assert_awaited_once()


@pytest.mark.asyncio
async def test_bulk_template_sequence_username_invalid(fake_state, admin, fake_message):
event = type(fake_message)(text="bad name")

await bulk_h.bulk_template_sequence_username(event, db=object(), state=fake_state, admin=admin)

event.reply.assert_awaited_once()


@pytest.mark.asyncio
async def test_bulk_template_start_number_invalid(fake_state, admin, fake_message):
event = type(fake_message)(text="-2")

await bulk_h.bulk_template_start_number(event, db=object(), state=fake_state, admin=admin)

event.reply.assert_awaited_once()


@pytest.mark.asyncio
async def test_perform_bulk_creation_success(monkeypatch, fake_state, admin, fake_message):
event = type(fake_message)(text="run")
monkeypatch.setattr(bulk_h, "delete_messages", AsyncMock())
monkeypatch.setattr(
bulk_h.user_operations,
"bulk_create_users_from_template",
AsyncMock(return_value=SimpleNamespace(created=2, subscription_urls=["https://a", "https://b"])),
)

await bulk_h._perform_bulk_creation(
event,
db=object(),
admin=admin,
state=fake_state,
template_id=1,
count=2,
strategy=UsernameGenerationStrategy.random,
)

assert event.answer.await_count >= 2


@pytest.mark.asyncio
async def test_delete_expired_sets_state(fake_state, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())

await bulk_h.delete_expired(event, state=fake_state)

assert await fake_state.get_state() == bulk_h.forms.DeleteExpired.expired_before


@pytest.mark.asyncio
async def test_process_expire_before_validation(fake_state, fake_message):
event = type(fake_message)(text="abc")

await bulk_h.process_expire_before(event, state=fake_state)

event.reply.assert_awaited_once()


@pytest.mark.asyncio
async def test_delete_expired_done(monkeypatch, admin, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())
callback_data = BulkActionPanel.Callback(action=BulkAction.delete_expired, amount="3")
monkeypatch.setattr(
bulk_h.user_operations,
"delete_expired_users",
AsyncMock(return_value=SimpleNamespace(count=5)),
)

await bulk_h.delete_expired_done(event, db=object(), admin=admin, callback_data=callback_data)

event.message.edit_text.assert_awaited_once()


@pytest.mark.asyncio
async def test_modify_expiry_sets_state(fake_state, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())

await bulk_h.modify_expiry(event, state=fake_state)

assert await fake_state.get_state() == bulk_h.forms.BulkModify.expiry


@pytest.mark.asyncio
async def test_process_expiry_invalid(fake_state, fake_message):
event = type(fake_message)(text="bad")

await bulk_h.process_expiry(event, state=fake_state)

event.reply.assert_awaited_once()


@pytest.mark.asyncio
async def test_modify_expiry_done(monkeypatch, admin, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())
callback_data = BulkActionPanel.Callback(action=BulkAction.modify_expiry, amount="5")
monkeypatch.setattr(bulk_h.user_operations, "bulk_modify_expire", AsyncMock(return_value=7))

await bulk_h.modify_expiry_done(event, db=object(), admin=admin, callback_data=callback_data)

event.message.edit_text.assert_awaited_once()


@pytest.mark.asyncio
async def test_modify_data_limit_sets_state(fake_state, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())

await bulk_h.modify_data_limit(event, state=fake_state)

assert await fake_state.get_state() == bulk_h.forms.BulkModify.data_limit


@pytest.mark.asyncio
async def test_process_data_limit_invalid(fake_state, fake_message):
event = type(fake_message)(text="bad")

await bulk_h.process_data_limit(event, state=fake_state)

event.reply.assert_awaited_once()


@pytest.mark.asyncio
async def test_modify_data_limit_done(monkeypatch, admin, fake_message, fake_callback):
event = type(fake_callback)(message=type(fake_message)())
callback_data = BulkActionPanel.Callback(action=BulkAction.modify_data_limit, amount="4")
monkeypatch.setattr(bulk_h.user_operations, "bulk_modify_datalimit", AsyncMock(return_value=9))

await bulk_h.modify_data_limit_done(event, db=object(), admin=admin, callback_data=callback_data)

event.message.edit_text.assert_awaited_once()
Loading