feat(renewal): per-subscription renewal date base setting for manual renewals

This commit is contained in:
Xiaolan Bot
2026-08-24 20:29:39 +08:00
parent 2338c375b9
commit 2c9d538f55
2 changed files with 109 additions and 19 deletions
+2
View File
@@ -5,6 +5,7 @@
## 功能特性 ## 功能特性
- ➕ 添加订阅(名称、费用、货币、分类、到期日、周期、续费方式、备注) - ➕ 添加订阅(名称、费用、货币、分类、到期日、周期、续费方式、备注)
- 🔁 手动续订日期起算方式可选:从续费日起算,或从到期日顺延(提前续费不损失天数)
- 📋 列出订阅并查看详情 - 📋 列出订阅并查看详情
- 🗂️ 按分类浏览订阅 - 🗂️ 按分类浏览订阅
- ✏️ 编辑订阅信息 - ✏️ 编辑订阅信息
@@ -103,6 +104,7 @@ python SubMind.py
- `frequency_value`(正整数) - `frequency_value`(正整数)
- `renewal_type``auto` / `manual` - `renewal_type``auto` / `manual`
- `notes`(可选) - `notes`(可选)
- `renewal_base`(可选,仅手动续费有效:`today`=从续费日起算,`due_date`=从到期日顺延;默认 `today`
示例: 示例:
+107 -19
View File
@@ -51,7 +51,7 @@ AUTO_UPDATE_BRANCH = os.getenv('AUTO_UPDATE_BRANCH', 'main').strip() or 'main'
# --- 对话处理器状态 --- # --- 对话处理器状态 ---
(ADD_NAME, ADD_COST, ADD_CURRENCY, ADD_CATEGORY, ADD_NEXT_DUE, (ADD_NAME, ADD_COST, ADD_CURRENCY, ADD_CATEGORY, ADD_NEXT_DUE,
ADD_FREQ_UNIT, ADD_FREQ_VALUE, ADD_RENEWAL_TYPE, ADD_NOTES) = range(9) ADD_FREQ_UNIT, ADD_FREQ_VALUE, ADD_RENEWAL_TYPE, ADD_RENEWAL_BASE, ADD_NOTES) = range(10)
(EDIT_SELECT_FIELD, EDIT_GET_NEW_VALUE, EDIT_FREQ_UNIT, EDIT_FREQ_VALUE) = range(4) (EDIT_SELECT_FIELD, EDIT_GET_NEW_VALUE, EDIT_FREQ_UNIT, EDIT_FREQ_VALUE) = range(4)
(REMIND_SELECT_ACTION, REMIND_GET_DAYS) = range(2) (REMIND_SELECT_ACTION, REMIND_GET_DAYS) = range(2)
(IMPORT_UPLOAD,) = range(1) (IMPORT_UPLOAD,) = range(1)
@@ -115,6 +115,7 @@ def init_db():
id INTEGER PRIMARY KEY, user_id INTEGER, name TEXT, cost REAL, currency TEXT, id INTEGER PRIMARY KEY, user_id INTEGER, name TEXT, cost REAL, currency TEXT,
category TEXT, next_due DATE, frequency TEXT, category TEXT, next_due DATE, frequency TEXT,
renewal_type TEXT DEFAULT 'auto', renewal_type TEXT DEFAULT 'auto',
renewal_base TEXT DEFAULT 'today',
reminders_enabled BOOLEAN DEFAULT TRUE, reminders_enabled BOOLEAN DEFAULT TRUE,
reminder_days INTEGER DEFAULT 3, reminder_days INTEGER DEFAULT 3,
reminder_on_due_date BOOLEAN DEFAULT TRUE, reminder_on_due_date BOOLEAN DEFAULT TRUE,
@@ -140,6 +141,8 @@ def init_db():
cursor.execute("ALTER TABLE subscriptions ADD COLUMN notes TEXT") cursor.execute("ALTER TABLE subscriptions ADD COLUMN notes TEXT")
if 'last_reminded_date' not in columns: if 'last_reminded_date' not in columns:
cursor.execute("ALTER TABLE subscriptions ADD COLUMN last_reminded_date DATE") cursor.execute("ALTER TABLE subscriptions ADD COLUMN last_reminded_date DATE")
if 'renewal_base' not in columns:
cursor.execute("ALTER TABLE subscriptions ADD COLUMN renewal_base TEXT DEFAULT 'today'")
cursor.execute(''' cursor.execute('''
CREATE TABLE IF NOT EXISTS categories ( CREATE TABLE IF NOT EXISTS categories (
@@ -257,6 +260,30 @@ def calculate_new_due_date(base_date, unit, value):
return base_date + delta if delta else None return base_date + delta if delta else None
# 手动续费的日期起算方式:
# - today: 以点击续费当天为起点重新计算
# - due_date: 以上一周期到期日为起点顺延(提前续费不损失天数;已过期则从今天起算)
VALID_RENEWAL_BASES = {'today', 'due_date'}
RENEWAL_BASE_LABELS = {
'today': '从续费日起算',
'due_date': '从到期日顺延',
}
def calculate_renewal_due_date(renewal_base, next_due, unit, value):
if str(renewal_base or 'today') == 'due_date':
try:
base = datetime.datetime.strptime(next_due, '%Y-%m-%d').date() if next_due else None
except (ValueError, TypeError):
base = None
today = datetime.date.today()
if base is None or base < today:
base = today
else:
base = datetime.date.today()
return calculate_new_due_date(base, unit, value)
def format_frequency(unit, value) -> str: def format_frequency(unit, value) -> str:
if not unit or value is None: if not unit or value is None:
return "未知" return "未知"
@@ -275,6 +302,7 @@ EDITABLE_SUB_FIELDS = {
'category': 'category', 'category': 'category',
'next_due': 'next_due', 'next_due': 'next_due',
'renewal_type': 'renewal_type', 'renewal_type': 'renewal_type',
'renewal_base': 'renewal_base',
'notes': 'notes' 'notes': 'notes'
} }
MAX_NAME_LEN = 128 MAX_NAME_LEN = 128
@@ -636,7 +664,7 @@ async def export_command(update: Update, context: CallbackContext):
def process_export(): def process_export():
with get_db_connection() as conn: with get_db_connection() as conn:
df = pd.read_sql_query( df = pd.read_sql_query(
"SELECT name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, notes FROM subscriptions WHERE user_id = ?", "SELECT name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, renewal_base, notes FROM subscriptions WHERE user_id = ?",
conn, params=(user_id,)) conn, params=(user_id,))
if df.empty: if df.empty:
return False, None return False, None
@@ -665,7 +693,7 @@ async def export_command(update: Update, context: CallbackContext):
async def import_start(update: Update, context: CallbackContext): async def import_start(update: Update, context: CallbackContext):
await update.message.reply_text( await update.message.reply_text(
"请上传一个 CSV 文件以导入订阅数据。\n文件应包含以下列:name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, notesnotes 可为空)。") "请上传一个 CSV 文件以导入订阅数据。\n文件应包含以下列:name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, notesnotes 可为空)。\n可选列:renewal_base(手动续费的日期起算方式,today=从续费日起算 / due_date=从到期日顺延,默认 today)。")
return IMPORT_UPLOAD return IMPORT_UPLOAD
@@ -690,6 +718,8 @@ async def import_upload_received(update: Update, context: CallbackContext):
valid_units = ['day', 'week', 'month', 'year'] valid_units = ['day', 'week', 'month', 'year']
valid_renewal_types = ['auto', 'manual'] valid_renewal_types = ['auto', 'manual']
valid_renewal_bases = ['today', 'due_date']
has_renewal_base_col = 'renewal_base' in df.columns
records = [] records = []
for _, row in df.iterrows(): for _, row in df.iterrows():
try: try:
@@ -711,6 +741,11 @@ async def import_upload_received(update: Update, context: CallbackContext):
renewal_type = str(row['renewal_type']).lower() renewal_type = str(row['renewal_type']).lower()
if renewal_type not in valid_renewal_types: if renewal_type not in valid_renewal_types:
raise ValueError(f"无效续费类型: {renewal_type}") raise ValueError(f"无效续费类型: {renewal_type}")
renewal_base = 'today'
if renewal_type == 'manual' and has_renewal_base_col and pd.notna(row['renewal_base']):
renewal_base = str(row['renewal_base']).lower().strip()
if renewal_base not in valid_renewal_bases:
raise ValueError(f"无效续订起算方式: {renewal_base}")
notes = str(row['notes']).strip() if pd.notna(row['notes']) else None notes = str(row['notes']).strip() if pd.notna(row['notes']) else None
if notes and len(notes) > MAX_NOTES_LEN: if notes and len(notes) > MAX_NOTES_LEN:
raise ValueError(f"备注过长(>{MAX_NOTES_LEN}") raise ValueError(f"备注过长(>{MAX_NOTES_LEN}")
@@ -726,7 +761,7 @@ async def import_upload_received(update: Update, context: CallbackContext):
raise ValueError(f"类别过长(>{MAX_CATEGORY_LEN}") raise ValueError(f"类别过长(>{MAX_CATEGORY_LEN}")
records.append(( records.append((
user_id, name, cost, currency, category, user_id, name, cost, currency, category,
next_due, frequency_unit, frequency_value, renewal_type, notes next_due, frequency_unit, frequency_value, renewal_type, renewal_base, notes
)) ))
except Exception as e: except Exception as e:
logger.error(f"Invalid row in CSV import, error: {e}") logger.error(f"Invalid row in CSV import, error: {e}")
@@ -736,8 +771,8 @@ async def import_upload_received(update: Update, context: CallbackContext):
with get_db_connection() as conn: with get_db_connection() as conn:
cursor = conn.cursor() cursor = conn.cursor()
cursor.executemany(''' cursor.executemany('''
INSERT INTO subscriptions (user_id, name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, notes) INSERT INTO subscriptions (user_id, name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, renewal_base, notes)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
''', records) ''', records)
for record in records: for record in records:
cursor.execute("INSERT OR IGNORE INTO categories (user_id, name) VALUES (?, ?)", (user_id, record[4])) cursor.execute("INSERT OR IGNORE INTO categories (user_id, name) VALUES (?, ?)", (user_id, record[4]))
@@ -916,6 +951,35 @@ async def add_renewal_type_received(update: Update, context: CallbackContext):
await query.edit_message_text("错误:无效的续费类型,请重试。") await query.edit_message_text("错误:无效的续费类型,请重试。")
return ConversationHandler.END return ConversationHandler.END
sub_data['renewal_type'] = renewal_type sub_data['renewal_type'] = renewal_type
if renewal_type == 'manual':
keyboard = [
[InlineKeyboardButton("从续费日起算", callback_data='renewbase_today'),
InlineKeyboardButton("从到期日顺延", callback_data='renewbase_due_date')]
]
await query.edit_message_text(
"第九步:请选择手动续费的<b>日期起算方式</b>\n\n"
"- <b>从续费日起算</b>:以您点击续费当天为起点重新计算\n"
"- <b>从到期日顺延</b>:以上一周期到期日为起点顺延(提前续费不损失天数)",
reply_markup=InlineKeyboardMarkup(keyboard), parse_mode='HTML')
return ADD_RENEWAL_BASE
sub_data['renewal_base'] = 'today'
await query.edit_message_text("最后一步(可选):需要添加备注吗?\n(如:共享账号、用途等。不需要请 /skip)")
return ADD_NOTES
async def add_renewal_base_received(update: Update, context: CallbackContext):
sub_data, _ = _get_new_sub_data_or_end(update, context)
query = update.callback_query
await query.answer()
if sub_data is None:
await query.edit_message_text("会话已过期,请重新使用 /add_sub 开始。")
return ConversationHandler.END
renewal_base = query.data.partition('renewbase_')[2]
if renewal_base not in VALID_RENEWAL_BASES:
await query.edit_message_text("错误:无效的日期起算方式,请重试。")
return ConversationHandler.END
sub_data['renewal_base'] = renewal_base
await query.edit_message_text("最后一步(可选):需要添加备注吗?\n(如:共享账号、用途等。不需要请 /skip)") await query.edit_message_text("最后一步(可选):需要添加备注吗?\n(如:共享账号、用途等。不需要请 /skip)")
return ADD_NOTES return ADD_NOTES
@@ -960,12 +1024,13 @@ def save_subscription(user_id, data):
with get_db_connection() as conn: with get_db_connection() as conn:
cursor = conn.cursor() cursor = conn.cursor()
cursor.execute(''' cursor.execute('''
INSERT INTO subscriptions (user_id, name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, notes) INSERT INTO subscriptions (user_id, name, cost, currency, category, next_due, frequency_unit, frequency_value, renewal_type, renewal_base, notes)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
''', ( ''', (
user_id, data.get('name'), data.get('cost'), data.get('currency'), data.get('category'), user_id, data.get('name'), data.get('cost'), data.get('currency'), data.get('category'),
data.get('next_due'), data.get('next_due'),
data.get('unit'), data.get('value'), data.get('renewal_type', 'auto'), data.get('notes') data.get('unit'), data.get('value'), data.get('renewal_type', 'auto'),
data.get('renewal_base', 'today'), data.get('notes')
)) ))
conn.commit() conn.commit()
@@ -1026,12 +1091,17 @@ async def show_subscription_view(update: Update, context: CallbackContext, sub_i
safe_name, safe_category, safe_freq = escape_html(name), escape_html(category), escape_html(freq_text) safe_name, safe_category, safe_freq = escape_html(name), escape_html(category), escape_html(freq_text)
cost_str, converted_cost_str = escape_html(f"{cost:.2f}"), escape_html(f"{converted_cost:.2f}") cost_str, converted_cost_str = escape_html(f"{cost:.2f}"), escape_html(f"{converted_cost:.2f}")
renewal_text = "手动续费" if renewal_type == 'manual' else "自动续费" renewal_text = "手动续费" if renewal_type == 'manual' else "自动续费"
renewal_base_text = ""
if renewal_type == 'manual':
base_label = RENEWAL_BASE_LABELS.get(sub['renewal_base'] or 'today', sub['renewal_base'])
renewal_base_text = f"\n- <b>续订起算</b>: <code>{escape_html(base_label)}</code>"
reminder_status = "开启" if reminders_enabled else "关闭" reminder_status = "开启" if reminders_enabled else "关闭"
text = (f"<b>订阅详情: {safe_name}</b>\n\n" text = (f"<b>订阅详情: {safe_name}</b>\n\n"
f"- <b>费用</b>: <code>{cost_str} {currency.upper()}</code> (~<code>{converted_cost_str} {main_currency.upper()}</code>)\n" f"- <b>费用</b>: <code>{cost_str} {currency.upper()}</code> (~<code>{converted_cost_str} {main_currency.upper()}</code>)\n"
f"- <b>类别</b>: <code>{safe_category}</code>\n" f"- <b>类别</b>: <code>{safe_category}</code>\n"
f"- <b>下次付款</b>: <code>{next_due}</code> (周期: {safe_freq})\n" f"- <b>下次付款</b>: <code>{next_due}</code> (周期: {safe_freq})\n"
f"- <b>续费方式</b>: <code>{renewal_text}</code>\n" f"- <b>续费方式</b>: <code>{renewal_text}</code>"
f"{renewal_base_text}\n"
f"- <b>提醒状态</b>: <code>{reminder_status}</code>") f"- <b>提醒状态</b>: <code>{reminder_status}</code>")
if notes: if notes:
text += f"\n- <b>备注</b>: {escape_html(notes)}" text += f"\n- <b>备注</b>: {escape_html(notes)}"
@@ -1122,13 +1192,13 @@ async def button_callback_handler(update: Update, context: CallbackContext):
with get_db_connection() as conn: with get_db_connection() as conn:
cursor = conn.cursor() cursor = conn.cursor()
cursor.execute( cursor.execute(
"SELECT frequency_unit, frequency_value FROM subscriptions WHERE id = ? AND user_id = ?", "SELECT next_due, frequency_unit, frequency_value, renewal_base FROM subscriptions WHERE id = ? AND user_id = ?",
(sub_id, user_id) (sub_id, user_id)
) )
sub = cursor.fetchone() sub = cursor.fetchone()
if sub: if sub:
today = datetime.date.today() new_due_date = calculate_renewal_due_date(
new_due_date = calculate_new_due_date(today, sub['frequency_unit'], sub['frequency_value']) sub['renewal_base'], sub['next_due'], sub['frequency_unit'], sub['frequency_value'])
if new_due_date: if new_due_date:
new_date_str = new_due_date.strftime('%Y-%m-%d') new_date_str = new_due_date.strftime('%Y-%m-%d')
cursor.execute( cursor.execute(
@@ -1147,13 +1217,13 @@ async def button_callback_handler(update: Update, context: CallbackContext):
with get_db_connection() as conn: with get_db_connection() as conn:
cursor = conn.cursor() cursor = conn.cursor()
cursor.execute( cursor.execute(
"SELECT name, frequency_unit, frequency_value FROM subscriptions WHERE id = ? AND user_id = ?", "SELECT name, next_due, frequency_unit, frequency_value, renewal_base FROM subscriptions WHERE id = ? AND user_id = ?",
(sub_id, user_id) (sub_id, user_id)
) )
sub = cursor.fetchone() sub = cursor.fetchone()
if sub: if sub:
today = datetime.date.today() new_due_date = calculate_renewal_due_date(
new_due_date = calculate_new_due_date(today, sub['frequency_unit'], sub['frequency_value']) sub['renewal_base'], sub['next_due'], sub['frequency_unit'], sub['frequency_value'])
if new_due_date: if new_due_date:
new_date_str = new_due_date.strftime('%Y-%m-%d') new_date_str = new_due_date.strftime('%Y-%m-%d')
cursor.execute( cursor.execute(
@@ -1254,8 +1324,9 @@ async def edit_start(update: Update, context: CallbackContext):
[InlineKeyboardButton("下次付款日", callback_data="editfield_next_due"), [InlineKeyboardButton("下次付款日", callback_data="editfield_next_due"),
InlineKeyboardButton("周期", callback_data="editfield_frequency")], InlineKeyboardButton("周期", callback_data="editfield_frequency")],
[InlineKeyboardButton("续费方式", callback_data="editfield_renewal_type"), [InlineKeyboardButton("续费方式", callback_data="editfield_renewal_type"),
InlineKeyboardButton("📝 备注", callback_data="editfield_notes")], InlineKeyboardButton("续订起算", callback_data="editfield_renewal_base")],
[InlineKeyboardButton("« 返回详情", callback_data=f'view_{sub_id}')] [InlineKeyboardButton("📝 备注", callback_data="editfield_notes"),
InlineKeyboardButton("« 返回详情", callback_data=f'view_{sub_id}')]
] ]
await query.edit_message_text("请选择您想编辑的字段:", reply_markup=InlineKeyboardMarkup(keyboard)) await query.edit_message_text("请选择您想编辑的字段:", reply_markup=InlineKeyboardMarkup(keyboard))
return EDIT_SELECT_FIELD return EDIT_SELECT_FIELD
@@ -1273,6 +1344,17 @@ async def edit_field_selected(update: Update, context: CallbackContext):
] ]
await query.edit_message_text("请选择新的续费方式:", reply_markup=InlineKeyboardMarkup(keyboard)) await query.edit_message_text("请选择新的续费方式:", reply_markup=InlineKeyboardMarkup(keyboard))
return EDIT_GET_NEW_VALUE return EDIT_GET_NEW_VALUE
if field_to_edit == 'renewal_base':
keyboard = [
[InlineKeyboardButton("从续费日起算", callback_data='editvalue_today'),
InlineKeyboardButton("从到期日顺延", callback_data='editvalue_due_date')]
]
await query.edit_message_text(
"请选择新的续订日期起算方式:\n\n"
"- <b>从续费日起算</b>:以点击续费当天为起点重新计算\n"
"- <b>从到期日顺延</b>:以上一周期到期日为起点顺延(提前续费不损失天数)",
reply_markup=InlineKeyboardMarkup(keyboard), parse_mode='HTML')
return EDIT_GET_NEW_VALUE
if field_to_edit == 'frequency': if field_to_edit == 'frequency':
keyboard = [ keyboard = [
[InlineKeyboardButton("", callback_data='freq_unit_day'), [InlineKeyboardButton("", callback_data='freq_unit_day'),
@@ -1362,7 +1444,7 @@ async def edit_new_value_received(update: Update, context: CallbackContext):
if update.message and update.message.text == '/empty' and field == 'notes': if update.message and update.message.text == '/empty' and field == 'notes':
new_value = None new_value = None
elif query: elif query:
new_value = query.data.split('_')[1] new_value = query.data.partition('_')[2]
elif update.message: elif update.message:
new_value = update.message.text new_value = update.message.text
else: else:
@@ -1409,6 +1491,11 @@ async def edit_new_value_received(update: Update, context: CallbackContext):
if message_to_reply: if message_to_reply:
await message_to_reply.reply_text("续费方式只能为 auto 或 manual。") await message_to_reply.reply_text("续费方式只能为 auto 或 manual。")
validation_failed = True validation_failed = True
elif field == 'renewal_base':
if str(new_value) not in VALID_RENEWAL_BASES:
if message_to_reply:
await message_to_reply.reply_text("续订起算方式无效。")
validation_failed = True
elif field == 'notes': elif field == 'notes':
note_val = str(new_value).strip() note_val = str(new_value).strip()
if note_val and len(note_val) > MAX_NOTES_LEN: if note_val and len(note_val) > MAX_NOTES_LEN:
@@ -1798,6 +1885,7 @@ def main():
ADD_FREQ_UNIT: [CallbackQueryHandler(add_freq_unit_received, pattern='^freq_unit_')], ADD_FREQ_UNIT: [CallbackQueryHandler(add_freq_unit_received, pattern='^freq_unit_')],
ADD_FREQ_VALUE: [MessageHandler(filters.TEXT & ~filters.COMMAND, add_freq_value_received)], ADD_FREQ_VALUE: [MessageHandler(filters.TEXT & ~filters.COMMAND, add_freq_value_received)],
ADD_RENEWAL_TYPE: [CallbackQueryHandler(add_renewal_type_received, pattern='^renewal_')], ADD_RENEWAL_TYPE: [CallbackQueryHandler(add_renewal_type_received, pattern='^renewal_')],
ADD_RENEWAL_BASE: [CallbackQueryHandler(add_renewal_base_received, pattern='^renewbase_')],
ADD_NOTES: [ ADD_NOTES: [
MessageHandler(filters.TEXT & ~filters.COMMAND, add_notes_received), MessageHandler(filters.TEXT & ~filters.COMMAND, add_notes_received),
CommandHandler('skip', skip_notes) CommandHandler('skip', skip_notes)