Files
auto_control/core/models.py
T
butubb f22263ab45 feat: 数据库连接层改造——库目标由 .env 装配 + 环境防呆(迁 MySQL 第一步)
为把数据库从单文件 SQLite 迁到 MySQL 5.7 铺路。本期不改后端:
DB_HOST 为空时仍走 SQLite,本地开发无感。

- config.py: 新增 DEPLOY_ENV(默认 dev)与 DB_HOST/PORT/USER/PASSWORD/NAME/
  CHARSET/COLLATION、DATABASE_URL、两个逃生阀(DB_ALLOW_ENV_MISMATCH /
  DB_ALLOW_SQLITE_FALLBACK)
- core/db_config.py(新增): URI 组装;按方言分叉的引擎参数(utf8mb4、
  pool_pre_ping、pool_recycle=1800、READ COMMITTED、STRICT_TRANS_TABLES);
  连接探活;app_meta 方言中立读写(MySQL 里 key 是保留字,需反引号)
- 防混库三层: ①库名与环境绑定(dev→auto_control_dev / prod→auto_control)
  ②库标签 app_meta.deployment_env 与 .env 声明比对 ③启动横幅打印当前库
  (生产用 WARNING 级)。不符直接拒绝启动并说明两边分别是什么
- web_server.py: 硬编码 sqlite URI → db_config;配置错在装配期就 exit 2;
  init_db 之后跑库标签校验 + 横幅
- core/models.py: PRAGMA 监听器加 sqlite 类型守卫——它挂在 Engine 基类上,
  MySQL 连接执行 PRAGMA 会直接导致建连失败
- requirements.txt 加 PyMySQL;scripts/start.sh 依赖守卫加 pymysql,
  并在 exec 前打印 DEPLOY_ENV/DB_NAME/DB_HOST
- .env.example 新增「数据库」段;DEVELOPMENT.md §4.1/4.2、DEPLOY.md §2.2 同步

验证: 用 DATABASE_URL 指向 users.db 的一致快照副本跑通主要只读接口
(health/devices/jobs/pool/groups/discovery/summary 全 200,app_meta 读写正常);
DEPLOY_ENV=prod 且无 DB_HOST 时退出码 2;开 DB_ALLOW_SQLITE_FALLBACK 后可回退。
2026-09-13 10:27:14 +08:00

