diff --git a/app.py b/app.py index 058ffd5..897999b 100644 --- a/app.py +++ b/app.py @@ -1488,6 +1488,8 @@ class AddUserRequest(BaseModel): protocol: Optional[str] = None connection_name: Optional[str] = None expiration_date: Optional[str] = None + expire_after_first_use: bool = False + expiration_days: int = 0 telemt_quota: Optional[str] = None telemt_max_ips: Optional[int] = None telemt_expiry: Optional[str] = None @@ -1564,6 +1566,8 @@ class UpdateUserRequest(BaseModel): traffic_limit: Optional[float] = 0 traffic_reset_strategy: Optional[str] = None expiration_date: Optional[str] = None + expire_after_first_use: Optional[bool] = None + expiration_days: Optional[int] = None password: Optional[str] = None @@ -1703,6 +1707,10 @@ async def startup(): migrated = True if 'expiration_date' not in u: u['expiration_date'] = None + if 'expire_after_first_use' not in u: + u['expire_after_first_use'] = False + if 'expiration_days' not in u: + u['expiration_days'] = 0 migrated = True if 'xui_email' not in u: u['xui_email'] = None @@ -3346,6 +3354,8 @@ async def api_list_users(request: Request, search: str = '', page: int = 1, size 'traffic_reset_strategy': u.get('traffic_reset_strategy', 'never'), 'last_reset_at': u.get('last_reset_at'), "expiration_date": u.get("expiration_date"), + "expire_after_first_use": bool(u.get("expire_after_first_use")), + "expiration_days": int(u.get("expiration_days") or 0), 'share_enabled': u.get('share_enabled', False), 'share_token': u.get('share_token'), 'has_share_password': bool(u.get('share_password_hash')), @@ -3389,7 +3399,9 @@ async def api_add_user(request: Request, req: AddUserRequest): 'traffic_used': 0, 'traffic_total': 0, 'last_reset_at': datetime.now().isoformat(), - 'expiration_date': req.expiration_date, + 'expiration_date': None if req.expire_after_first_use else req.expiration_date, + 'expire_after_first_use': bool(req.expire_after_first_use), + 'expiration_days': int(req.expiration_days or 0) if req.expire_after_first_use else 0, 'enabled': True, 'created_at': datetime.now().isoformat(), 'remnawave_uuid': None, @@ -3398,6 +3410,8 @@ async def api_add_user(request: Request, req: AddUserRequest): 'share_token': secrets.token_urlsafe(16), 'share_password_hash': None, } + if new_user['expire_after_first_use'] and new_user['expiration_days'] <= 0: + return JSONResponse({'error': 'expiration_days must be > 0 when start-after-first-use is enabled'}, status_code=400) data['users'].append(new_user) save_data(data) @@ -3472,8 +3486,29 @@ async def api_update_user(request: Request, user_id: str, req: UpdateUserRequest user['last_reset_at'] = datetime.now().isoformat() req_fields = getattr(req, 'model_fields_set', getattr(req, '__fields_set__', set())) - if 'expiration_date' in req_fields: + if 'expire_after_first_use' in req_fields or 'expiration_days' in req_fields or 'expiration_date' in req_fields: + after_first = bool(req.expire_after_first_use) if req.expire_after_first_use is not None else bool(user.get('expire_after_first_use')) + days = int(req.expiration_days) if req.expiration_days is not None else int(user.get('expiration_days') or 0) + if after_first: + if days <= 0: + return JSONResponse({'error': 'expiration_days must be > 0 when start-after-first-use is enabled'}, status_code=400) + user['expire_after_first_use'] = True + user['expiration_days'] = days + # Keep existing absolute date if countdown already started; otherwise wait for first use + if not user.get('expiration_date'): + user['expiration_date'] = None + elif 'expiration_date' in req_fields and req.expiration_date: + # Admin can still override absolute end date after start + user['expiration_date'] = req.expiration_date or None + else: + user['expire_after_first_use'] = False + user['expiration_days'] = 0 + if 'expiration_date' in req_fields: + user['expiration_date'] = req.expiration_date or None + elif 'expiration_date' in req_fields: user['expiration_date'] = req.expiration_date or None + user['expire_after_first_use'] = False + user['expiration_days'] = 0 if req.password: user['password_hash'] = hash_password(req.password) @@ -3735,6 +3770,12 @@ async def api_share_config(token: str, connection_id: str, request: Request): return JSONResponse({'error': 'Not found'}, status_code=404) try: + from managers.user_expiration import maybe_start_user_expiration, user_is_expired + if user_is_expired(user): + return JSONResponse({'error': 'Subscription expired'}, status_code=403) + if maybe_start_user_expiration(data, user['id']): + save_data(data) + if protocol_base(conn.get('protocol', '')) == 'xui': return JSONResponse({'error': '3x-ui support removed'}, status_code=410) @@ -3749,7 +3790,7 @@ async def api_share_config(token: str, connection_id: str, request: Request): config = _manager_call(manager, 'get_client_config', conn['protocol'], conn['client_id'], server['host'], port) ssh.disconnect() vpn_link = generate_vpn_link(config) if config else '' - return {'config': config, 'vpn_link': vpn_link} + return {'config': config, 'vpn_link': vpn_link, 'expires_at': user.get('expiration_date')} except Exception as e: logger.exception("Error getting shared config") return JSONResponse({'error': str(e)}, status_code=500) @@ -3873,6 +3914,12 @@ async def api_guest_config(token: str, connection_id: str, request: Request): if not conn: return JSONResponse({'error': 'Not found'}, status_code=404) try: + from managers.user_expiration import maybe_start_user_expiration, user_is_expired + if user_is_expired(holder): + return JSONResponse({'error': 'Subscription expired'}, status_code=403) + if maybe_start_user_expiration(data, holder['id']): + save_data(data) + if protocol_base(conn.get('protocol', '')) == 'xui': return JSONResponse({'error': '3x-ui support removed'}, status_code=410) @@ -3886,7 +3933,7 @@ async def api_guest_config(token: str, connection_id: str, request: Request): config = _manager_call(manager, 'get_client_config', conn['protocol'], conn['client_id'], server['host'], port) ssh.disconnect() vpn_link = generate_vpn_link(config) if config else '' - return {'config': config, 'vpn_link': vpn_link} + return {'config': config, 'vpn_link': vpn_link, 'expires_at': holder.get('expiration_date')} except Exception as e: logger.exception("Error getting guest config") return JSONResponse({'error': str(e)}, status_code=500) @@ -3909,6 +3956,10 @@ async def api_guest_create(token: str, req: GuestCreateRequest, request: Request name = f"{name}_{secrets.token_hex(3)}" try: + from managers.user_expiration import maybe_start_user_expiration, user_is_expired + if user_is_expired(holder): + return JSONResponse({'error': 'Subscription expired'}, status_code=403) + sid = int(guest.get('create_server_id') or 0) if sid >= len(data['servers']): return JSONResponse({'error': 'Guest server not found'}, status_code=400) @@ -3944,6 +3995,7 @@ async def api_guest_create(token: str, req: GuestCreateRequest, request: Request async with DATA_LOCK: data = load_data() data['user_connections'].append(conn) + maybe_start_user_expiration(data, holder['id']) save_data(data) config = result.get('config') or '' @@ -3955,6 +4007,7 @@ async def api_guest_create(token: str, req: GuestCreateRequest, request: Request 'config': config, 'subscription_url': subscription_url, 'vpn_link': vpn_link, + 'expires_at': next((u.get('expiration_date') for u in data.get('users', []) if u.get('id') == holder['id']), None), } except Exception as e: logger.exception("Error creating guest config") @@ -4323,7 +4376,14 @@ async def api_my_connection_config(request: Request, connection_id: str): if not user: return JSONResponse({'error': 'Forbidden'}, status_code=403) try: + from managers.user_expiration import maybe_start_user_expiration, user_is_expired data = load_data() + panel_user = next((u for u in data.get('users', []) if u['id'] == user['id']), None) + if panel_user and user_is_expired(panel_user): + return JSONResponse({'error': 'Subscription expired'}, status_code=403) + if panel_user and maybe_start_user_expiration(data, user['id']): + save_data(data) + conn = next( (c for c in data.get('user_connections', []) if c['id'] == connection_id and c['user_id'] == user['id']), None @@ -4347,7 +4407,11 @@ async def api_my_connection_config(request: Request, connection_id: str): config = _manager_call(manager, 'get_client_config', conn['protocol'], conn['client_id'], server['host'], port) ssh.disconnect() vpn_link = generate_vpn_link(config) if config else '' - return {'config': config, 'vpn_link': vpn_link} + expires_at = None + panel_user = next((u for u in data.get('users', []) if u['id'] == user['id']), None) + if panel_user: + expires_at = panel_user.get('expiration_date') + return {'config': config, 'vpn_link': vpn_link, 'expires_at': expires_at} except Exception as e: logger.exception("Error getting my connection config") return JSONResponse({'error': str(e)}, status_code=500) diff --git a/db/connection.py b/db/connection.py index 9d7cd86..41362f2 100644 --- a/db/connection.py +++ b/db/connection.py @@ -93,6 +93,12 @@ def init_schema(): cur.execute( "ALTER TABLE user_connections ADD COLUMN IF NOT EXISTS xui_panel_id TEXT NOT NULL DEFAULT ''" ) + cur.execute( + "ALTER TABLE users ADD COLUMN IF NOT EXISTS expire_after_first_use BOOLEAN NOT NULL DEFAULT FALSE" + ) + cur.execute( + "ALTER TABLE users ADD COLUMN IF NOT EXISTS expiration_days INTEGER NOT NULL DEFAULT 0" + ) conn.commit() _schema_ready = True logger.info('PostgreSQL schema ready') diff --git a/db/schema.sql b/db/schema.sql index d08f635..6dce1cb 100644 --- a/db/schema.sql +++ b/db/schema.sql @@ -28,6 +28,8 @@ CREATE TABLE IF NOT EXISTS users ( traffic_reset_strategy TEXT NOT NULL DEFAULT 'never', last_reset_at TIMESTAMPTZ, expiration_date TIMESTAMPTZ, + expire_after_first_use BOOLEAN NOT NULL DEFAULT FALSE, + expiration_days INTEGER NOT NULL DEFAULT 0, remnawave_uuid TEXT, xui_email TEXT, share_enabled BOOLEAN NOT NULL DEFAULT FALSE, diff --git a/db/store.py b/db/store.py index bbea24d..9edf9b3 100644 --- a/db/store.py +++ b/db/store.py @@ -151,6 +151,8 @@ def _row_to_user(row) -> dict: 'traffic_reset_strategy': row['traffic_reset_strategy'] or 'never', 'last_reset_at': _ts_iso(row['last_reset_at']), 'expiration_date': _ts_iso(row['expiration_date']), + 'expire_after_first_use': bool(row.get('expire_after_first_use') if hasattr(row, 'get') else row['expire_after_first_use']), + 'expiration_days': int((row.get('expiration_days') if hasattr(row, 'get') else row['expiration_days']) or 0), 'remnawave_uuid': row['remnawave_uuid'], 'xui_email': row.get('xui_email'), 'share_enabled': bool(row['share_enabled']), @@ -222,7 +224,8 @@ def load_data() -> dict: 'SELECT id, username, password_hash, role, enabled, created_at, ' 'telegram_id, email, description, traffic_limit, traffic_used, ' 'traffic_total, traffic_reset_strategy, last_reset_at, ' - 'expiration_date, remnawave_uuid, xui_email, share_enabled, share_token, ' + 'expiration_date, expire_after_first_use, expiration_days, ' + 'remnawave_uuid, xui_email, share_enabled, share_token, ' 'share_password_hash FROM users ORDER BY created_at NULLS LAST, username' ) users = [_row_to_user(r) for r in cur.fetchall()] @@ -316,10 +319,10 @@ def save_data(data: dict) -> None: 'id, username, password_hash, role, enabled, created_at, ' 'telegram_id, email, description, traffic_limit, traffic_used, ' 'traffic_total, traffic_reset_strategy, last_reset_at, ' - 'expiration_date, remnawave_uuid, xui_email, share_enabled, share_token, ' + 'expiration_date, expire_after_first_use, expiration_days, remnawave_uuid, xui_email, share_enabled, share_token, ' 'share_password_hash' ') VALUES (' - '%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s' + '%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s' ')', ( _as_uuid(user['id']), @@ -337,6 +340,8 @@ def save_data(data: dict) -> None: user.get('traffic_reset_strategy') or 'never', _parse_ts(user.get('last_reset_at')), _parse_ts(user.get('expiration_date')), + bool(user.get('expire_after_first_use', False)), + int(user.get('expiration_days') or 0), user.get('remnawave_uuid'), user.get('xui_email'), bool(user.get('share_enabled', False)), diff --git a/managers/user_expiration.py b/managers/user_expiration.py new file mode 100644 index 0000000..cd25a9a --- /dev/null +++ b/managers/user_expiration.py @@ -0,0 +1,42 @@ +"""User subscription expiration helpers.""" + +from __future__ import annotations + +from datetime import datetime, timedelta +from typing import Optional + + +def user_is_expired(user: dict, *, now: Optional[datetime] = None) -> bool: + """True when absolute expiration_date is in the past.""" + exp_str = user.get('expiration_date') + if not exp_str: + return False + try: + exp_date = datetime.fromisoformat(str(exp_str)) + except Exception: + return False + return (now or datetime.now()) > exp_date + + +def maybe_start_user_expiration(data: dict, user_id: str, *, now: Optional[datetime] = None) -> Optional[str]: + """If user waits for first config use, start the countdown. + + Returns the new expiration_date ISO string when activated, else None. + """ + if not user_id: + return None + user = next((u for u in data.get('users', []) if u.get('id') == user_id), None) + if not user: + return None + if not user.get('expire_after_first_use'): + return None + if user.get('expiration_date'): + return None + days = int(user.get('expiration_days') or 0) + if days <= 0: + return None + stamp = now or datetime.now() + expires = stamp + timedelta(days=days) + iso = expires.isoformat() + user['expiration_date'] = iso + return iso diff --git a/telegram_bot.py b/telegram_bot.py index 48fe59b..87210a3 100644 --- a/telegram_bot.py +++ b/telegram_bot.py @@ -499,7 +499,7 @@ async def _handle_refresh(api: TelegramAPI, chat_id: int, message_id: int, callb await api.edit_message(chat_id, message_id, f"Your connections ({len(conns)}) — tap to get config:", reply_markup=kb) -async def _handle_get_config(api: TelegramAPI, chat_id: int, message_id: int, callback_id: str, conn_id: str, tg_id: str, load_data_fn: Callable, generate_vpn_link_fn: Callable): +async def _handle_get_config(api: TelegramAPI, chat_id: int, message_id: int, callback_id: str, conn_id: str, tg_id: str, load_data_fn: Callable, generate_vpn_link_fn: Callable, save_data_fn: Optional[Callable] = None): await api.answer_callback(callback_id, "Fetching config...") panel_user = _find_user(load_data_fn, tg_id) @@ -508,6 +508,18 @@ async def _handle_get_config(api: TelegramAPI, chat_id: int, message_id: int, ca return data = load_data_fn() + try: + from managers.user_expiration import maybe_start_user_expiration, user_is_expired + owner = next((u for u in data.get("users", []) if u.get("id") == panel_user.get("id")), panel_user) + if user_is_expired(owner): + await api.send_message(chat_id, "❌ Subscription expired.") + return + if save_data_fn and maybe_start_user_expiration(data, owner.get("id")): + save_data_fn(data) + data = load_data_fn() + except Exception: + logger.exception("Bot: failed to start expiration on first use") + conn = next((c for c in data.get("user_connections", []) if c.get("id") == conn_id and (_is_admin(panel_user) or c.get("user_id") == panel_user.get("id"))), None) if not conn: await api.send_message(chat_id, "❌ Connection not found.") @@ -1072,7 +1084,7 @@ async def _dispatch(api: TelegramAPI, update: dict, load_data_fn: Callable, gene await _handle_refresh(api, chat_id, message_id, callback_id, tg_id, load_data_fn) return if data_str.startswith("cfg:"): - await _handle_get_config(api, chat_id, message_id, callback_id, data_str[4:], tg_id, load_data_fn, generate_vpn_link_fn) + await _handle_get_config(api, chat_id, message_id, callback_id, data_str[4:], tg_id, load_data_fn, generate_vpn_link_fn, save_data_fn) return panel_user = _require_admin(load_data_fn, tg_id) diff --git a/templates/users.html b/templates/users.html index 3485d80..0e6bc02 100644 --- a/templates/users.html +++ b/templates/users.html @@ -109,10 +109,21 @@