Add rate limiting and penetration check

This commit is contained in:
IgorVolochay
2026-08-24 18:44:18 +03:00
parent 68fcc89153
commit 48fd5ea19e
8 changed files with 614 additions and 48 deletions
+24 -11
View File
@@ -7,7 +7,7 @@ import asyncio
from dotenv import load_dotenv from dotenv import load_dotenv
from fastapi import FastAPI, Depends, Response, Header, HTTPException, status from fastapi import FastAPI, Depends, Response, Header, HTTPException, status
from guard import SecurityMiddleware, SecurityConfig from guard import SecurityMiddleware, SecurityConfig, SecurityDecorator
from typing import Optional from typing import Optional
@@ -33,16 +33,30 @@ app: FastAPI = FastAPI(
) )
config = SecurityConfig( config = SecurityConfig(
enable_rate_limiting=True, enable_rate_limiting=True,
rate_limit=120, rate_limit=10, # TODO: check rate limits in real usage
rate_limit_window=60, rate_limit_window=3, # TODO: check rate limits in real usage
enable_redis=False, enable_redis=False,
enable_ip_banning=True, enable_ip_banning=True,
auto_ban_threshold=3,
auto_ban_duration=3600,
custom_log_file="security.log", custom_log_file="security.log",
enable_penetration_detection=True,
auto_ban_threshold=3,
auto_ban_duration=3600,
detection_compiler_timeout=2.0,
detection_max_content_length=10000,
detection_preserve_attack_patterns=True,
detection_semantic_threshold=0.7,
detection_anomaly_threshold=3.0,
detection_slow_pattern_threshold=0.1,
detection_monitor_history_size=1000,
detection_max_tracked_patterns=1000,
) )
guard_deco = SecurityDecorator(config)
app.add_middleware(SecurityMiddleware, config=config) app.add_middleware(SecurityMiddleware, config=config)
app.state.guard_decorator = guard_deco
mongo_worker = MongoWorker() mongo_worker = MongoWorker()
_rabbit_worker: Optional[RabbitWorker] = None _rabbit_worker: Optional[RabbitWorker] = None
@@ -93,6 +107,7 @@ async def get_user(
@app.post("/add_user", status_code=201) @app.post("/add_user", status_code=201)
@guard_deco.rate_limit(requests=3, window=60)
async def add_user( async def add_user(
new_user: AddUserBody, new_user: AddUserBody,
response: Response, response: Response,
@@ -122,7 +137,9 @@ async def get_card(
return BaseResponse(result="There is no card with this card_id", error=True) return BaseResponse(result="There is no card with this card_id", error=True)
@app.get("/get_random_cards", status_code=200) @app.get("/get_random_cards", status_code=200)
@guard_deco.rate_limit(requests=5, window=60)
async def get_random_cards( async def get_random_cards(
user_id: int, user_id: int,
response: Response, response: Response,
@@ -132,8 +149,6 @@ async def get_random_cards(
if cards_visited.error: if cards_visited.error:
response.status_code = status.HTTP_404_NOT_FOUND response.status_code = status.HTTP_404_NOT_FOUND
return cards_visited return cards_visited
# Передаём exclude_ids напрямую в запрос — один round-trip к БД вместо цикла
exclude_ids = cards_visited.result.cards_visited or None exclude_ids = cards_visited.result.cards_visited or None
random_cards = await mongo.get_random_cards(10, True, exclude_ids=exclude_ids) random_cards = await mongo.get_random_cards(10, True, exclude_ids=exclude_ids)
@@ -145,14 +160,13 @@ async def get_random_cards(
@app.post("/add_card", status_code=201) @app.post("/add_card", status_code=201)
@guard_deco.rate_limit(requests=3, window=60)
async def add_card( async def add_card(
new_card: AddCardBody, new_card: AddCardBody,
response: Response, response: Response,
mongo: MongoWorker = Depends(lambda: mongo_worker),) -> BaseResponse: mongo: MongoWorker = Depends(lambda: mongo_worker),) -> BaseResponse:
if moderate_text(new_card.choice_A) and moderate_text(new_card.choice_B): if moderate_text(new_card.choice_A) and moderate_text(new_card.choice_B):
card = await mongo.add_card_by_api(new_card.choice_A, new_card.choice_B, new_card.author_id) card = await mongo.add_card_by_api(new_card.choice_A, new_card.choice_B, new_card.author_id)
# Отправляем карточку в RabbitMQ на ручную модерацию админом
try: try:
await get_rabbit_worker().send_to_moderation(card) await get_rabbit_worker().send_to_moderation(card)
except Exception as exc: except Exception as exc:
@@ -168,7 +182,6 @@ async def card_accept(
card_id: int, card_id: int,
response: Response, response: Response,
mongo: MongoWorker = Depends(lambda: mongo_worker),) -> BaseResponse: mongo: MongoWorker = Depends(lambda: mongo_worker),) -> BaseResponse:
"""Принимает карточку — доступ только с секретным ключом."""
result = await mongo.accept_card(card_id) result = await mongo.accept_card(card_id)
if result.error: if result.error:
response.status_code = status.HTTP_404_NOT_FOUND response.status_code = status.HTTP_404_NOT_FOUND
@@ -180,7 +193,6 @@ async def card_reject(
card_id: int, card_id: int,
response: Response, response: Response,
mongo: MongoWorker = Depends(lambda: mongo_worker),) -> BaseResponse: mongo: MongoWorker = Depends(lambda: mongo_worker),) -> BaseResponse:
"""Отклоняет карточку — доступ только с секретным ключом."""
result = await mongo.reject_card(card_id) result = await mongo.reject_card(card_id)
if result.error: if result.error:
response.status_code = status.HTTP_404_NOT_FOUND response.status_code = status.HTTP_404_NOT_FOUND
@@ -234,6 +246,7 @@ async def dislike_card(
@app.post("/comment", status_code=201) @app.post("/comment", status_code=201)
@guard_deco.rate_limit(requests=5, window=20)
async def comment( async def comment(
comment_info: AddCommentBody, comment_info: AddCommentBody,
response: Response, response: Response,
+5 -5
View File
@@ -34,7 +34,7 @@ class MongoWorker:
self.comments_data = self.db["comments"] self.comments_data = self.db["comments"]
async def create_indexes(self) -> None: async def create_indexes(self) -> None:
"""Создаёт индексы при старте приложения.""" """Creates indexes on application startup."""
await self.users_data.create_index("user_id", unique=True) await self.users_data.create_index("user_id", unique=True)
await self.game_data.create_index("card_id", unique=True) await self.game_data.create_index("card_id", unique=True)
await self.game_data.create_index("active_status") await self.game_data.create_index("active_status")
@@ -66,7 +66,7 @@ class MongoWorker:
async def get_and_update_counter(self, counter_name: str) -> int: async def get_and_update_counter(self, counter_name: str) -> int:
"""Атомарно инкрементирует счётчик и возвращает новое значение.""" """Atomically increments the counter and returns the new value."""
counter = await self.counters.find_one_and_update( counter = await self.counters.find_one_and_update(
{"counter_name": counter_name}, {"counter_name": counter_name},
{"$inc": {"counter": 1}}, {"$inc": {"counter": 1}},
@@ -101,7 +101,7 @@ class MongoWorker:
return None return None
async def get_random_cards(self, amount: int,active_status: bool,exclude_ids: Optional[set[int]] = None,) -> Optional[list[Card]]: async def get_random_cards(self, amount: int,active_status: bool,exclude_ids: Optional[set[int]] = None,) -> Optional[list[Card]]:
"""Возвращает случайные карточки, исключая уже просмотренные (одним запросом).""" """Returns random cards, excluding already visited ones (in a single query)."""
match_filter: dict = {"active_status": active_status} match_filter: dict = {"active_status": active_status}
if exclude_ids: if exclude_ids:
match_filter["card_id"] = {"$nin": list(exclude_ids)} match_filter["card_id"] = {"$nin": list(exclude_ids)}
@@ -142,7 +142,7 @@ class MongoWorker:
raise raise
async def accept_card(self, card_id: int) -> BaseResponse: async def accept_card(self, card_id: int) -> BaseResponse:
"""Принимает карточку: ставит active_status=True и moderation_date=сейчас.""" """Accepts a card: sets active_status=True and moderation_date=now."""
result = await self.game_data.find_one_and_update( result = await self.game_data.find_one_and_update(
{"card_id": card_id}, {"card_id": card_id},
{"$set": { {"$set": {
@@ -156,7 +156,7 @@ class MongoWorker:
return BaseResponse(result=Card.model_validate(result)) return BaseResponse(result=Card.model_validate(result))
async def reject_card(self, card_id: int) -> BaseResponse: async def reject_card(self, card_id: int) -> BaseResponse:
"""Отклоняет карточку: удаляет её из БД.""" """Rejects a card: deletes it from the database."""
result = await self.game_data.delete_one({"card_id": card_id}) result = await self.game_data.delete_one({"card_id": card_id})
if result.deleted_count == 0: if result.deleted_count == 0:
return BaseResponse(result="Card doesn't exist", error=True) return BaseResponse(result="Card doesn't exist", error=True)
+2 -4
View File
@@ -20,10 +20,9 @@ class RabbitWorker:
self.url = ( self.url = (
f"amqp://{os.getenv('RABBIT_USER')}:{os.getenv('RABBIT_PASS')}" f"amqp://{os.getenv('RABBIT_USER')}:{os.getenv('RABBIT_PASS')}"
f"@{os.getenv('RABBIT_HOST')}:{os.getenv('RABBIT_PORT')}" f"@{os.getenv('RABBIT_HOST')}:{os.getenv('RABBIT_PORT')}"
) )
async def send_to_moderation(self, card: Card) -> None: async def send_to_moderation(self, card: Card) -> None:
"""Отправляет карточку в очередь модерации."""
connection = await aio_pika.connect_robust(self.url) connection = await aio_pika.connect_robust(self.url)
async with connection: async with connection:
channel = await connection.channel() channel = await connection.channel()
@@ -41,7 +40,6 @@ class RabbitWorker:
self, self,
callback: Callable[[Card], Awaitable[None]], callback: Callable[[Card], Awaitable[None]],
) -> None: ) -> None:
"""Бесконечно слушает очередь модерации и вызывает callback для каждой карточки."""
connection = await aio_pika.connect_robust(self.url) connection = await aio_pika.connect_robust(self.url)
async with connection: async with connection:
channel = await connection.channel() channel = await connection.channel()
@@ -61,7 +59,7 @@ class RabbitWorker:
await queue.consume(on_message) await queue.consume(on_message)
# Держим consumer живым, но позволяем отмену (Ctrl+C) # Keep consumer alive while allowing cancellation (Ctrl+C)
stop_event = asyncio.Event() stop_event = asyncio.Event()
try: try:
await stop_event.wait() await stop_event.wait()
+48
View File
@@ -0,0 +1,48 @@
"""
conftest.py — fixtures for resetting rate-limiter state and IP-ban
between tests so functional tests do not hit 429 errors.
"""
import pytest
from guard import ip_ban_manager
from guard_core.handlers.ratelimit_handler import RateLimitManager
def _clear_middleware_suspicious_counts():
"""Search for SecurityMiddleware in middleware stack and clear suspicious_request_counts."""
from guard.middleware import SecurityMiddleware
from main import app
current = app
visited = set()
while current is not None and id(current) not in visited:
visited.add(id(current))
if isinstance(current, SecurityMiddleware):
current.suspicious_request_counts.clear()
break
current = getattr(current, 'app', None)
def _reset_all():
"""Full reset of rate-limiter, IP-ban, and suspicious counts."""
# Rate limit timestamps
rl: RateLimitManager | None = RateLimitManager._instance
if rl is not None:
rl.request_timestamps.clear()
# IP bans
ip_ban_manager.banned_ips.clear()
ip_ban_manager.banned_networks.clear()
# Suspicious request counts
_clear_middleware_suspicious_counts()
@pytest.fixture(autouse=True)
def reset_guard_state():
"""
Synchronous fixture (autouse) that resets guard middleware
state before and after each test.
"""
_reset_all()
yield
_reset_all()
+12 -2
View File
@@ -107,9 +107,14 @@ async def test_add_card_malformed_json():
@pytest.mark.asyncio(loop_scope="session") @pytest.mark.asyncio(loop_scope="session")
async def test_async_card_creation(): async def test_async_card_creation():
"""
Test parallel card creation.
Limit on /add_card — 3 requests/60s (decorator).
Send only 2 parallel requests to avoid exceeding the limit.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
tasks = [] tasks = []
num_cards = 8 num_cards = 2 # at most 3 (decorator limit), leaving a margin
for i in range(num_cards): for i in range(num_cards):
payload = { payload = {
"choice_A": f"Async Option A {i}", "choice_A": f"Async Option A {i}",
@@ -118,7 +123,7 @@ async def test_async_card_creation():
} }
tasks.append(client.post("/add_card", json=payload)) tasks.append(client.post("/add_card", json=payload))
responses = await asyncio.gather(*tasks) responses = await asyncio.gather(*tasks)
card_ids = [] card_ids = []
for idx, response in enumerate(responses): for idx, response in enumerate(responses):
print(f"\nAsync creation {idx}: status={response.status_code}, response={response.json()}") print(f"\nAsync creation {idx}: status={response.status_code}, response={response.json()}")
@@ -139,16 +144,21 @@ async def test_async_card_creation():
@pytest.mark.asyncio(loop_scope="session") @pytest.mark.asyncio(loop_scope="session")
async def test_get_card_valid(): async def test_get_card_valid():
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
# First create a card
payload = { payload = {
"choice_A": "GetTest A", "choice_A": "GetTest A",
"choice_B": "GetTest B", "choice_B": "GetTest B",
"author_id": EXIST_AUTHOR "author_id": EXIST_AUTHOR
} }
create_resp = await client.post("/add_card", json=payload) create_resp = await client.post("/add_card", json=payload)
assert create_resp.status_code in (200, 201), (
f"Failed to create card: {create_resp.status_code} {create_resp.text}"
)
base_create = BaseResponse.model_validate(create_resp.json()) base_create = BaseResponse.model_validate(create_resp.json())
card = Card.model_validate(base_create.result) card = Card.model_validate(base_create.result)
card_id = card.card_id card_id = card.card_id
# Now retrieve it
response = await client.get("/get_card", params={"card_id": card_id}) response = await client.get("/get_card", params={"card_id": card_id})
print(f"\nINPUT: endpoint=/get_card | params={{'card_id': {card_id}}}\nOUTPUT: status={response.status_code} | json={response.json()}") print(f"\nINPUT: endpoint=/get_card | params={{'card_id': {card_id}}}\nOUTPUT: status={response.status_code} | json={response.json()}")
assert response.status_code == 200 assert response.status_code == 200
+483
View File
@@ -0,0 +1,483 @@
"""
test_security.py — tests for checking rate limiting and penetration detection.
Rate limiting settings from main.py:
- Global: 10 requests / 3 sec (middleware)
- /add_user: 3 requests / 60 sec (decorator)
- /add_card: 3 requests / 60 sec (decorator)
- /get_random_cards: 5 requests / 60 sec (decorator)
- /comment: 5 requests / 20 sec (decorator)
Penetration detection:
- enable_penetration_detection=True
- auto_ban_threshold=3 (ban after 3 suspicious requests)
- auto_ban_duration=3600 (ban for 1 hour)
"""
import random
import asyncio
import pytest
from httpx import AsyncClient, ASGITransport
from main import app
# ========================================================================
# RATE LIMIT TESTS
# ========================================================================
class TestGlobalRateLimit:
"""Tests for global rate limit: 10 requests / 3 seconds."""
@pytest.mark.asyncio(loop_scope="session")
async def test_global_rate_limit_allows_under_threshold(self):
"""Requests within the limit (<=10) should pass."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
for i in range(9):
resp = await client.get("/check_user", params={"user_id": 1})
assert resp.status_code == 200, (
f"Request {i+1}/9 returned {resp.status_code}, expected 200: {resp.text}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_global_rate_limit_blocks_over_threshold(self):
"""
After exceeding global limit (10 requests/3s) -> 429.
/check_user does not have a rate_limit decorator, so only global limit applies.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
# Send 10 requests (fill the limit)
for i in range(10):
await client.get("/check_user", params={"user_id": 1})
# 11th request should return 429
resp = await client.get("/check_user", params={"user_id": 1})
assert resp.status_code == 429, (
f"Expected 429 after exceeding global limit, got {resp.status_code}"
)
assert "Too many requests" in resp.text
@pytest.mark.asyncio(loop_scope="session")
async def test_global_rate_limit_response_format(self):
"""Verify response format on rate limit."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
# Exhaust limit
for _ in range(10):
await client.get("/check_user", params={"user_id": 1})
resp = await client.get("/check_user", params={"user_id": 1})
assert resp.status_code == 429
assert resp.text == "Too many requests"
class TestDecoratorRateLimit:
"""Tests for rate limit via @guard_deco.rate_limit() decorator."""
@pytest.mark.asyncio(loop_scope="session")
async def test_add_user_rate_limit(self):
"""
/add_user: limit 3 requests / 60 sec.
First 3 requests pass (422 due to invalid data is OK, main point is not 429).
4th request -> 429.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
data = {
"user_id": random.randint(100000000, 999999999),
"username": "RateTest",
"first_name": "F",
"last_name": "L",
"photo_url": "http://test.test/photo.jpg"
}
# First 3 requests — not 429
for i in range(3):
resp = await client.post("/add_user", json=data)
assert resp.status_code != 429, (
f"Request {i+1}/3 returned 429, limit should not be exceeded yet"
)
# 4th request -> 429
resp = await client.post("/add_user", json=data)
assert resp.status_code == 429, (
f"Expected 429 after 3 requests to /add_user, got {resp.status_code}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_add_card_rate_limit(self):
"""
/add_card: limit 3 requests / 60 sec.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
author_id = random.randint(100000000, 999999999)
for i in range(3):
payload = {
"choice_A": f"Rate A {i}",
"choice_B": f"Rate B {i}",
"author_id": author_id
}
resp = await client.post("/add_card", json=payload)
assert resp.status_code != 429, (
f"Request {i+1}/3 to /add_card returned 429 prematurely"
)
payload = {
"choice_A": "Rate A overflow",
"choice_B": "Rate B overflow",
"author_id": author_id
}
resp = await client.post("/add_card", json=payload)
assert resp.status_code == 429, (
f"Expected 429 after 3 requests to /add_card, got {resp.status_code}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_get_random_cards_rate_limit(self):
"""
/get_random_cards: limit 5 requests / 60 sec.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
user_id = random.randint(100000000, 999999999)
for i in range(5):
resp = await client.get("/get_random_cards", params={"user_id": user_id})
# Can be 200 or 404 (if no cards/user), but not 429
assert resp.status_code != 429, (
f"Request {i+1}/5 to /get_random_cards returned 429 prematurely"
)
resp = await client.get("/get_random_cards", params={"user_id": user_id})
assert resp.status_code == 429, (
f"Expected 429 after 5 requests to /get_random_cards, got {resp.status_code}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_comment_rate_limit(self):
"""
/comment: limit 5 requests / 20 sec.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
for i in range(5):
payload = {
"author_id": random.randint(100000000, 999999999),
"card_id": 1,
"comment_text": f"Rate test comment {i}"
}
resp = await client.post("/comment", json=payload)
# Can be 201, 400 (moderation), 404 (card/user not found) — but not 429
assert resp.status_code != 429, (
f"Request {i+1}/5 to /comment returned 429 prematurely"
)
payload = {
"author_id": random.randint(100000000, 999999999),
"card_id": 1,
"comment_text": "Overflow comment"
}
resp = await client.post("/comment", json=payload)
assert resp.status_code == 429, (
f"Expected 429 after 5 requests to /comment, got {resp.status_code}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_different_endpoints_have_independent_limits(self):
"""
Decorator rate limit is tracked separately for each endpoint.
Requests to /check_user should not affect /add_card limit.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
# 5 requests to /check_user (no decorator, but global limit 10/3s)
for _ in range(5):
await client.get("/check_user", params={"user_id": 1})
# First request to /add_card — should pass (its own separate limit)
payload = {
"choice_A": "IndepA",
"choice_B": "IndepB",
"author_id": random.randint(100000000, 999999999)
}
resp = await client.post("/add_card", json=payload)
assert resp.status_code != 429, (
f"Request to /add_card blocked after requests to /check_user: {resp.status_code}"
)
class TestRateLimitParallel:
"""Rate limit tests with parallel requests."""
@pytest.mark.asyncio(loop_scope="session")
async def test_parallel_requests_hit_rate_limit(self):
"""
Multiple parallel requests should lead to 429 for some of them.
Send 15 parallel requests with global limit of 10/3s.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
tasks = [
client.get("/check_user", params={"user_id": 1})
for _ in range(15)
]
responses = await asyncio.gather(*tasks)
statuses = [r.status_code for r in responses]
count_200 = statuses.count(200)
count_429 = statuses.count(429)
print(f"\nParallel requests: 200={count_200}, 429={count_429}")
assert count_429 > 0, (
f"No request received 429 during 15 parallel requests: {statuses}"
)
assert count_200 > 0, (
f"All requests were blocked, none passed: {statuses}"
)
# ========================================================================
# PENETRATION DETECTION TESTS
# ========================================================================
class TestPenetrationDetection:
"""
Tests for malicious request detection.
enable_penetration_detection=True
auto_ban_threshold=3
auto_ban_duration=3600
"""
@pytest.mark.asyncio(loop_scope="session")
async def test_sql_injection_detected(self):
"""SQL injection in query parameters should be detected."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
resp = await client.get(
"/get_card",
params={"card_id": "1 OR 1=1; DROP TABLE users;--"}
)
print(f"\nSQL injection test: status={resp.status_code} | text={resp.text[:200]}")
# Expect: 400 (suspicious activity) or 422 (validation) — but NOT 200
assert resp.status_code in (400, 403, 422), (
f"SQL injection was not blocked, got {resp.status_code}: {resp.text[:200]}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_xss_in_query_params_detected(self):
"""XSS attack in query parameters should be detected."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
resp = await client.get(
"/get_card",
params={"card_id": "<script>alert('XSS')</script>"}
)
print(f"\nXSS in params test: status={resp.status_code} | text={resp.text[:200]}")
assert resp.status_code in (400, 403, 422), (
f"XSS attack was not detected, got {resp.status_code}: {resp.text[:200]}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_path_traversal_detected(self):
"""Path traversal attack should be detected."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
resp = await client.get("/get_card/../../../etc/passwd")
print(f"\nPath traversal test: status={resp.status_code} | text={resp.text[:200]}")
# Can be 400, 403, 404, or 422 — but MUST NOT expose file contents
assert resp.status_code != 200 or "root:" not in resp.text, (
"Path traversal not detected — system file accessed!"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_xss_in_post_body_detected(self):
"""XSS attack in POST body should be detected."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
payload = {
"choice_A": "<script>document.cookie</script>",
"choice_B": "Normal option",
"author_id": random.randint(100000000, 999999999)
}
resp = await client.post("/add_card", json=payload)
print(f"\nXSS in body test: status={resp.status_code} | text={resp.text[:200]}")
# 400 (suspicious), 403 (banned), or 422 — but not 201
assert resp.status_code in (400, 403, 422), (
f"XSS in request body was not detected, got {resp.status_code}: {resp.text[:200]}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_command_injection_detected(self):
"""Command injection attempt should be detected."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
payload = {
"choice_A": "; cat /etc/passwd; echo",
"choice_B": "$(whoami)",
"author_id": random.randint(100000000, 999999999)
}
resp = await client.post("/add_card", json=payload)
print(f"\nCommand injection test: status={resp.status_code} | text={resp.text[:200]}")
# 400 (suspicious), 403 (banned) — not 201
assert resp.status_code in (400, 403, 422), (
f"Command injection was not detected, got {resp.status_code}: {resp.text[:200]}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_sql_union_injection_detected(self):
"""UNION-based SQL injection should be detected."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
resp = await client.get(
"/get_card",
params={"card_id": "1 UNION SELECT password FROM users"}
)
print(f"\nUNION SQL injection test: status={resp.status_code} | text={resp.text[:200]}")
assert resp.status_code in (400, 403, 422), (
f"UNION SQL injection was not blocked, got {resp.status_code}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_legitimate_request_not_blocked(self):
"""Legitimate request with normal data should not be blocked as suspicious."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
resp = await client.get("/get_card", params={"card_id": 1})
print(f"\nLegitimate request test: status={resp.status_code}")
# 200 (card found) or 404 (not found) — but not 400/403
assert resp.status_code in (200, 404), (
f"Legitimate request blocked: {resp.status_code}: {resp.text[:200]}"
)
class TestAutoIPBan:
"""
Tests for automatic IP banning after repeated suspicious requests.
auto_ban_threshold=3, auto_ban_duration=3600
"""
@pytest.mark.asyncio(loop_scope="session")
async def test_repeated_attacks_trigger_ip_ban(self):
"""
After auto_ban_threshold (3) suspicious requests, the IP should be banned.
Subsequent requests (even legitimate ones) should return 403.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
# Send suspicious requests sequentially (SQL injection variants)
injection_payloads = [
"1' OR '1'='1",
"1; DROP TABLE cards;--",
"1 UNION SELECT * FROM users;--",
"1' AND 1=CONVERT(int,(SELECT TOP 1 name FROM sysobjects));--",
]
detected_as_suspicious = 0
for payload in injection_payloads:
resp = await client.get("/get_card", params={"card_id": payload})
if resp.status_code in (400, 403):
detected_as_suspicious += 1
print(f" Attack attempt: status={resp.status_code} | payload={payload[:50]}")
print(f"\nSuspicious requests detected: {detected_as_suspicious}/{len(injection_payloads)}")
if detected_as_suspicious >= 3:
# Threshold reached — verify ban on legitimate request
resp = await client.get("/check_user", params={"user_id": 1})
print(f"Post-attack legitimate request: status={resp.status_code}")
assert resp.status_code == 403, (
f"IP should be banned after {detected_as_suspicious} suspicious requests, "
f"but legitimate request returned {resp.status_code}: {resp.text[:200]}"
)
else:
pytest.skip(
f"Only {detected_as_suspicious} of {len(injection_payloads)} attacks detected, "
f"ban threshold (3) not reached"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_banned_ip_returns_403_on_all_endpoints(self):
"""
If IP is banned, all endpoints should return 403.
"""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
# Send various attacks to guarantee hitting the threshold
attacks = [
"1' OR '1'='1; --",
"1; DROP TABLE cards; --",
"<script>alert(1)</script>",
"../../etc/shadow",
"1 UNION SELECT password FROM users",
]
detected_count = 0
for payload in attacks:
resp = await client.get("/get_card", params={"card_id": payload})
if resp.status_code in (400, 403):
detected_count += 1
print(f" [{payload[:40]}] status={resp.status_code}")
print(f"\nDetected: {detected_count}/{len(attacks)}")
if detected_count >= 3:
# Check ban on different endpoints
endpoints = [
("GET", "/check_user", {"user_id": 99999}),
("GET", "/get_user", {"user_id": 99999}),
("GET", "/get_card", {"card_id": 1}),
]
for method, path, params in endpoints:
resp = await client.get(path, params=params)
assert resp.status_code == 403, (
f"IP is banned, but {method} {path} returned {resp.status_code}"
)
else:
pytest.skip(
f"Only {detected_count} attacks detected, ban threshold (3) not reached"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_banned_ip_message(self):
"""Banned IP should receive 'IP address banned' message."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
attacks = [
"1' OR '1'='1; --",
"1; DROP TABLE cards; --",
"1 UNION SELECT password FROM users",
"<script>alert(1)</script>",
]
detected = 0
for payload in attacks:
resp = await client.get("/get_card", params={"card_id": payload})
if resp.status_code in (400, 403):
detected += 1
if detected >= 3:
resp = await client.get("/check_user", params={"user_id": 1})
assert resp.status_code == 403
assert "IP address banned" in resp.text, (
f"Expected message 'IP address banned', got: {resp.text[:200]}"
)
else:
pytest.skip(f"Only {detected} attacks detected, threshold not reached")
class TestSuspiciousHeaders:
"""Tests for suspicious header detection."""
@pytest.mark.asyncio(loop_scope="session")
async def test_suspicious_user_agent(self):
"""Request with suspicious User-Agent may be blocked."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
resp = await client.get(
"/check_user",
params={"user_id": 1},
headers={"User-Agent": "sqlmap/1.6.12#stable (http://sqlmap.org)"}
)
print(f"\nSuspicious UA test: status={resp.status_code}")
# sqlmap is a known SQL injection tool
# Expect block (403) or pass (200 — if UA is not in blocklist)
assert resp.status_code in (200, 400, 403), (
f"Unexpected status code for suspicious User-Agent: {resp.status_code}"
)
@pytest.mark.asyncio(loop_scope="session")
async def test_xss_in_headers(self):
"""XSS attack via custom headers."""
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
resp = await client.get(
"/check_user",
params={"user_id": 1},
headers={"X-Forwarded-For": "<script>alert(1)</script>"}
)
print(f"\nXSS in headers test: status={resp.status_code}")
# Header may be ignored or detected as suspicious
assert resp.status_code in (200, 400, 403), (
f"Unexpected status code for XSS in headers: {resp.status_code}"
)
+26 -12
View File
@@ -63,30 +63,44 @@ async def test_get_random_cards_valid():
@pytest.mark.asyncio(loop_scope="session") @pytest.mark.asyncio(loop_scope="session")
async def test_get_random_cards_randomness(): async def test_get_random_cards_randomness():
if NO_ACTIVE_CARDS_STATUS:
pytest.skip(reason="No active cards in MongoDB")
elif ACTIVE_CARDS_LESS_THAN_TEN:
pytest.skip(reason="The number of active cards is less than 10 in MongoDB")
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
params = {"user_id": EXIST_USER} params = {"user_id": EXIST_USER}
response1 = await client.get("/get_random_cards", params=params) response1 = await client.get("/get_random_cards", params=params)
response2 = await client.get("/get_random_cards", params=params) response2 = await client.get("/get_random_cards", params=params)
print(f"\nRandomness check: r1={response1.status_code}, r2={response2.status_code}")
assert response1.status_code == 200, f"First request returned {response1.status_code}: {response1.text}"
assert response2.status_code == 200, f"Second request returned {response2.status_code}: {response2.text}"
result1 = response1.json().get("result") result1 = response1.json().get("result")
result2 = response2.json().get("result") result2 = response2.json().get("result")
print(f"\nINPUT: endpoint=/get_random_cards (двойной вызов)\nOUTPUT 1: {result1}\nOUTPUT 2: {result2}") print(f"\nINPUT: endpoint=/get_random_cards (double call)\nOUTPUT 1: {result1}\nOUTPUT 2: {result2}")
if len(result1) == 10 and len(result2) == 10: if len(result1) == 10 and len(result2) == 10:
assert result1 != result2 assert result1 != result2
@pytest.mark.asyncio(loop_scope="session") @pytest.mark.asyncio(loop_scope="session")
async def test_get_random_cards_parallel_requests(): async def test_get_random_cards_parallel_requests():
"""
Parallel requests to /get_random_cards.
Decorator limit: 5 requests/60s.
Make 3 parallel requests to stay within limit.
"""
if NO_ACTIVE_CARDS_STATUS: if NO_ACTIVE_CARDS_STATUS:
pytest.skip(reason="No active cards in MongoDB") pytest.skip(reason="No active cards in MongoDB")
elif ACTIVE_CARDS_LESS_THAN_TEN: elif ACTIVE_CARDS_LESS_THAN_TEN:
pytest.skip(reason="The number of active cards is less than 10 in MongoDB") pytest.skip(reason="The number of active cards is less than 10 in MongoDB")
else:
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
params = {"user_id": EXIST_USER} params = {"user_id": EXIST_USER}
tasks = [client.get("/get_random_cards", params=params) for _ in range(5)] tasks = [client.get("/get_random_cards", params=params) for _ in range(3)]
responses = await asyncio.gather(*tasks) responses = await asyncio.gather(*tasks)
for resp in responses: for resp in responses:
print(f"\nParallel call: status={resp.status_code} | json={resp.json()}") print(f"\nParallel call: status={resp.status_code} | text={resp.text[:200]}")
assert resp.status_code == 200 assert resp.status_code == 200, (
result = resp.json().get("result") f"Expected 200, got {resp.status_code}: {resp.text}"
assert isinstance(result, list) )
assert len(result) == 10 result = resp.json().get("result")
assert isinstance(result, list)
+14 -14
View File
@@ -1,11 +1,11 @@
""" """
Telegram-бот модерации карточек. Telegram card moderation bot.
Слушает очередь RabbitMQ «moderation» и отправляет карточки Listens to the RabbitMQ "moderation" queue and sends cards
в чат администратору с inline-кнопками «Принять ✅» / «Отклонить ❌». to the admin chat with inline buttons "Accept ✅" / "Reject ❌".
При нажатии кнопки бот вызывает защищённые эндпоинты When a button is pressed, the bot calls protected endpoints
/card_accept или /card_reject с секретным заголовком. /card_accept or /card_reject with a secret header.
""" """
import os import os
@@ -27,7 +27,7 @@ load_dotenv()
logging.basicConfig(level=logging.INFO) logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# ── Конфигурация ────────────────────────────────────────────── # ── Configuration ──────────────────────────────────────────────
BOT_TOKEN = os.getenv("TG_BOT_TOKEN") BOT_TOKEN = os.getenv("TG_BOT_TOKEN")
ADMIN_CHAT_ID = int(os.getenv("TG_ADMIN_CHAT_ID", "0")) ADMIN_CHAT_ID = int(os.getenv("TG_ADMIN_CHAT_ID", "0"))
API_BASE_URL = os.getenv("API_BASE_URL", "http://localhost:5000") API_BASE_URL = os.getenv("API_BASE_URL", "http://localhost:5000")
@@ -38,9 +38,9 @@ dp = Dispatcher()
rabbit = RabbitWorker() rabbit = RabbitWorker()
# ── Отправка карточки администратору ────────────────────────── # ── Sending card to admin ──────────────────────────
async def _get_author_username(author_id: int) -> str: async def _get_author_username(author_id: int) -> str:
"""Запрашивает username автора через API.""" """Fetches the author's username via API."""
try: try:
async with aiohttp.ClientSession() as session: async with aiohttp.ClientSession() as session:
async with session.get( async with session.get(
@@ -57,7 +57,7 @@ async def _get_author_username(author_id: int) -> str:
def _format_date(iso_date: str) -> str: def _format_date(iso_date: str) -> str:
"""Преобразует ISO-дату в формат ДД.ММ.ГГГГ ЧЧ:ММ:СС.""" """Converts ISO date to DD.MM.YYYY HH:MM:SS format."""
try: try:
dt = datetime.fromisoformat(iso_date) dt = datetime.fromisoformat(iso_date)
return dt.strftime("%d.%m.%Y %H:%M:%S") return dt.strftime("%d.%m.%Y %H:%M:%S")
@@ -66,7 +66,7 @@ def _format_date(iso_date: str) -> str:
async def send_card_to_admin(card: Card) -> None: async def send_card_to_admin(card: Card) -> None:
"""Формирует сообщение и inline-клавиатуру для карточки.""" """Formats message and inline keyboard for a card."""
author_display = await _get_author_username(card.author_id) author_display = await _get_author_username(card.author_id)
date_display = _format_date(card.creation_date) date_display = _format_date(card.creation_date)
@@ -100,10 +100,10 @@ async def send_card_to_admin(card: Card) -> None:
logger.info("Sent card %s to admin chat", card.card_id) logger.info("Sent card %s to admin chat", card.card_id)
# ── Вызов защищённых эндпоинтов API ────────────────────────── # ── Calling protected API endpoints ──────────────────────────
async def call_moderation_api(action: str, card_id: int) -> dict: async def call_moderation_api(action: str, card_id: int) -> dict:
""" """
Вызывает /card_accept или /card_reject с секретным заголовком. Calls /card_accept or /card_reject with secret header.
action: 'accept' | 'reject' action: 'accept' | 'reject'
""" """
endpoint = f"{API_BASE_URL}/card_{action}" endpoint = f"{API_BASE_URL}/card_{action}"
@@ -116,7 +116,7 @@ async def call_moderation_api(action: str, card_id: int) -> dict:
return data return data
# ── Обработчики callback-кнопок ─────────────────────────────── # ── Callback button handlers ───────────────────────────────
@dp.callback_query(F.data.startswith("accept:")) @dp.callback_query(F.data.startswith("accept:"))
async def on_accept(callback: CallbackQuery) -> None: async def on_accept(callback: CallbackQuery) -> None:
card_id = int(callback.data.split(":")[1]) card_id = int(callback.data.split(":")[1])
@@ -150,7 +150,7 @@ async def on_reject(callback: CallbackQuery) -> None:
await callback.answer("Карточка отклонена!") await callback.answer("Карточка отклонена!")
logger.info("Card %s rejected by admin", card_id) logger.info("Card %s rejected by admin", card_id)
# ── Lifecycle-хуки aiogram ──────────────────────────────────── # ── aiogram Lifecycle hooks ────────────────────────────────────
_rabbit_task: asyncio.Task | None = None _rabbit_task: asyncio.Task | None = None