Files
auto_control/core/models.py
T
butubb 838886e043 security: 环境变量密钥、密码加盐哈希、CSRF 防护
- config.py/web_server.py 支持环境变量注入 STF_TOKEN/WEB_SECRET_KEY,新增 .env 加载和 .env.example;.env 加入 gitignore
- 用户密码从裸 SHA-256 改为 werkzeug 加盐哈希(scrypt),旧哈希登录时自动升级
- CSRF:session token + X-CSRF-Token 请求头校验非 GET 请求,前端自动携带
2026-08-08 21:29:39 +08:00

334 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""SQLAlchemy 数据模型 + db 初始化(用户/设备分组/任务计划)。
为什么放这里:
web_server 和 task_manager 都要访问这些模型,放 core/ 避免循环依赖。
web_server 负责初始化 db(app context),task_manager 只读写数据。
模型说明:
User — 后台用户(Flask-Login 认证)
DeviceGroup — 设备分组(持久化,替代旧 groups.json)
TaskJob — 任务计划(持久化,替代旧 jobs.json)
字段设计:
DeviceGroup.serials 用 JSON 存列表(SQLite 无数组类型)
TaskJob.target/params/schedule/retry 用 JSON 存嵌套结构
Flask-Admin 默认用 TextArea 编辑 JSON,够用且通用
"""
import os
import json
import hashlib
from flask_sqlalchemy import SQLAlchemy
from flask_login import UserMixin
from sqlalchemy import event, text
from sqlalchemy.engine import Engine
from werkzeug.security import generate_password_hash, check_password_hash
from core.logger import get_logger
_log = get_logger("core.models")
db = SQLAlchemy()
@event.listens_for(Engine, "connect")
def _sqlite_pragma(dbapi_connection, connection_record):
"""SQLite 并发写优化:WAL 模式 + 忙等待超时 + 降同步级别。
多 worker 后台线程同时写库(任务参数/分组)时,避免 database is locked。
"""
cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA journal_mode=WAL")
cursor.execute("PRAGMA busy_timeout=5000")
cursor.execute("PRAGMA synchronous=NORMAL")
cursor.close()
class User(UserMixin, db.Model):
"""后台用户。"""
id = db.Column(db.Integer, primary_key=True)
username = db.Column(db.String(80), unique=True, nullable=False)
password_hash = db.Column(db.String(255), nullable=False)
is_admin = db.Column(db.Boolean, default=True)
def set_password(self, password):
self.password_hash = generate_password_hash(password)
def check_password(self, password):
"""校验密码。兼容旧 SHA-256 哈希(匹配则自动升级为新哈希)。"""
ph = self.password_hash or ""
if ph and not ph.startswith(("pbkdf2:", "scrypt:")):
# 旧版裸 SHA-256:比对通过则升级为加盐哈希
if hashlib.sha256(password.encode()).hexdigest() == ph:
self.set_password(password)
try:
db.session.commit()
except Exception:
pass
return True
return False
return check_password_hash(ph, password)
def __repr__(self):
return f"<User {self.username}>"
class DeviceGroup(db.Model):
"""设备分组(替代旧 groups.json 的 DeviceGroup 类)。"""
id = db.Column(db.Integer, primary_key=True)
name = db.Column(db.String(80), unique=True, nullable=False)
serials = db.Column(db.Text, default="[]") # JSON 列表
description = db.Column(db.Text, default="")
def get_serials(self):
try:
return json.loads(self.serials or "[]")
except Exception:
return []
def set_serials(self, lst):
self.serials = json.dumps(lst or [], ensure_ascii=False)
def to_dict(self):
return {"name": self.name, "serials": self.get_serials(),
"description": self.description or ""}
def __repr__(self):
return f"<Group {self.name}>"
class TaskJob(db.Model):
"""任务计划(替代旧 jobs.json 的 TaskJob 类)。
字段含义和旧 TaskJob 一致,只是持久化方式从 JSON 文件改到 SQLite。
"""
id = db.Column(db.String(32), primary_key=True) # uuid 前 8 位
name = db.Column(db.String(120), nullable=False)
task_type = db.Column(db.String(60), default="douyin_nurture")
target = db.Column(db.Text, default='{"mode":"all"}') # JSON
params = db.Column(db.Text, default="{}") # JSON
schedule = db.Column(db.Text, default='{"mode":"once"}') # JSON
retry = db.Column(db.Text, default='{"max_attempts":1,"delay":60}') # JSON
enabled = db.Column(db.Boolean, default=True)
def _load_json(self, field, default):
try:
return json.loads(getattr(self, field) or default)
except Exception:
return json.loads(default)
def _dump_json(self, field, value):
setattr(self, field, json.dumps(value or {}, ensure_ascii=False))
def get_target(self): return self._load_json("target", '{"mode":"all"}')
def set_target(self, v): self._dump_json("target", v)
def get_params(self): return self._load_json("params", "{}")
def set_params(self, v): self._dump_json("params", v)
def get_schedule(self): return self._load_json("schedule", '{"mode":"once"}')
def set_schedule(self, v): self._dump_json("schedule", v)
def get_retry(self): return self._load_json("retry", '{"max_attempts":1,"delay":60}')
def set_retry(self, v): self._dump_json("retry", v)
def to_dict(self):
return {"id": self.id, "name": self.name, "task_type": self.task_type,
"target": self.get_target(), "params": self.get_params(),
"schedule": self.get_schedule(), "retry": self.get_retry(),
"enabled": self.enabled}
def __repr__(self):
return f"<TaskJob {self.name}>"
class CustomAction(db.Model):
"""自定义动作(把一系列步骤打包成一个可复用的动作)。
steps 字段存 JSON 数组,格式同 generic_steps 的 step schema。
前端拖拽到画布时,展开为 group 步骤(type='group')。
"""
id = db.Column(db.String(32), primary_key=True)
name = db.Column(db.String(120), nullable=False)
icon = db.Column(db.String(4), default="📦")
steps = db.Column(db.Text, default="[]") # JSON 数组
created_at = db.Column(db.String(20), default="")
def get_steps(self):
try:
return json.loads(self.steps or "[]")
except Exception:
return []
def set_steps(self, v):
self.steps = json.dumps(v or [], ensure_ascii=False)
def to_dict(self):
return {"id": self.id, "name": self.name, "icon": self.icon or "📦",
"steps": self.get_steps(), "created_at": self.created_at or ""}
def __repr__(self):
return f"<CustomAction {self.name}>"
class ApkFile(db.Model):
"""上传的 APK 文件元信息(应用管理功能)。"""
id = db.Column(db.String(32), primary_key=True) # uuid 前 8 位
filename = db.Column(db.String(255), nullable=False) # 磁盘文件名 (id.apk)
display_name = db.Column(db.String(120), default="") # 应用名
package_name = db.Column(db.String(200), default="") # 包名
version_name = db.Column(db.String(50), default="") # 版本号
version_code = db.Column(db.Integer, default=0) # 版本码
size = db.Column(db.Integer, default=0) # 文件大小(字节)
upload_time = db.Column(db.String(20), default="") # 上传时间
def to_dict(self):
return {"id": self.id, "filename": self.filename,
"display_name": self.display_name or "",
"package_name": self.package_name or "",
"version_name": self.version_name or "",
"version_code": self.version_code or 0,
"size": self.size or 0,
"upload_time": self.upload_time or ""}
def __repr__(self):
return f"<ApkFile {self.display_name}>"
# 版本化 schema 迁移:新增结构变更时在此追加 (版本号, 说明, SQL)
# 版本号单调递增,只执行比当前 schema_version 新的迁移。
SCHEMA_MIGRATIONS = [
# (1, "初始 schema(由 create_all 建立,版本标记从 1 开始)", None),
]
def init_db(app):
"""在 Flask app context 里初始化数据库 + 创建默认管理员。
web_server 启动时调用。自动迁移旧 groups.json/jobs.json 到 SQLite。
"""
db.init_app(app)
with app.app_context():
db.create_all()
_migrate_schema()
_ensure_default_admin()
_migrate_old_json()
def _migrate_schema():
"""按 SCHEMA_MIGRATIONS 顺序执行版本化迁移,记录当前 schema_version。
create_all 只负责首次建表;结构变更必须走迁移,避免改了模型后老库对不上。
"""
try:
db.session.execute(text(
"CREATE TABLE IF NOT EXISTS app_meta (key TEXT PRIMARY KEY, value TEXT)"))
db.session.commit()
cur = db.session.execute(
text("SELECT value FROM app_meta WHERE key='schema_version'")).scalar()
current = int(cur) if cur else 0
for version, desc, sql in SCHEMA_MIGRATIONS:
if version <= current:
continue
if sql:
db.session.execute(text(sql))
db.session.execute(
text("INSERT OR REPLACE INTO app_meta(key,value) VALUES('schema_version',:v)"),
{"v": str(version)})
db.session.commit()
_log.info(f"schema 迁移到版本 {version}: {desc}")
except Exception as e:
_log.error(f"schema 迁移失败(不阻塞启动): {e}")
def _ensure_default_admin():
"""首次启动创建默认管理员 admin/admin123。"""
if not User.query.filter_by(username="admin").first():
u = User(username="admin", is_admin=True)
u.set_password("admin123")
db.session.add(u)
db.session.commit()
_log.info("已创建默认管理员 admin/admin123,请及时改密码")
def _migrate_old_json():
"""把旧 groups.json / jobs.json 迁移到 SQLite(仅首次)。
迁移策略:
1. 仅当数据库对应表为空时才迁移(首次启动场景)
2. 迁移成功后立即把 JSON 文件重命名为 <name>.json.migrated,
保留备份但永不再迁移——避免"用户删完全部任务后重启又从旧文件复原"
3. 若数据库已有数据但 JSON 文件仍在(历史残留),直接归档,
避免将来数据库被清空后又触发迁移导致已删任务复原
4. 迁移任务时顺手剔除已废弃的 comment action 配置,
避免 _load 阶段还要再写回一次
"""
from config import DATA_DIR
# 迁移分组
groups_file = os.path.join(DATA_DIR, "groups.json")
if os.path.exists(groups_file) and DeviceGroup.query.count() == 0:
try:
with open(groups_file, encoding="utf-8") as f:
groups = json.load(f)
for g in groups:
if not DeviceGroup.query.filter_by(name=g["name"]).first():
row = DeviceGroup(name=g["name"],
description=g.get("description", ""))
row.set_serials(g.get("serials", []))
db.session.add(row)
db.session.commit()
_log.info(f"已迁移 {len(groups)} 个分组到数据库")
_archive_migrated(groups_file)
except Exception as e:
_log.error(f"迁移 groups.json 失败: {e}")
# 迁移任务
jobs_file = os.path.join(DATA_DIR, "jobs.json")
if os.path.exists(jobs_file) and TaskJob.query.count() == 0:
try:
with open(jobs_file, encoding="utf-8") as f:
jobs = json.load(f)
for j in jobs:
if TaskJob.query.get(j["id"]):
continue
params = j.get("params", {})
# 剔除已废弃的 comment action,避免带入数据库
actions = params.get("actions", {})
if "comment" in actions:
del actions["comment"]
_log.info(f"迁移任务 {j.get('name')}: 已剔除废弃的 comment 配置")
row = TaskJob(id=j["id"], name=j["name"],
task_type=j.get("task_type", "douyin_nurture"),
enabled=j.get("enabled", True))
row.set_target(j.get("target", {"mode": "all"}))
row.set_params(params)
row.set_schedule(j.get("schedule", {"mode": "once"}))
row.set_retry(j.get("retry", {"max_attempts": 1, "delay": 60}))
db.session.add(row)
db.session.commit()
_log.info(f"已迁移 {len(jobs)} 个任务到数据库")
_archive_migrated(jobs_file)
except Exception as e:
_log.error(f"迁移 jobs.json 失败: {e}")
# 清理历史残留:数据库已有数据但 JSON 文件仍在(修复前遗留下来的文件)。
# 不归档的话,用户哪天删光所有任务/分组,count==0 又会触发迁移导致已删数据复原。
if DeviceGroup.query.count() > 0 and os.path.exists(groups_file):
_archive_migrated(groups_file)
if TaskJob.query.count() > 0 and os.path.exists(jobs_file):
_archive_migrated(jobs_file)
def _archive_migrated(file_path):
"""把已迁移的 JSON 文件重命名为 <name>.migrated,避免下次重启再迁移。
保留备份以便排查,但 _migrate_old_json 的 exists 判断会跳过它。
重命名失败只警告不抛出,不影响启动。
"""
archived = file_path + ".migrated"
try:
if os.path.exists(archived):
os.remove(archived)
os.rename(file_path, archived)
_log.info(f"已归档迁移文件: {os.path.basename(file_path)} -> {os.path.basename(archived)}")
except Exception as e:
_log.warning(f"归档 {file_path} 失败(不影响运行): {e}")