Skip to content
Merged
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
74 changes: 70 additions & 4 deletions bot/application/command_handlers/qr/generate_qr_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,21 @@ def __init__(
self.dynamic_setup = dynamic_setup


class _LimitReservation:
"""Tracks pre-allocated atomic limiter counters for rollback."""

__slots__ = ("keys",)

def __init__(self) -> None:
self.keys: list[str] = []

def add(self, key: str) -> None:
self.keys.append(key)

def is_empty(self) -> bool:
return len(self.keys) == 0


class GenerateQRHandler(ICommandHandler[GenerateQRCommand, tuple[bytes, int | None]]):
"""
Command handler for QR code generation.
Expand Down Expand Up @@ -146,14 +161,16 @@ async def handle(
RuntimeError: If unique short code cannot be generated
"""
start_time = time.time()
limit_reservation: _LimitReservation | None = None
persisted_successfully = False

try:
# ============================================
# PHASE 0: Enforce generation limits (SECURITY)
# ============================================
# Limits MUST be checked here (not only in QRGenerationService)
# to prevent bypass via direct mediator.send(GenerateQRCommand)
await self._enforce_limits(command.user_id)
limit_reservation = await self._enforce_limits(command.user_id)

# ============================================
# PHASE 1: Validation and Preparation (no I/O)
Expand All @@ -172,6 +189,7 @@ async def handle(
qr_data, qr_code_id, pending_events = await self._persist_with_retry(
command, content, qr_data, generation_time_ms
)
persisted_successfully = True

# ============================================
# PHASE 4: Side effects (AFTER successful commit)
Expand All @@ -194,11 +212,23 @@ async def handle(
return qr_data, qr_code_id

except ValueError:
if (
limit_reservation
and not persisted_successfully
and not limit_reservation.is_empty()
):
await self._rollback_limit_reservation(limit_reservation)
logger.error(
f"❌ Validation error in GenerateQRCommand for user {command.user_id}"
)
raise
except Exception:
if (
limit_reservation
and not persisted_successfully
and not limit_reservation.is_empty()
):
await self._rollback_limit_reservation(limit_reservation)
logger.error(
f"❌ Error handling GenerateQRCommand for user {command.user_id}"
)
Expand All @@ -208,7 +238,7 @@ async def handle(
# Private: Phase 0 — Limit Enforcement (SECURITY)
# ------------------------------------------------------------------

async def _enforce_limits(self, user_id: int) -> None:
async def _enforce_limits(self, user_id: int) -> _LimitReservation:
"""
Enforce QR generation limits at the handler level.

Expand All @@ -225,6 +255,8 @@ async def _enforce_limits(self, user_id: int) -> None:
PermissionError: If user is banned or limits exceeded
ValueError: If user not found
"""
reservation = _LimitReservation()

user = await self._user_repo.get_by_telegram_id(user_id)
if not user:
raise ValueError(f"User {user_id} not found")
Expand Down Expand Up @@ -255,7 +287,13 @@ async def _enforce_limits(self, user_id: int) -> None:
daily_key,
ttl=86400, # Expire at end of day (24h max)
)
if daily_count > 0 and daily_count > tariff.daily_limit:
if daily_count <= 0:
await self._rollback_limit_reservation(reservation)
raise PermissionError("Rate limiter is unavailable. Try again shortly.")

reservation.add(daily_key)
if daily_count > tariff.daily_limit:
await self._rollback_limit_reservation(reservation)
raise PermissionError(
f"Daily QR generation limit exceeded ({tariff.daily_limit})"
)
Expand All @@ -267,7 +305,13 @@ async def _enforce_limits(self, user_id: int) -> None:
monthly_key,
ttl=86400 * 31, # ~1 month
)
if monthly_count > 0 and monthly_count > tariff.monthly_limit:
if monthly_count <= 0:
await self._rollback_limit_reservation(reservation)
raise PermissionError("Rate limiter is unavailable. Try again shortly.")

reservation.add(monthly_key)
if monthly_count > tariff.monthly_limit:
await self._rollback_limit_reservation(reservation)
raise PermissionError(
f"Monthly QR generation limit exceeded ({tariff.monthly_limit})"
)
Expand All @@ -292,6 +336,24 @@ async def _enforce_limits(self, user_id: int) -> None:
f"Monthly QR generation limit exceeded ({tariff.monthly_limit})"
)

return reservation

async def _rollback_limit_reservation(self, reservation: _LimitReservation) -> None:
"""Rollback pre-allocated limiter counters.

Called when generation fails after successful pre-allocation.
"""
if not self._cache or reservation.is_empty():
return

for key in reversed(reservation.keys):
try:
await self._cache.decrement(key)
except Exception:
logger.exception(
"Failed to rollback limiter reservation for key '%s'", key
)

# ------------------------------------------------------------------
# Private: Phase 1 — Preparation (pure logic, no I/O)
# ------------------------------------------------------------------
Expand Down Expand Up @@ -522,6 +584,10 @@ async def _do_persist(
event = event.with_updated(qr_code_id=qr_code_id)
pending_events.append(event)

# Keep daily_usage in sync with successful generation inside the same
# transaction so limit checks have a consistent source of truth.
await uow.qr_codes.increment_daily_usage(command.user_id, date.today())

# Increment user's QR count (Rich Domain Model)
user = await uow.users.get_by_telegram_id(command.user_id)
if user:
Expand Down
17 changes: 17 additions & 0 deletions bot/application/ports/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,3 +101,20 @@ async def increment(self, key: str, ttl: int | None = None) -> int:
Returns:
New counter value after increment
"""

@abstractmethod
async def decrement(self, key: str) -> int:
"""Decrement a counter and return the new value.

Used to rollback quota reservations if generation fails after
a successful pre-allocation.

Implementations must be resilient: if key is missing or operation
fails, they should return 0 and avoid negative counters.

Args:
key: Counter key

Returns:
New counter value after decrement, or 0 if key is absent/failure
"""
35 changes: 26 additions & 9 deletions bot/application/services/qr_generation_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -241,9 +241,7 @@ async def generate(

except Exception as e:
logger.error(f"QR generation error: {e}", exc_info=True)
return GenerationResult.internal_error(
"Service temporarily unavailable. Please try again."
)
return GenerationResult.internal_error(str(e))

async def _check_limits(
self,
Expand Down Expand Up @@ -352,20 +350,39 @@ async def _get_remaining_quota(
limit_query = CheckUserLimitQuery(user_id=user_id, is_admin=is_admin)
limit_check = await self._mediator.send(limit_query)

daily_used = self._extract_limit_usage(
limit_check, "daily_used", "used_today"
)
monthly_used = self._extract_limit_usage(
limit_check, "monthly_used", "used_this_month"
)

remaining_daily = None
if limit_check.daily_limit is not None:
remaining_daily = max(
0, limit_check.daily_limit - limit_check.daily_used - 1
)
remaining_daily = max(0, limit_check.daily_limit - daily_used - 1)

remaining_monthly = None
if limit_check.monthly_limit is not None:
remaining_monthly = max(
0, limit_check.monthly_limit - limit_check.monthly_used - 1
)
remaining_monthly = max(0, limit_check.monthly_limit - monthly_used - 1)

return remaining_daily, remaining_monthly

except Exception as e:
logger.warning(f"Failed to get remaining quota: {e}")
return None, None

@staticmethod
def _extract_limit_usage(
limit_check,
primary_attr: str,
fallback_attr: str,
) -> int:
"""Extract integer usage counter from limit check object."""
value = getattr(limit_check, primary_attr, None)
if not isinstance(value, (int, float)):
value = getattr(limit_check, fallback_attr, 0)

try:
return int(value)
except (TypeError, ValueError):
return 0
12 changes: 12 additions & 0 deletions bot/di/modules/enhanced_cache_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,10 +124,22 @@ async def delete(self, key: str) -> bool:
"""Always return False (nothing to delete)."""
return False

async def invalidate(self, key: str) -> bool:
"""Always return False (nothing to invalidate)."""
return False

async def invalidate_pattern(self, pattern: str) -> int:
"""Always return 0 (nothing to invalidate)."""
return 0

async def increment(self, key: str, ttl: int | None = None) -> int:
"""Always return 0 (no counters)."""
return 0

async def decrement(self, key: str) -> int:
"""Always return 0 (no counters)."""
return 0


class CacheHealthChecker:
"""
Expand Down
6 changes: 6 additions & 0 deletions bot/domain/repositories/qr_history_repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,12 @@ async def get_user_history(
async def get_daily_count(self, user_id: int, date: date) -> int:
"""Get count of QR codes generated by user on specific date"""

@abstractmethod
async def increment_daily_usage(
self, user_id: int, usage_date: date, increment: int = 1
) -> int:
"""Atomically increment daily usage counter and return new value"""

@abstractmethod
async def get_monthly_count(self, user_id: int, year: int, month: int) -> int:
"""Get count of QR codes generated by user in specific month"""
Expand Down
19 changes: 19 additions & 0 deletions bot/infrastructure/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -239,6 +239,25 @@ async def increment(self, key: str, ttl: int | None = None) -> int:
logger.error(f"❌ Cache increment error for key '{key}': {e}")
return 0

async def decrement(self, key: str) -> int:
"""Decrement a counter.

Returns:
New counter value after decrement (0 if missing/failure)
"""
if not self._enabled or not self._redis:
return 0

try:
value = await self._redis.decr(key)
if value <= 0:
await self._redis.delete(key)
return 0
return value
except Exception as e:
logger.error(f"❌ Cache decrement error for key '{key}': {e}")
return 0

async def ping(self) -> bool:
"""
Check if Redis connection is alive.
Expand Down
30 changes: 30 additions & 0 deletions bot/infrastructure/cache/enhanced_cache_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -532,6 +532,36 @@ async def increment(self, key: str, ttl: int | None = None) -> int:
logger.warning(f"Cache increment failed for key '{key}': {e}")
return 0

async def decrement(self, key: str) -> int:
"""Decrement a counter and return the new value.

Used to rollback quota reservations when QR generation fails after
successful pre-allocation. Never returns a negative value.

Args:
key: Counter key

Returns:
New counter value, or 0 if key is absent/failure
"""
try:
if self._is_circuit_open():
return 0

value = await self._redis.decr(key)
if value <= 0:
# Avoid negative counters (e.g., DECR on missing key => -1)
await self._redis.delete(key)
return 0

self._metrics.total_operations += 1
return value

except Exception as e:
self._record_failure(e)
logger.warning(f"Cache decrement failed for key '{key}': {e}")
return 0

async def _get_internal(self, key: str) -> Any | None:
"""
Internal get method with decompression and deserialization.
Expand Down
27 changes: 27 additions & 0 deletions bot/infrastructure/repositories/qr_history_repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,33 @@ async def get_daily_count(self, user_id: int, target_date: date) -> int:

return row["count"] if row else 0

async def increment_daily_usage(
self, user_id: int, usage_date: date, increment: int = 1
) -> int:
"""Atomically increment daily usage and return updated count."""
if increment <= 0:
raise ValueError("increment must be positive")

query = """
WITH seq AS (
SELECT COALESCE(
pg_get_serial_sequence('daily_usage', 'id'),
'daily_usage_id_seq'
) AS seq_name
)
INSERT INTO daily_usage (id, user_id, usage_date, qr_count)
VALUES (nextval((SELECT seq_name FROM seq)::regclass), $1, $2, $3)
ON CONFLICT (user_id, usage_date)
DO UPDATE SET
qr_count = daily_usage.qr_count + EXCLUDED.qr_count,
updated_at = CURRENT_TIMESTAMP
RETURNING qr_count
"""

async with self._pool.acquire() as conn:
value = await conn.fetchval(query, user_id, usage_date, increment)
return int(value) if value is not None else 0

async def get_monthly_count(self, user_id: int, year: int, month: int) -> int:
"""
Get number of QR codes generated by user in specific month.
Expand Down
3 changes: 0 additions & 3 deletions bot/presentation/handlers/qr/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@
├── input_handlers.py # URL, Text, Email input handlers
├── preview_handlers.py # Preview screen and customization handlers
├── format_selection.py # Format choice and QR generation
├── url_fix.py # URL suggestion fix handler
└── shared.py # Shared utilities and helpers

Each module has its own router that is combined in __init__.py.
Expand All @@ -28,7 +27,6 @@
from .input_handlers import router as input_router
from .payment_flow import router as payment_router
from .preview_handlers import router as preview_router
from .url_fix import router as url_fix_router
from .vcard_flow import router as vcard_router
from .wifi_flow import router as wifi_router

Expand All @@ -48,6 +46,5 @@
router.include_router(frame_router) # Frame with text customization handlers
router.include_router(preview_router) # Preview and customization
router.include_router(format_router) # Format selection and generation
router.include_router(url_fix_router) # URL fix handler

__all__ = ["router"]
Loading