@@ -0,0 +1,780 @@
from fastapi import FastAPI , Body , Request , HTTPException , UploadFile , File , Form , Header , Depends
from fastapi . staticfiles import StaticFiles
from pydantic import BaseModel
import os
import sqlite3
from pathlib import Path , PurePosixPath
from dotenv import load_dotenv
from contextlib import asynccontextmanager
import db
import minio_tool
import hashlib
import secrets
import base64
import json
from datetime import datetime , timezone , timedelta
from fastapi . middleware . cors import CORSMiddleware
import socket
import subprocess
import time
import shutil
import tempfile
import urllib . request
from cryptography . hazmat . primitives . serialization import load_pem_private_key
from cryptography . hazmat . primitives . asymmetric import padding
from cryptography . hazmat . primitives import hashes
# 加载环境变量
load_dotenv ( )
SERVER_HOST = os . getenv ( " SERVER_HOST " )
SERVER_PORT = int ( os . getenv ( " SERVER_PORT " ) )
VALID_CLIENT_TOKEN = os . getenv ( " CLIENT_API_TOKEN " )
# 大型 multipart 上传使用内存文件系统暂存,避免与 MinIO 对象双重占用根分区。
UPLOAD_SPOOL_DIR = Path ( os . getenv ( " UPLOAD_SPOOL_DIR " , " /dev/shm/marsco-upload " ) )
UPLOAD_SPOOL_DIR . mkdir ( parents = True , exist_ok = True )
tempfile . tempdir = str ( UPLOAD_SPOOL_DIR )
MINIO_DATA_DIR = Path ( os . getenv ( " MINIO_DATA_DIR " , " minio_data " ) )
UPLOAD_SPACE_RESERVE = 256 * 1024 * 1024
# 后台管理令牌:仅持久化 SHA-256 哈希,不保存明文
ADMIN_TOKEN_HASH_PATH = Path ( os . getenv ( " ADMIN_TOKEN_HASH_PATH " , " admin_token.sha256 " ) )
DEFAULT_ADMIN_TOKEN = os . getenv ( " ADMIN_TOKEN " , " AdminSecret2026 " )
def token_digest ( token : str ) - > str :
return hashlib . sha256 ( token . encode ( " utf-8 " ) ) . hexdigest ( )
def load_admin_token_hash ( ) - > str :
if ADMIN_TOKEN_HASH_PATH . exists ( ) :
stored = ADMIN_TOKEN_HASH_PATH . read_text ( encoding = " utf-8 " ) . strip ( )
if len ( stored ) == 64 :
return stored
return token_digest ( DEFAULT_ADMIN_TOKEN )
ADMIN_TOKEN_HASH = load_admin_token_hash ( )
print ( " 读取到的CLIENT_API_TOKEN: " , repr ( VALID_CLIENT_TOKEN ) )
# ========== 生命周期初始化 ==========
@asynccontextmanager
async def lifespan ( app : FastAPI ) :
print ( " ===== 进程启动,开始初始化数据库 ===== " )
db_file = " mini.db "
sql_file_path = Path ( " ./tables.sql " )
conn = sqlite3 . connect ( db_file , check_same_thread = False )
conn . row_factory = sqlite3 . Row
cur = conn . cursor ( )
if sql_file_path . exists ( ) :
with open ( sql_file_path , " r " , encoding = " utf-8 " ) as f :
full_sql = f . read ( )
cur . executescript ( full_sql )
conn . commit ( )
print ( f " ✅ 成功执行tables.sql,创建数据表与测试数据 " )
else :
print ( f " ❌ 错误:当前目录找不到 tables.sql 文件! " )
conn . close ( )
print ( " ===== 数据库初始化完成 ===== " )
start_minio_if_needed ( )
yield
print ( " 服务进程正常退出 " )
def is_minio_running ( endpoint : str ) - > bool :
try :
if endpoint . startswith ( " http:// " ) or endpoint . startswith ( " https:// " ) :
base_url = endpoint . rstrip ( " / " )
else :
base_url = f " http:// { endpoint . rstrip ( ' / ' ) } "
health_url = base_url + " /minio/health/live "
with urllib . request . urlopen ( health_url , timeout = 2 ) as resp :
return resp . status == 200
except Exception :
return False
def download_minio_binary ( target_path : Path ) - > str | None :
url = " https://dl.min.io/server/minio/release/linux-amd64/minio "
try :
print ( f " MinIO 二进制未找到,尝试下载到 { target_path } " )
target_path . parent . mkdir ( parents = True , exist_ok = True )
with urllib . request . urlopen ( url , timeout = 30 ) as response , open ( target_path , ' wb ' ) as out_file :
out_file . write ( response . read ( ) )
target_path . chmod ( 0o755 )
abs_path = str ( target_path . resolve ( ) )
print ( f " 已下载 MinIO 二进制到 { abs_path } " )
return abs_path
except Exception as err :
print ( f " 自动下载 MinIO 失败: { err } " )
return None
def start_minio_if_needed ( ) :
endpoint = os . getenv ( ' MINIO_ENDPOINT ' )
if not endpoint :
return
if is_minio_running ( endpoint ) :
print ( f " MinIO 已在 { endpoint } 运行,跳过启动 " )
return
minio_cmd = shutil . which ( ' minio ' )
if not minio_cmd :
minio_bin = Path ( os . getenv ( ' MINIO_BIN_PATH ' , ' minio ' ) )
minio_cmd = download_minio_binary ( minio_bin )
if not minio_cmd :
print ( " MinIO 二进制未找到且自动下载失败,无法启动 MinIO " )
return
data_dir = Path ( os . getenv ( ' MINIO_DATA_DIR ' , ' minio_data ' ) )
data_dir . mkdir ( parents = True , exist_ok = True )
print ( f " 尝试自动启动 MinIO: { minio_cmd } server { data_dir } --address ' :9000 ' " )
try :
subprocess . Popen ( [ minio_cmd , ' server ' , str ( data_dir ) , ' --address ' , ' :9000 ' ] , stdout = subprocess . DEVNULL , stderr = subprocess . DEVNULL )
time . sleep ( 5 )
if is_minio_running ( endpoint ) :
print ( " MinIO 已启动成功 " )
else :
print ( " MinIO 启动失败,仍然无法连接 " )
except Exception as err :
print ( f " 启动 MinIO 失败: { err } " )
# Manifest signing key path
MANIFEST_PRIVATE_KEY_PATH = os . getenv ( " MANIFEST_PRIVATE_KEY_PATH " , str ( Path ( __file__ ) . resolve ( ) . parent / " keys " / " manifest_private_key.pem " ) )
# 创建APP绑定生命周期
app = FastAPI ( title = " 最小自动更新服务 " , lifespan = lifespan )
# 本地上传备份文件目录(MinIO 不通时会保存到这里)
app . mount ( " /static " , StaticFiles ( directory = " local_uploads " ) , name = " static " )
# 跨域配置(给admin前端页面用)
app . add_middleware (
CORSMiddleware ,
allow_origins = [ " * " ] ,
allow_credentials = True ,
allow_methods = [ " * " ] ,
allow_headers = [ " * " ] ,
)
# ===================== 全局客户端鉴权中间件 =====================
@app.middleware ( " http " )
async def auth_middleware ( request : Request , call_next ) :
path = request . url . path
# 后台接口、文档接口、静态回退链接全部跳过客户端token校验
skip_paths = ( " /admin " , " /docs " , " /openapi.json " , " /redoc " , " /static " )
if path . startswith ( skip_paths ) :
response = await call_next ( request )
return response
client_token = request . headers . get ( " X-Client-Token " )
print ( " 本次请求携带token: " , repr ( client_token ) )
if client_token != VALID_CLIENT_TOKEN :
raise HTTPException ( status_code = 401 , detail = " 非法客户端,令牌校验失败 " )
response = await call_next ( request )
return response
# ========== 请求体模型 ==========
class CheckUpdateReq ( BaseModel ) :
app_id : str
current_version : str
channel : str
class DownloadUrlReq ( BaseModel ) :
app_id : str
channel : str
version : str
version_id : int
files : list [ str ]
class ManifestReq ( BaseModel ) :
app_id : str
channel : str
version : str
version_id : int
class ReportReq ( BaseModel ) :
device_id : str
from_version : str
to_version : str
result : str
class ChangeAdminTokenReq ( BaseModel ) :
new_token : str
# 统一后台鉴权:全部接口从Header读取token,不再区分表单/头
def admin_auth ( X_Admin_Token : str = Header ( " " ) ) :
if not secrets . compare_digest ( token_digest ( X_Admin_Token ) , ADMIN_TOKEN_HASH ) :
raise HTTPException ( status_code = 403 , detail = " 后台密钥错误,禁止访问 " )
return True
# Manifest signing helpers
def load_manifest_private_key ( ) :
key_path = Path ( MANIFEST_PRIVATE_KEY_PATH )
if not key_path . exists ( ) :
raise FileNotFoundError ( f " Manifest private key not found: { key_path } " )
key_data = key_path . read_bytes ( )
return load_pem_private_key ( key_data , password = None )
def canonical_manifest_bytes ( manifest_obj : dict ) - > bytes :
copy = { k : manifest_obj [ k ] for k in manifest_obj if k != " signature " }
return json . dumps ( copy , sort_keys = True , separators = ( " , " , " : " ) , ensure_ascii = False ) . encode ( " utf-8 " )
def sign_manifest ( manifest_obj : dict ) - > str :
json_text = canonical_manifest_bytes ( manifest_obj )
private_key = load_manifest_private_key ( )
signature = private_key . sign (
json_text ,
padding . PKCS1v15 ( ) ,
hashes . SHA256 ( )
)
return base64 . b64encode ( signature ) . decode ( " ascii " )
def canonical_signed_bytes ( obj : dict ) - > bytes :
copy = { k : obj [ k ] for k in obj if k != " signature " }
return json . dumps ( copy , sort_keys = True , separators = ( " , " , " : " ) , ensure_ascii = False ) . encode ( " utf-8 " )
def sign_policy ( policy_obj : dict ) - > tuple [ str , str ] :
payload = canonical_signed_bytes ( policy_obj )
signature = load_manifest_private_key ( ) . sign ( payload , padding . PKCS1v15 ( ) , hashes . SHA256 ( ) )
return payload . decode ( " utf-8 " ) , base64 . b64encode ( signature ) . decode ( " ascii " )
def policy_row_to_dict ( row ) - > dict :
if not row :
return {
" policy_seq " : 1 , " force_update " : False , " allow_rollback " : False ,
" offline_allowed " : True , " valid_until " : " 2099-12-31T23:59:59Z " ,
" min_supported_version " : " " , " disabled_versions " : [ ] , " message " : " "
}
try :
disabled = json . loads ( row [ " disabled_versions " ] or " [] " )
except ( TypeError , json . JSONDecodeError ) :
disabled = [ ]
return {
" policy_seq " : int ( row [ " policy_seq " ] ) , " force_update " : bool ( row [ " force_update " ] ) ,
" allow_rollback " : bool ( row [ " allow_rollback " ] ) , " offline_allowed " : bool ( row [ " offline_allowed " ] ) ,
" valid_until " : row [ " valid_until " ] , " min_supported_version " : row [ " min_supported_version " ] or " " ,
" disabled_versions " : disabled if isinstance ( disabled , list ) else [ ] , " message " : row [ " message " ] or " "
}
def is_executable_path ( path : str ) - > bool :
return path . lower ( ) . endswith ( ( " .exe " , " .dll " ) )
def version_key ( version : str ) :
parts = [ ]
for part in version . split ( " . " ) :
digits = " " . join ( ch for ch in part if ch . isdigit ( ) )
parts . append ( int ( digits ) if digits else 0 )
return tuple ( parts )
def normalize_relative_path ( raw_path : str ) - > str :
if not isinstance ( raw_path , str ) :
raise HTTPException ( status_code = 400 , detail = " 文件相对路径无效 " )
normalized = raw_path . replace ( chr ( 92 ) , " / " )
if normalized . startswith ( " / " ) :
raise HTTPException ( status_code = 400 , detail = f " 文件路径必须是相对路径: { raw_path } " )
normalized = normalized . strip ( " / " )
if not normalized or chr ( 0 ) in normalized :
raise HTTPException ( status_code = 400 , detail = " 文件相对路径为空或包含非法字符 " )
path = PurePosixPath ( normalized )
if path . is_absolute ( ) or any ( part in ( " " , " . " , " .. " ) for part in path . parts ) :
raise HTTPException ( status_code = 400 , detail = f " 不安全的文件路径: { raw_path } " )
if path . parts and " : " in path . parts [ 0 ] :
raise HTTPException ( status_code = 400 , detail = f " 文件路径不能包含盘符: { raw_path } " )
return path . as_posix ( )
# 文件sha256工具
def calc_sha256 ( data : bytes ) :
sha = hashlib . sha256 ( )
sha . update ( data )
return sha . hexdigest ( )
# ==================== 原有客户端接口 完整保留 ====================
# 1. 查询更新接口
@app.post ( " /api/v1/update/check " )
async def check_update ( body : CheckUpdateReq = Body ( . . . ) ) :
print ( " 收到客户端版本检测请求,参数: " , body . model_dump ( ) )
app_id , cur_ver , channel = body . app_id , body . current_version , body . channel
conn = db . get_conn ( )
ver = conn . execute (
" SELECT * FROM versions WHERE app_id=? AND channel=? AND latest=1 " , ( app_id , channel )
) . fetchone ( )
policy_row = conn . execute (
" SELECT * FROM version_policies WHERE app_id=? AND channel=? " , ( app_id , channel )
) . fetchone ( )
conn . close ( )
settings = policy_row_to_dict ( policy_row )
latest_ver = ver [ " version " ] if ver else " "
version_id = int ( ver [ " id " ] ) if ver else 0
upgrade_available = bool ( latest_ver ) and version_key ( latest_ver ) > version_key ( cur_ver )
rollback_candidate = bool ( latest_ver ) and version_key ( latest_ver ) < version_key ( cur_ver )
rollback_allowed = rollback_candidate and bool ( settings [ " allow_rollback " ] )
need_update = upgrade_available or rollback_allowed
disabled = cur_ver in settings [ " disabled_versions " ]
below_min = bool ( settings [ " min_supported_version " ] ) and version_key ( cur_ver ) < version_key ( settings [ " min_supported_version " ] )
allow_run = not disabled and not below_min
force_update = bool ( settings [ " force_update " ] and upgrade_available ) or disabled or below_min
if disabled :
action , message = " blocked " , settings [ " message " ] or f " 当前版本 { cur_ver } 已被管理员禁用 "
elif below_min :
action , message = " force_update " , settings [ " message " ] or f " 当前版本低于最低支持版本 { settings [ ' min_supported_version ' ] } "
elif force_update :
action , message = " force_update " , settings [ " message " ] or " 必须升级到最新版本后才能继续使用 "
elif rollback_allowed :
action , message = " rollback_allowed " , settings [ " message " ] or f " 管理员要求回退到版本 { latest_ver } "
elif rollback_candidate :
action , message = " rollback_denied " , settings [ " message " ] or f " 渠道目标版本 { latest_ver } 低于当前版本,策略禁止降级 "
elif upgrade_available :
action , message = " optional_update " , settings [ " message " ] or " 发现可用新版本 "
else :
action , message = " allow " , settings [ " message " ] or " 当前版本允许使用 "
issued_at = datetime . now ( timezone . utc ) . replace ( microsecond = 0 ) . isoformat ( ) . replace ( " +00:00 " , " Z " )
policy = {
" app_id " : app_id , " channel " : channel , " current_version " : cur_ver ,
" policy_seq " : settings [ " policy_seq " ] , " allow_run " : allow_run ,
" force_update " : force_update , " allow_rollback " : settings [ " allow_rollback " ] ,
" rollback_targets " : [ ] , " offline_allowed " : settings [ " offline_allowed " ] ,
" issued_at " : issued_at , " valid_until " : settings [ " valid_until " ] ,
" latest_version " : latest_ver , " min_supported_version " : settings [ " min_supported_version " ] ,
" disabled_versions " : settings [ " disabled_versions " ] , " message " : message ,
" signature_alg " : " RSA-2048-SHA256 " , " key_id " : " manifest-key-v1 "
}
policy_text , policy_signature = sign_policy ( policy )
policy [ " signature " ] = policy_signature
return {
" allow_run " : allow_run , " action " : action , " message " : message ,
" need_update " : need_update , " release_available " : bool ( ver ) ,
" current_version " : cur_ver , " latest_version " : latest_ver , " version_id " : version_id ,
" force_update " : force_update , " allow_rollback " : settings [ " allow_rollback " ] ,
" rollback " : rollback_allowed ,
" policy_seq " : settings [ " policy_seq " ] , " policy_valid_until " : settings [ " valid_until " ] ,
" policy " : policy , " policy_text " : policy_text , " policy_signature " : policy_signature
}
# 2. 获取文件下载链接
@app.post ( " /api/v1/update/download-url " )
async def get_download_url ( body : DownloadUrlReq = Body ( . . . ) ) :
ver_id = body . version_id
conn = db . get_conn ( )
rows = conn . execute (
" SELECT path, sha256, size FROM version_files WHERE version_id=? " ,
( ver_id , )
) . fetchall ( )
conn . close ( )
res = [ ]
base_path = f " { body . app_id } / { body . channel } / { body . version } /files/ "
for r in rows :
full_path = base_path + r [ " path " ]
url = minio_tool . get_url ( full_path )
res . append ( {
" path " : r [ " path " ] ,
" url " : url ,
" sha256 " : r [ " sha256 " ] ,
" size " : r [ " size " ]
} )
return { " files " : res }
# 2b. 获取版本 manifest
@app.post ( " /api/v1/update/manifest " )
async def get_manifest ( body : ManifestReq = Body ( . . . ) ) :
ver_id = body . version_id
conn = db . get_conn ( )
ver = conn . execute (
" SELECT app_id, channel, version, create_time FROM versions WHERE id=? " ,
( ver_id , )
) . fetchone ( )
if not ver :
conn . close ( )
raise HTTPException ( status_code = 404 , detail = " Version not found " )
rows = conn . execute (
" SELECT path, sha256, size FROM version_files WHERE version_id=? " ,
( ver_id , )
) . fetchall ( )
conn . close ( )
files = [ ]
for r in rows :
files . append ( {
" path " : r [ " path " ] ,
" sha256 " : r [ " sha256 " ] ,
" size " : r [ " size " ] ,
" executable " : is_executable_path ( r [ " path " ] )
} )
manifest = {
" app_id " : ver [ " app_id " ] ,
" version " : ver [ " version " ] ,
" channel " : ver [ " channel " ] ,
" platform " : " windows " ,
" arch " : " x64 " ,
" manifest_seq " : int ( ver_id ) ,
" created_at " : ver [ " create_time " ] ,
" files " : files
}
manifest_text = canonical_manifest_bytes ( manifest ) . decode ( " utf-8 " )
signature = sign_manifest ( manifest )
manifest [ " signature " ] = signature
return {
" manifest " : manifest ,
" manifest_text " : manifest_text ,
" signature " : signature
}
# 3. 升级结果上报
@app.post ( " /api/v1/update/report " )
async def report ( body : ReportReq = Body ( . . . ) ) :
conn = db . get_conn ( )
insert_sql = """
INSERT INTO upgrade_logs(device_id,from_ver,to_ver,result,create_time)
VALUES (?,?,?,?,datetime( ' now ' ))
"""
conn . execute ( insert_sql , ( body . device_id , body . from_version , body . to_version , body . result ) )
conn . commit ( )
conn . close ( )
return { " code " : 0 , " msg " : " 上报成功 " }
# ==================== 管理后台 Admin 全套接口(统一Header鉴权) ====================
@app.post ( " /admin/token/change " )
def admin_change_token ( body : ChangeAdminTokenReq , auth = Depends ( admin_auth ) ) :
global ADMIN_TOKEN_HASH
new_token = body . new_token . strip ( )
if len ( new_token ) < 8 :
raise HTTPException ( status_code = 400 , detail = " 新令牌至少需要 8 个字符 " )
if len ( new_token ) > 128 :
raise HTTPException ( status_code = 400 , detail = " 新令牌不能超过 128 个字符 " )
new_hash = token_digest ( new_token )
ADMIN_TOKEN_HASH_PATH . write_text ( new_hash , encoding = " utf-8 " )
ADMIN_TOKEN_HASH = new_hash
return { " code " : 0 , " msg " : " 管理员令牌已更新 " }
# 应用列表 GET
@app.get ( " /admin/app/list " )
def admin_get_app_list ( auth = Depends ( admin_auth ) ) :
conn = db . get_conn ( )
rows = conn . execute ( " SELECT app_id, app_name FROM apps " ) . fetchall ( )
conn . close ( )
return { " list " : [ { " app_id " : r [ " app_id " ] , " app_name " : r [ " app_name " ] } for r in rows ] }
# 新增应用 POST json
@app.post ( " /admin/app/add " )
def admin_add_app ( body : dict , auth = Depends ( admin_auth ) ) :
aid = body [ " app_id " ]
aname = body [ " app_name " ]
conn = db . get_conn ( )
try :
conn . execute ( " INSERT INTO apps(app_id, app_name) VALUES (?, ?) " , ( aid , aname ) )
conn . commit ( )
except Exception as e :
return { " msg " : f " 创建失败: { str ( e ) } " }
conn . close ( )
return { " msg " : " 应用创建成功 " }
# 查询和保存应用渠道策略
@app.get ( " /admin/policy " )
def admin_get_policy ( app_id : str , channel : str = " stable " , auth = Depends ( admin_auth ) ) :
conn = db . get_conn ( )
row = conn . execute ( " SELECT * FROM version_policies WHERE app_id=? AND channel=? " , ( app_id , channel ) ) . fetchone ( )
conn . close ( )
result = policy_row_to_dict ( row )
result . update ( { " app_id " : app_id , " channel " : channel } )
return result
@app.post ( " /admin/policy/save " )
def admin_save_policy ( body : dict , auth = Depends ( admin_auth ) ) :
app_id = str ( body . get ( " app_id " ) or " " ) . strip ( )
channel = str ( body . get ( " channel " ) or " stable " ) . strip ( )
if not app_id or channel not in ( " stable " , " preview " , " dev " ) :
raise HTTPException ( status_code = 400 , detail = " App ID 或渠道无效 " )
disabled = body . get ( " disabled_versions " ) or [ ]
if not isinstance ( disabled , list ) or any ( not isinstance ( v , str ) for v in disabled ) :
raise HTTPException ( status_code = 400 , detail = " disabled_versions 必须是版本字符串数组 " )
valid_until = str ( body . get ( " valid_until " ) or " " ) . strip ( )
try :
datetime . fromisoformat ( valid_until . replace ( " Z " , " +00:00 " ) )
except ValueError :
raise HTTPException ( status_code = 400 , detail = " valid_until 必须是 ISO-8601 时间 " )
conn = db . get_conn ( )
row = conn . execute ( " SELECT policy_seq FROM version_policies WHERE app_id=? AND channel=? " , ( app_id , channel ) ) . fetchone ( )
next_seq = ( int ( row [ " policy_seq " ] ) + 1 ) if row else 1
values = ( next_seq , int ( bool ( body . get ( " force_update " ) ) ) , int ( bool ( body . get ( " allow_rollback " ) ) ) ,
int ( bool ( body . get ( " offline_allowed " ) ) ) , valid_until ,
str ( body . get ( " min_supported_version " ) or " " ) . strip ( ) ,
json . dumps ( disabled , ensure_ascii = False ) , str ( body . get ( " message " ) or " " ) . strip ( ) , app_id , channel )
conn . execute ( """
INSERT INTO version_policies(policy_seq,force_update,allow_rollback,offline_allowed,valid_until,
min_supported_version,disabled_versions,message,app_id,channel)
VALUES(?,?,?,?,?,?,?,?,?,?)
ON CONFLICT(app_id,channel) DO UPDATE SET policy_seq=excluded.policy_seq,
force_update=excluded.force_update,allow_rollback=excluded.allow_rollback,
offline_allowed=excluded.offline_allowed,valid_until=excluded.valid_until,
min_supported_version=excluded.min_supported_version,disabled_versions=excluded.disabled_versions,
message=excluded.message,updated_at=datetime( ' now ' )
""" , values )
conn . commit ( ) ; conn . close ( )
return { " msg " : " 策略已保存并递增 policy_seq " , " policy_seq " : next_seq }
# 发布新版本(文件上传)
@app.post ( " /admin/publish " )
async def admin_publish_version ( request : Request , auth = Depends ( admin_auth ) ) :
# 在 Starlette 解析 multipart 前检查空间;否则大文件 rollover 会直接抛出 Errno 28。
content_length_header = request . headers . get ( " content-length " )
try :
content_length = int ( content_length_header or 0 )
except ValueError :
raise HTTPException ( status_code = 400 , detail = " 无效的 Content-Length " )
if content_length < = 0 :
raise HTTPException ( status_code = 411 , detail = " 发布请求必须提供 Content-Length " )
spool_free = shutil . disk_usage ( UPLOAD_SPOOL_DIR ) . free
storage_free = shutil . disk_usage ( MINIO_DATA_DIR ) . free
if spool_free < content_length + 64 * 1024 * 1024 :
raise HTTPException ( status_code = 507 , detail = {
" error " : " upload_temp_space_insufficient " ,
" msg " : " 上传临时空间不足,请减小发布包或扩容 /dev/shm " ,
" required_bytes " : content_length + 64 * 1024 * 1024 ,
" available_bytes " : spool_free ,
} )
if storage_free < content_length + UPLOAD_SPACE_RESERVE :
raise HTTPException ( status_code = 507 , detail = {
" error " : " storage_space_insufficient " ,
" msg " : " 版本存储空间不足,请扩容服务器根卷或删除不再需要的历史版本 " ,
" required_bytes " : content_length + UPLOAD_SPACE_RESERVE ,
" available_bytes " : storage_free ,
} )
# 手动解析表单,增加兼容性并便于调试前端上传问题
try :
form = await request . form ( )
except OSError as err :
if getattr ( err , " errno " , None ) == 28 :
raise HTTPException ( status_code = 507 , detail = " 上传临时空间已满,请扩容后重试 " )
raise
# 尝试从 form 中获取字段
app_id = form . get ( " app_id " )
channel = form . get ( " channel " ) or " stable "
version = form . get ( " version " )
# 收集 files(可能是多个同名字段)
files = [ ]
if hasattr ( form , " getlist " ) :
files = form . getlist ( " files " )
else :
# 兼容性降级:遍历所有表单项
for k , v in form . multi_items ( ) :
if k == " files " :
files . append ( v )
# 如果关键信息缺失,返回详细调试信息,方便前端定位问题
if not app_id or not version :
headers = { k : v for k , v in request . headers . items ( ) }
detail = {
" error " : " missing form fields " ,
" have_app_id " : bool ( app_id ) ,
" have_version " : bool ( version ) ,
" content_type " : request . headers . get ( " content-type " ) ,
" headers " : headers ,
" form_keys " : list ( form . keys ( ) ) ,
" app_id_value " : app_id ,
" version_value " : version ,
" app_id_len " : len ( app_id or " " ) ,
" version_len " : len ( version or " " )
}
raise HTTPException ( status_code = 422 , detail = detail )
if not files :
raise HTTPException ( status_code = 400 , detail = " 未选择任何文件,请选择程序文件夹 " )
relative_paths = form . getlist ( " relative_paths " ) if hasattr ( form , " getlist " ) else [ ]
if relative_paths and len ( relative_paths ) != len ( files ) :
raise HTTPException ( status_code = 400 , detail = " 文件数量与相对路径数量不一致 " )
upload_items = [ ]
seen_paths = set ( )
for index , file in enumerate ( files ) :
if not hasattr ( file , " filename " ) :
raise HTTPException ( status_code = 400 , detail = " 上传内容中包含无效文件 " )
raw_path = relative_paths [ index ] if relative_paths else file . filename
rel_path = normalize_relative_path ( str ( raw_path ) )
path_key = rel_path . lower ( )
if path_key in seen_paths :
raise HTTPException ( status_code = 400 , detail = f " 存在重复文件路径: { rel_path } " )
seen_paths . add ( path_key )
upload_items . append ( ( file , rel_path ) )
if " mainapp.exe " not in seen_paths :
nested_main = next ( ( path for path in seen_paths if path . endswith ( " /mainapp.exe " ) ) , None )
if nested_main :
raise HTTPException ( status_code = 400 , detail = f " MainApp.exe 不在发布根级,请改为选择其所在目录: { nested_main } " )
raise HTTPException ( status_code = 400 , detail = " 发布根目录中缺少 MainApp.exe " )
conn = db . get_conn ( )
publish_started = False
publish_prefix = f " { app_id } / { channel } / { version } / "
try :
cur = conn . cursor ( )
exists = cur . execute (
" SELECT 1 FROM versions WHERE app_id=? AND channel=? AND version=? " ,
( app_id , channel , version )
) . fetchone ( )
if exists :
raise HTTPException ( status_code = 400 , detail = " 该版本已存在,请更换版本号或先删除旧版本 " )
cur . execute ( " UPDATE versions SET latest=0 WHERE app_id=? AND channel=? " , ( app_id , channel ) )
cur . execute (
" INSERT INTO versions(app_id, channel, version, latest) VALUES (?,?,?,1) " ,
( app_id , channel , version )
)
new_vid = cur . lastrowid
publish_started = True
for file , rel_path in upload_items :
content = await file . read ( )
f_size = len ( content )
f_sha = calc_sha256 ( content )
obj_path = f " { app_id } / { channel } / { version } /files/ { rel_path } "
put_result = minio_tool . put_file ( obj_path , content , f_size )
if put_result . get ( ' storage ' ) == ' local ' :
print ( f " MinIO 不可用,已保存到本地回退: { put_result . get ( ' path ' ) } " )
cur . execute (
" INSERT INTO version_files(version_id, path, sha256, size) VALUES (?,?,?,?) " ,
( new_vid , rel_path , f_sha , f_size )
)
conn . commit ( )
except sqlite3 . IntegrityError as err :
conn . rollback ( )
if publish_started :
minio_tool . remove_prefix ( publish_prefix )
if ' UNIQUE constraint failed: versions.app_id, versions.channel, versions.version ' in str ( err ) :
raise HTTPException ( status_code = 400 , detail = " 该版本已存在,请更换版本号或先删除旧版本 " )
raise HTTPException ( status_code = 500 , detail = str ( err ) )
except sqlite3 . OperationalError as err :
conn . rollback ( )
if publish_started :
minio_tool . remove_prefix ( publish_prefix )
raise HTTPException ( status_code = 500 , detail = { " error " : " database locked " , " msg " : str ( err ) } )
except HTTPException :
conn . rollback ( )
if publish_started :
minio_tool . remove_prefix ( publish_prefix )
raise
except Exception as err :
conn . rollback ( )
if publish_started :
minio_tool . remove_prefix ( publish_prefix )
if isinstance ( err , OSError ) and getattr ( err , " errno " , None ) == 28 :
raise HTTPException ( status_code = 507 , detail = " 发布过程中存储空间耗尽,已清理未完成版本 " )
raise HTTPException ( status_code = 500 , detail = { " error " : " publish_failed " , " msg " : str ( err ) } )
finally :
conn . close ( )
return { " msg " : f " 版本 { version } 发布完成,共上传 { len ( upload_items ) } 个文件 " }
# 获取版本列表 GET
@app.get ( " /admin/version/list " )
def admin_get_version_list ( app_id : str , auth = Depends ( admin_auth ) ) :
conn = db . get_conn ( )
rows = conn . execute ( """
SELECT id, version, channel, latest, create_time
FROM versions WHERE app_id=? ORDER BY create_time DESC
""" , ( app_id , ) ) . fetchall ( )
conn . close ( )
out = [ ]
for r in rows :
out . append ( {
" id " : r [ " id " ] ,
" version " : r [ " version " ] ,
" channel " : r [ " channel " ] ,
" latest " : bool ( r [ " latest " ] ) ,
" create_time " : r [ " create_time " ]
} )
return { " list " : out }
# 设置指定版本为渠道最新 POST json
@app.post ( " /admin/version/set-latest " )
def admin_set_latest ( body : dict , auth = Depends ( admin_auth ) ) :
vid = body [ " version_id " ]
conn = db . get_conn ( )
cur = conn . cursor ( )
v_info = cur . execute ( " SELECT app_id, channel FROM versions WHERE id=? " , ( vid , ) ) . fetchone ( )
aid , ch = v_info [ " app_id " ] , v_info [ " channel " ]
cur . execute ( " UPDATE versions SET latest=0 WHERE app_id=? AND channel=? " , ( aid , ch ) )
cur . execute ( " UPDATE versions SET latest=1 WHERE id=? " , ( vid , ) )
conn . commit ( )
conn . close ( )
return { " msg " : " 已切换为渠道最新版本 " }
# 删除版本 POST json
@app.post ( " /admin/version/delete " )
def admin_delete_version ( body : dict , auth = Depends ( admin_auth ) ) :
vid = body [ " version_id " ]
conn = db . get_conn ( )
cur = conn . cursor ( )
v_info = cur . execute ( " SELECT app_id, channel, version, latest FROM versions WHERE id=? " , ( vid , ) ) . fetchone ( )
if not v_info :
conn . close ( )
raise HTTPException ( status_code = 404 , detail = " 版本不存在 " )
aid , ch , ver = v_info [ " app_id " ] , v_info [ " channel " ] , v_info [ " version " ]
was_latest = bool ( v_info [ " latest " ] )
# 删除minio下该版本全部文件
prefix = f " { aid } / { ch } / { ver } /files/ "
minio_tool . remove_prefix ( prefix )
# 删除数据库记录
cur . execute ( " DELETE FROM version_files WHERE version_id=? " , ( vid , ) )
cur . execute ( " DELETE FROM versions WHERE id=? " , ( vid , ) )
promoted_version = None
if was_latest :
replacement = cur . execute (
" SELECT id, version FROM versions WHERE app_id=? AND channel=? ORDER BY create_time DESC, id DESC LIMIT 1 " ,
( aid , ch )
) . fetchone ( )
if replacement :
cur . execute ( " UPDATE versions SET latest=1 WHERE id=? " , ( replacement [ " id " ] , ) )
promoted_version = replacement [ " version " ]
conn . commit ( )
conn . close ( )
return {
" msg " : " 版本、云端文件、数据库记录全部删除完成 " ,
" promoted_latest " : promoted_version
}
# 获取升级日志 GET
@app.get ( " /admin/report/list " )
def admin_get_report_log ( auth = Depends ( admin_auth ) ) :
conn = db . get_conn ( )
rows = conn . execute ( """
SELECT device_id,from_ver,to_ver,result,create_time
FROM upgrade_logs ORDER BY create_time DESC LIMIT 200
""" ) . fetchall ( )
conn . close ( )
res = [ ]
for r in rows :
res . append ( {
" device_id " : r [ " device_id " ] ,
" app_id " : " " ,
" from_version " : r [ " from_ver " ] ,
" to_version " : r [ " to_ver " ] ,
" result " : r [ " result " ] ,
" create_time " : r [ " create_time " ]
} )
return { " list " : res }
if __name__ == " __main__ " :
import uvicorn
uvicorn . run ( app , host = SERVER_HOST , port = SERVER_PORT , workers = 1 )