467 lines
20 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
import sqlite3
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。
注意:这个监听器挂在 Engine 基类上,对**所有方言**的连接都会触发;
MySQL 下执行 PRAGMA 会直接报语法错误导致连接失败,所以必须先判类型。
"""
if not isinstance(dbapi_connection, sqlite3.Connection):
return
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):
"""后台用户。
权限模型:
- is_admin=True:管理员,拥有全部权限,不受 perms 限制
- is_admin=False:按 perms 授权(JSON 数组,如 ["tasks","devices"])
- 权限位:tasks=任务管理, devices=设备控制, apks=应用管理, logs=日志查看
- 用户管理本身只有管理员能做(web_server 层强制),不设普通权限位
"""
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)
perms = db.Column(db.Text, default="[]") # JSON 数组:业务权限位
def get_perms(self):
try:
return json.loads(self.perms or "[]")
except Exception:
return []
def set_perms(self, lst):
"""设置权限位。管理员强制为全部权限(内部不区分)。"""
self.perms = json.dumps(list(lst or []), ensure_ascii=False)
def has_perm(self, perm):
"""是否拥有指定权限。管理员恒为 True。"""
if self.is_admin:
return True
return perm in self.get_perms()
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="generic_steps")
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}>"
class Device(db.Model):
"""设备池(本地设备清单,替代 STF 池作为调度数据源)。
serial 即 adb 序列号(IP:5555 或 USB 序列号);enabled=False 不参与调度。
model 为在线时自动采集的型号(如 Redmi 12C),供管理页/监控页区分设备。
"""
serial = db.Column(db.String(120), primary_key=True)
name = db.Column(db.String(80), default="") # 设备名(唯一,人可读标识)
model = db.Column(db.String(120), default="") # 型号(自动采集)
enabled = db.Column(db.Boolean, default=True) # 是否参与调度
note = db.Column(db.Text, default="") # 备注
created_at = db.Column(db.String(20), default="") # 添加时间
# 设备指纹(ro.serialno):识别"同一台物理设备"的稳定标识。
# 网络设备(serial=IP:5555)换 IP 后靠它认领回原记录,名称/分组/任务引用都不丢。
fingerprint = db.Column(db.String(120), default="")
def to_dict(self):
return {"serial": self.serial, "name": self.name or "",
"model": self.model or "", "enabled": bool(self.enabled),
"note": self.note or "", "created_at": self.created_at or "",
"fingerprint": self.fingerprint or ""}
def __repr__(self):
return f"<Device {self.serial}>"
class PendingDevice(db.Model):
"""自动发现待连接池:扫描验证通过的设备,等待用户确认后才入正式池。
与 Device 表的区别:只代表"被扫描到、可连接",不参与任务调度。
"""
serial = db.Column(db.String(120), primary_key=True) # 如 192.168.20.5:5555
source = db.Column(db.String(20), default="") # lan / tailscale
first_seen = db.Column(db.String(20), default="") # 首次发现时间
last_seen = db.Column(db.String(20), default="") # 最近一次扫描仍可见的时间
# 扫描时顺带读取的设备指纹:与设备池中已有记录比对,用于提示
# 「这台其实就是 <名称>(原 IP 变了)」而不是让用户在一堆陌生 IP 里猜
fingerprint = db.Column(db.String(120), default="")
def to_dict(self):
return {"serial": self.serial, "source": self.source or "",
"first_seen": self.first_seen or "", "last_seen": self.last_seen or "",
"fingerprint": self.fingerprint or ""}
# 版本化 schema 迁移:新增结构变更时在此追加 (版本号, 说明, SQL)
# 版本号单调递增,只执行比当前 schema_version 新的迁移。
SCHEMA_MIGRATIONS = [
(1, "用户权限位:user 表新增 perms 列(JSON 数组,默认空=无业务权限,管理员不受限)",
"ALTER TABLE user ADD COLUMN perms TEXT DEFAULT '[]'"),
(2, "设备池:device 表(本地设备清单,替代 STF 池)",
"CREATE TABLE IF NOT EXISTS device ("
"serial VARCHAR(120) PRIMARY KEY,"
"name VARCHAR(80) DEFAULT '',"
"enabled BOOLEAN DEFAULT 1,"
"note TEXT DEFAULT '',"
"created_at VARCHAR(20) DEFAULT '')"),
(3, "设备池:device 表新增 model 列(型号,在线时自动采集)",
"ALTER TABLE device ADD COLUMN model TEXT DEFAULT ''"),
(4, "自动发现:pending_device 待连接池表(扫描发现的设备,用户确认后才入正式池)",
"CREATE TABLE IF NOT EXISTS pending_device ("
"serial VARCHAR(120) PRIMARY KEY,"
"source VARCHAR(20) DEFAULT '',"
"first_seen VARCHAR(20) DEFAULT '',"
"last_seen VARCHAR(20) DEFAULT '')"),
(5, "设备池:device 表新增 fingerprint 列(设备指纹 ro.serialno,换 IP 后认领回原记录)",
"ALTER TABLE device ADD COLUMN fingerprint VARCHAR(120) DEFAULT ''"),
(6, "自动发现:pending_device 表新增 fingerprint 列(扫描时读取,用于提示是已有设备换了 IP)",
"ALTER TABLE pending_device ADD COLUMN fingerprint VARCHAR(120) DEFAULT ''"),
]
# 唯一索引(部分索引:空值不参与唯一约束,兼容历史未命名/未采指纹的老数据)
# 名称唯一 = 设备的人可读标识;指纹唯一 = 一台物理设备在池中只能有一条记录
_UNIQUE_INDEXES = (
("ux_device_name", "CREATE UNIQUE INDEX IF NOT EXISTS ux_device_name "
"ON device(name) WHERE name IS NOT NULL AND name <> ''"),
("ux_device_fingerprint", "CREATE UNIQUE INDEX IF NOT EXISTS ux_device_fingerprint "
"ON device(fingerprint) WHERE fingerprint IS NOT NULL "
"AND fingerprint <> ''"),
)
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:
try:
db.session.execute(text(sql))
except Exception as e:
# 幂等兜底:目标结构已存在也视为该迁移生效。典型场景——
# create_all 已按当前模型把列/表直接建好(如 perms、device.model),
# 而 schema_version 又因历史中断没记录,导致每次启动重复 ALTER 报错。
# duplicate column name 说明列已存在=迁移目标已达成:回滚本次语句后
# 仍记录版本号,一次启动即自愈;其它异常才中止本批迁移。
if "duplicate column name" not in str(e).lower():
raise
db.session.rollback()
_log.info(f"schema 迁移 {version} 目标已存在,跳过: {desc}")
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}")
_ensure_unique_indexes()
except Exception as e:
_log.error(f"schema 迁移失败(不阻塞启动): {e}")
def _ensure_unique_indexes():
"""建设备池的唯一索引(幂等)。
单独抽出来是因为索引不属于某个版本迁移:老库升级后也要补建。
历史数据若存在重复(名称/指纹撞车),建索引会失败——只告警不回滚、
不阻塞启动,由管理页提示用户改名(唯一约束从此刻起对新数据生效)。
"""
for name, sql in _UNIQUE_INDEXES:
try:
db.session.execute(text(sql))
db.session.commit()
except Exception as e:
db.session.rollback()
_log.warning(f"唯一索引 {name} 创建失败(历史数据可能有重复): {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", "generic_steps"),
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}")