Enhance backend functionality and frontend UI
- Updated main.py to include authentication for brand, model, and OTA routers. - Added new OTA schemas in schemas.py for version management. - Enhanced model retrieval with sorting options in models.py. - Improved model update functionality to support multipart/form-data uploads. - Updated frontend layout and styles for a more modern look, including new font integration. - Implemented login route and authentication checks in router/index.js. - Added sorting capabilities in model table and improved file handling in model view. - Updated requirements.txt to include PyJWT for token management.
This commit is contained in:
+9
-5
@@ -1,11 +1,13 @@
|
||||
from fastapi import FastAPI
|
||||
from fastapi import FastAPI, Depends
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
import os
|
||||
import logging
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from database import engine, Base
|
||||
from routes import brands_router, models_router
|
||||
from routes import brands_router, models_router, ota_router
|
||||
from routes.auth import router as auth_router
|
||||
from security import get_current_user
|
||||
|
||||
# Load environment variables
|
||||
load_dotenv()
|
||||
@@ -46,9 +48,11 @@ app.add_middleware(
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# Include routers
|
||||
app.include_router(brands_router)
|
||||
app.include_router(models_router)
|
||||
# Include routers(业务接口需登录;登录接口除外)
|
||||
app.include_router(auth_router)
|
||||
app.include_router(brands_router, dependencies=[Depends(get_current_user)])
|
||||
app.include_router(models_router, dependencies=[Depends(get_current_user)])
|
||||
app.include_router(ota_router, dependencies=[Depends(get_current_user)])
|
||||
|
||||
|
||||
@app.get("/")
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
from sqlalchemy import Column, Integer, String, DateTime, Boolean, SmallInteger, Text
|
||||
from database import Base
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
class Ota(Base):
|
||||
"""OTA 升级版本模型"""
|
||||
__tablename__ = "ota"
|
||||
__table_args__ = {"comment": "OTA 升级版本"}
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True, comment="ID")
|
||||
verCode = Column(Integer, nullable=False, comment="版本号(整数)")
|
||||
verName = Column(String(20), nullable=False, comment="版本名称")
|
||||
url = Column(String(255), nullable=False, comment="升级包 URL")
|
||||
md5 = Column(String(32), nullable=False, comment="升级包 MD5")
|
||||
force = Column(SmallInteger, nullable=False, default=0, comment="是否强升;0-否")
|
||||
desc = Column(String(255), nullable=True, comment="描述")
|
||||
model = Column(String(100), nullable=True, comment="对应的设备型号")
|
||||
hw = Column(Integer, nullable=False, default=0, comment="硬件版本号")
|
||||
target = Column(SmallInteger, nullable=False, default=0, comment="是否定向,1-是,0-否;否表示面向所有用户")
|
||||
beta = Column(SmallInteger, nullable=False, default=0, comment="是否灰度,1-是,0-否")
|
||||
pawVerCode = Column(Integer, nullable=False, default=0, comment="配对版本号")
|
||||
pawVerName = Column(String(20), nullable=False, default='', comment="配对版本名称")
|
||||
pawUrl = Column(String(255), nullable=False, default='', comment="配对版本 URL")
|
||||
pawMd5 = Column(String(32), nullable=False, default='', comment="配对版本 MD5")
|
||||
startTime = Column(DateTime, nullable=True, comment="升级开始时间")
|
||||
endTime = Column(DateTime, nullable=True, comment="升级结束时间")
|
||||
status = Column(SmallInteger, nullable=False, default=1, comment="是否可用")
|
||||
|
||||
def __repr__(self):
|
||||
return f"<Ota(id={self.id}, verCode={self.verCode}, verName='{self.verName}')>"
|
||||
|
||||
def to_dict(self):
|
||||
"""Convert to dictionary"""
|
||||
return {
|
||||
"id": self.id,
|
||||
"verCode": self.verCode,
|
||||
"verName": self.verName,
|
||||
"url": self.url,
|
||||
"md5": self.md5,
|
||||
"force": self.force,
|
||||
"desc": self.desc,
|
||||
"model": self.model,
|
||||
"hw": self.hw,
|
||||
"target": self.target,
|
||||
"beta": self.beta,
|
||||
"pawVerCode": self.pawVerCode,
|
||||
"pawVerName": self.pawVerName,
|
||||
"pawUrl": self.pawUrl,
|
||||
"pawMd5": self.pawMd5,
|
||||
"startTime": self.startTime.isoformat() if self.startTime else None,
|
||||
"endTime": self.endTime.isoformat() if self.endTime else None,
|
||||
"status": self.status
|
||||
}
|
||||
@@ -7,3 +7,4 @@ uvicorn==0.24.0
|
||||
python-dotenv==1.0.0
|
||||
requests==2.31.0
|
||||
python-multipart==0.0.6
|
||||
PyJWT==2.8.0
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from routes.brands import router as brands_router
|
||||
from routes.models import router as models_router
|
||||
from routes.ota import router as ota_router
|
||||
|
||||
__all__ = ["brands_router", "models_router"]
|
||||
__all__ = ["brands_router", "models_router", "ota_router"]
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from response import ApiResponse
|
||||
from security import create_access_token, verify_credentials
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||
|
||||
|
||||
class LoginBody(BaseModel):
|
||||
username: str = Field(..., min_length=1, max_length=64)
|
||||
password: str = Field(..., min_length=1, max_length=128)
|
||||
|
||||
|
||||
@router.post("/login", response_model=ApiResponse)
|
||||
def login(body: LoginBody):
|
||||
"""控制台登录,成功后返回 JWT(有效期 12 小时)"""
|
||||
if not verify_credentials(body.username.strip(), body.password):
|
||||
logger.warning("Login failed for username=%s", body.username)
|
||||
return ApiResponse(code=0, msg="用户名或密码错误", data=None)
|
||||
token = create_access_token()
|
||||
ttl_seconds = 12 * 60 * 60
|
||||
logger.info("User %s logged in", body.username.strip())
|
||||
return ApiResponse(
|
||||
code=1,
|
||||
msg="success",
|
||||
data={
|
||||
"access_token": token,
|
||||
"token_type": "bearer",
|
||||
"expires_in": ttl_seconds,
|
||||
},
|
||||
)
|
||||
+80
-36
@@ -1,10 +1,10 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Body, UploadFile, File, Form
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy import asc, desc
|
||||
from typing import List, Optional
|
||||
import logging
|
||||
from database import get_db
|
||||
from models.model import Model
|
||||
from schemas import ModelCreate, ModelUpdate, ModelResponse
|
||||
from response import ApiResponse, PageData
|
||||
import requests
|
||||
import os
|
||||
@@ -36,18 +36,34 @@ def get_models(
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
|
||||
brand_name: Optional[str] = Query(None, description="按品牌名称模糊查询"),
|
||||
name: Optional[str] = Query(None, description="按型号名称模糊查询"),
|
||||
sort_by: str = Query("id", description="排序字段:id 或 create_at"),
|
||||
sort_order: str = Query("desc", description="排序方向:asc 或 desc"),
|
||||
db: Session = Depends(get_db)
|
||||
):
|
||||
"""获取所有型号列表(支持按品牌名称和型号名称模糊查询)"""
|
||||
"""获取所有型号列表(支持按品牌名称和型号名称模糊查询;支持按 id、create_at 排序)"""
|
||||
try:
|
||||
logger.info(f"Getting models: skip={skip}, limit={limit}, brand_name={brand_name}, name={name}")
|
||||
logger.info(
|
||||
f"Getting models: skip={skip}, limit={limit}, brand_name={brand_name}, name={name}, "
|
||||
f"sort_by={sort_by}, sort_order={sort_order}"
|
||||
)
|
||||
query = db.query(Model)
|
||||
if brand_name:
|
||||
query = query.filter(Model.brand_name.like(f"%{brand_name}%"))
|
||||
if name:
|
||||
query = query.filter(Model.name.like(f"%{name}%"))
|
||||
|
||||
|
||||
total = query.count()
|
||||
|
||||
sort_columns = {"id": Model.id, "create_at": Model.create_at}
|
||||
order_col = sort_columns.get(sort_by, Model.id)
|
||||
order_dir = (sort_order or "desc").lower()
|
||||
if order_dir not in ("asc", "desc"):
|
||||
order_dir = "desc"
|
||||
if order_dir == "desc":
|
||||
query = query.order_by(desc(order_col))
|
||||
else:
|
||||
query = query.order_by(asc(order_col))
|
||||
|
||||
models = query.offset(skip).limit(limit).all()
|
||||
|
||||
logger.info(f"Found {len(models)} models, total={total}")
|
||||
@@ -149,43 +165,71 @@ def create_model(
|
||||
|
||||
|
||||
@router.put("/{model_id}", response_model=ApiResponse)
|
||||
def update_model(model_id: int, model: ModelUpdate, db: Session = Depends(get_db)):
|
||||
"""更新型号"""
|
||||
def update_model(
|
||||
model_id: int,
|
||||
brand_name: Optional[str] = Form(None),
|
||||
name: Optional[str] = Form(None),
|
||||
form: Optional[str] = Form(None),
|
||||
rig: Optional[str] = Form(None),
|
||||
source: Optional[str] = Form(None),
|
||||
eq_key: Optional[str] = Form(None),
|
||||
measurement_file: UploadFile = File(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""更新型号(与 POST 一致,支持 multipart/form-data,可选上传频响文件)"""
|
||||
try:
|
||||
logger.info(f"Updating model: id={model_id}, data={model.dict()}")
|
||||
logger.info(
|
||||
f"Updating model: id={model_id}, brand_name={brand_name}, name={name}, form={form}"
|
||||
)
|
||||
db_model = db.query(Model).filter(Model.id == model_id).first()
|
||||
if not db_model:
|
||||
logger.warning(f"Model not found: id={model_id}")
|
||||
return ApiResponse(code=2, msg="empty", data=None)
|
||||
|
||||
# 如果更新品牌或型号名称,检查是否冲突
|
||||
if model.brand_name or model.name:
|
||||
new_brand_name = model.brand_name or db_model.brand_name
|
||||
new_name = model.name or db_model.name
|
||||
|
||||
if new_brand_name != db_model.brand_name or new_name != db_model.name:
|
||||
existing = db.query(Model).filter(
|
||||
Model.brand_name == new_brand_name,
|
||||
Model.name == new_name
|
||||
).first()
|
||||
if existing:
|
||||
logger.warning(f"Model already exists: brand_name={new_brand_name}, name={new_name}")
|
||||
return ApiResponse(code=0, msg="该品牌下型号名称已存在", data=None)
|
||||
|
||||
# 更新字段
|
||||
if model.brand_name:
|
||||
db_model.brand_name = model.brand_name
|
||||
if model.name:
|
||||
db_model.name = model.name
|
||||
if model.form is not None:
|
||||
db_model.form = model.form
|
||||
if model.rig is not None:
|
||||
db_model.rig = model.rig
|
||||
if model.source is not None:
|
||||
db_model.source = model.source
|
||||
if model.eq_key is not None:
|
||||
db_model.eq_key = model.eq_key
|
||||
|
||||
|
||||
new_brand_name = db_model.brand_name if brand_name is None else brand_name
|
||||
new_name = db_model.name if name is None else name
|
||||
if new_brand_name != db_model.brand_name or new_name != db_model.name:
|
||||
existing = (
|
||||
db.query(Model)
|
||||
.filter(Model.brand_name == new_brand_name, Model.name == new_name)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
logger.warning(
|
||||
f"Model already exists: brand_name={new_brand_name}, name={new_name}"
|
||||
)
|
||||
return ApiResponse(code=0, msg="该品牌下型号名称已存在", data=None)
|
||||
|
||||
eff_source = db_model.source if source is None else source
|
||||
eff_form = db_model.form if form is None else form
|
||||
if measurement_file and measurement_file.filename:
|
||||
logger.info(f"Uploading measurement file: {measurement_file.filename}")
|
||||
file_ext = os.path.splitext(measurement_file.filename)[1].lower()
|
||||
if file_ext not in ALLOWED_EXTENSIONS:
|
||||
logger.error(f"Unsupported file format: {file_ext}")
|
||||
return ApiResponse(code=0, msg=f"不支持的文件格式:{file_ext}", data=None)
|
||||
if not eff_source or not eff_form:
|
||||
return ApiResponse(code=0, msg="上传频响文件需要来源与形式字段", data=None)
|
||||
save_dir = UPLOAD_FOLDER / eff_source / "data" / eff_form
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
file_path = save_dir / measurement_file.filename
|
||||
with open(file_path, "wb") as buffer:
|
||||
shutil.copyfileobj(measurement_file.file, buffer)
|
||||
logger.info(f"File saved: {file_path}")
|
||||
|
||||
if brand_name is not None:
|
||||
db_model.brand_name = brand_name
|
||||
if name is not None:
|
||||
db_model.name = name
|
||||
if form is not None:
|
||||
db_model.form = form
|
||||
if rig is not None:
|
||||
db_model.rig = rig
|
||||
if source is not None:
|
||||
db_model.source = source
|
||||
if eq_key is not None:
|
||||
db_model.eq_key = eq_key
|
||||
|
||||
db.commit()
|
||||
db.refresh(db_model)
|
||||
logger.info(f"Model updated successfully: id={db_model.id}")
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Body
|
||||
from sqlalchemy.orm import Session
|
||||
from typing import List, Optional
|
||||
import logging
|
||||
from database import get_db
|
||||
from models.ota import Ota
|
||||
from schemas import OtaCreate, OtaUpdate, OtaResponse
|
||||
from response import ApiResponse, PageData
|
||||
|
||||
# Configure logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/ota", tags=["ota"])
|
||||
|
||||
|
||||
@router.get("/", response_model=ApiResponse)
|
||||
def get_ota_list(
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
|
||||
verCode: Optional[int] = Query(None, description="按版本号查询"),
|
||||
verName: Optional[str] = Query(None, description="按版本名称模糊查询"),
|
||||
model: Optional[str] = Query(None, description="按设备型号模糊查询"),
|
||||
status: Optional[int] = Query(None, ge=0, le=1, description="按状态查询"),
|
||||
db: Session = Depends(get_db)
|
||||
):
|
||||
"""获取 OTA 版本列表(支持多种筛选条件)"""
|
||||
try:
|
||||
logger.info(f"Getting OTA list: skip={skip}, limit={limit}, verCode={verCode}, verName={verName}, model={model}, status={status}")
|
||||
query = db.query(Ota)
|
||||
|
||||
# 应用筛选条件
|
||||
if verCode is not None:
|
||||
query = query.filter(Ota.verCode == verCode)
|
||||
if verName:
|
||||
query = query.filter(Ota.verName.like(f"%{verName}%"))
|
||||
if model:
|
||||
query = query.filter(Ota.model.like(f"%{model}%"))
|
||||
if status is not None:
|
||||
query = query.filter(Ota.status == status)
|
||||
|
||||
total = query.count()
|
||||
query = query.order_by(Ota.verCode.desc())
|
||||
ota_list = query.offset(skip).limit(limit).all()
|
||||
|
||||
logger.info(f"Found {len(ota_list)} OTA records, total={total}")
|
||||
|
||||
if not ota_list:
|
||||
logger.warning("No OTA records found")
|
||||
return ApiResponse(code=2, msg="empty", data=None)
|
||||
|
||||
# 转换为字典列表
|
||||
ota_data = [ota.to_dict() for ota in ota_list]
|
||||
|
||||
return ApiResponse(code=1, msg="success", data={"items": ota_data, "total": total, "skip": skip, "limit": limit})
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting OTA list: {str(e)}", exc_info=True)
|
||||
return ApiResponse(code=0, msg="error", data=None)
|
||||
|
||||
|
||||
@router.get("/{ota_id}", response_model=ApiResponse)
|
||||
def get_ota(ota_id: int, db: Session = Depends(get_db)):
|
||||
"""获取单个 OTA 版本"""
|
||||
try:
|
||||
logger.info(f"Getting OTA: id={ota_id}")
|
||||
ota = db.query(Ota).filter(Ota.id == ota_id).first()
|
||||
if not ota:
|
||||
logger.warning(f"OTA not found: id={ota_id}")
|
||||
return ApiResponse(code=2, msg="empty", data=None)
|
||||
logger.info(f"OTA found: {ota.to_dict()}")
|
||||
return ApiResponse(code=1, msg="success", data=ota.to_dict())
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting OTA {ota_id}: {str(e)}", exc_info=True)
|
||||
return ApiResponse(code=0, msg="error", data=None)
|
||||
|
||||
|
||||
@router.post("/", response_model=ApiResponse)
|
||||
def create_ota(ota: OtaCreate, db: Session = Depends(get_db)):
|
||||
"""创建新 OTA 版本"""
|
||||
try:
|
||||
logger.info(f"Creating OTA: verCode={ota.verCode}, verName={ota.verName}, model={ota.model}")
|
||||
|
||||
# 检查版本号是否已存在
|
||||
existing = db.query(Ota).filter(
|
||||
Ota.verCode == ota.verCode,
|
||||
Ota.model == ota.model
|
||||
).first()
|
||||
if existing:
|
||||
logger.warning(f"OTA version already exists: verCode={ota.verCode}, model={ota.model}")
|
||||
return ApiResponse(code=0, msg="该版本已存在", data=None)
|
||||
|
||||
db_ota = Ota(**ota.model_dump())
|
||||
db.add(db_ota)
|
||||
db.commit()
|
||||
db.refresh(db_ota)
|
||||
logger.info(f"OTA created successfully: id={db_ota.id}, verCode={db_ota.verCode}")
|
||||
return ApiResponse(code=1, msg="success", data=db_ota.to_dict())
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating OTA: {str(e)}", exc_info=True)
|
||||
db.rollback()
|
||||
return ApiResponse(code=0, msg="error", data=None)
|
||||
|
||||
|
||||
@router.put("/{ota_id}", response_model=ApiResponse)
|
||||
def update_ota(ota_id: int, ota: OtaUpdate, db: Session = Depends(get_db)):
|
||||
"""更新 OTA 版本"""
|
||||
try:
|
||||
logger.info(f"Updating OTA: id={ota_id}, data={ota.model_dump()}")
|
||||
db_ota = db.query(Ota).filter(Ota.id == ota_id).first()
|
||||
if not db_ota:
|
||||
logger.warning(f"OTA not found: id={ota_id}")
|
||||
return ApiResponse(code=2, msg="empty", data=None)
|
||||
|
||||
# 如果更新版本号,检查是否冲突
|
||||
if ota.verCode or ota.model:
|
||||
new_verCode = ota.verCode if ota.verCode is not None else db_ota.verCode
|
||||
new_model = ota.model if ota.model is not None else db_ota.model
|
||||
|
||||
if new_verCode != db_ota.verCode or new_model != db_ota.model:
|
||||
existing = db.query(Ota).filter(
|
||||
Ota.verCode == new_verCode,
|
||||
Ota.model == new_model
|
||||
).first()
|
||||
if existing:
|
||||
logger.warning(f"OTA version already exists: verCode={new_verCode}, model={new_model}")
|
||||
return ApiResponse(code=0, msg="该版本已存在", data=None)
|
||||
|
||||
# 更新字段
|
||||
update_data = ota.model_dump(exclude_unset=True)
|
||||
for field, value in update_data.items():
|
||||
setattr(db_ota, field, value)
|
||||
|
||||
db.commit()
|
||||
db.refresh(db_ota)
|
||||
logger.info(f"OTA updated successfully: id={db_ota.id}")
|
||||
return ApiResponse(code=1, msg="success", data=db_ota.to_dict())
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating OTA {ota_id}: {str(e)}", exc_info=True)
|
||||
db.rollback()
|
||||
return ApiResponse(code=0, msg="error", data=None)
|
||||
|
||||
|
||||
@router.delete("/{ota_id}", response_model=ApiResponse)
|
||||
def delete_ota(ota_id: int, db: Session = Depends(get_db)):
|
||||
"""删除 OTA 版本"""
|
||||
try:
|
||||
logger.info(f"Deleting OTA: id={ota_id}")
|
||||
db_ota = db.query(Ota).filter(Ota.id == ota_id).first()
|
||||
if not db_ota:
|
||||
logger.warning(f"OTA not found: id={ota_id}")
|
||||
return ApiResponse(code=2, msg="empty", data=None)
|
||||
|
||||
db.delete(db_ota)
|
||||
db.commit()
|
||||
logger.info(f"OTA deleted successfully: id={ota_id}")
|
||||
return ApiResponse(code=1, msg="success", data=None)
|
||||
except Exception as e:
|
||||
logger.error(f"Error deleting OTA {ota_id}: {str(e)}", exc_info=True)
|
||||
db.rollback()
|
||||
return ApiResponse(code=0, msg="error", data=None)
|
||||
|
||||
|
||||
@router.get("/latest/check", response_model=ApiResponse)
|
||||
def check_latest_ta(
|
||||
currentVerCode: int = Query(..., description="当前版本号"),
|
||||
model: str = Query(..., description="设备型号"),
|
||||
hw: Optional[int] = Query(None, description="硬件版本号"),
|
||||
db: Session = Depends(get_db)
|
||||
):
|
||||
"""检查是否有可用的 OTA 升级"""
|
||||
try:
|
||||
logger.info(f"Checking latest OTA: currentVerCode={currentVerCode}, model={model}, hw={hw}")
|
||||
|
||||
# 构建查询条件
|
||||
query = db.query(Ota).filter(
|
||||
Ota.status == 1,
|
||||
Ota.verCode > currentVerCode,
|
||||
Ota.model == model
|
||||
)
|
||||
|
||||
# 如果提供了硬件版本号,添加筛选条件
|
||||
if hw is not None:
|
||||
query = query.filter(Ota.hw == hw)
|
||||
|
||||
# 按版本号降序排列,获取最新版本
|
||||
latest_ota = query.order_by(Ota.verCode.desc()).first()
|
||||
|
||||
if not latest_ota:
|
||||
logger.info(f"No available OTA found for model={model}, currentVerCode={currentVerCode}")
|
||||
return ApiResponse(code=2, msg="empty", data=None)
|
||||
|
||||
logger.info(f"Latest OTA found: verCode={latest_ota.verCode}, verName={latest_ota.verName}")
|
||||
return ApiResponse(code=1, msg="success", data=latest_ota.to_dict())
|
||||
except Exception as e:
|
||||
logger.error(f"Error checking latest OTA: {str(e)}", exc_info=True)
|
||||
return ApiResponse(code=0, msg="error", data=None)
|
||||
@@ -52,3 +52,55 @@ class ModelResponse(ModelBase):
|
||||
|
||||
class Config:
|
||||
from_attributes = True
|
||||
|
||||
|
||||
# OTA Schemas
|
||||
class OtaBase(BaseModel):
|
||||
verCode: int = Field(..., description="版本号(整数)")
|
||||
verName: str = Field(..., min_length=1, max_length=20, description="版本名称")
|
||||
url: str = Field(..., max_length=255, description="升级包 URL")
|
||||
md5: str = Field(..., min_length=32, max_length=32, description="升级包 MD5")
|
||||
force: Optional[int] = Field(0, ge=0, le=1, description="是否强升;0-否")
|
||||
desc: Optional[str] = Field(None, max_length=255, description="描述")
|
||||
model: Optional[str] = Field(None, max_length=100, description="对应的设备型号")
|
||||
hw: Optional[int] = Field(0, description="硬件版本号")
|
||||
target: Optional[int] = Field(0, ge=0, le=1, description="是否定向,1-是,0-否")
|
||||
beta: Optional[int] = Field(0, ge=0, le=1, description="是否灰度,1-是,0-否")
|
||||
pawVerCode: Optional[int] = Field(0, description="配对版本号")
|
||||
pawVerName: Optional[str] = Field("", max_length=20, description="配对版本名称")
|
||||
pawUrl: Optional[str] = Field("", max_length=255, description="配对版本 URL")
|
||||
pawMd5: Optional[str] = Field("", min_length=32, max_length=32, description="配对版本 MD5")
|
||||
startTime: Optional[datetime] = Field(None, description="升级开始时间")
|
||||
endTime: Optional[datetime] = Field(None, description="升级结束时间")
|
||||
status: Optional[int] = Field(1, ge=0, le=1, description="是否可用")
|
||||
|
||||
|
||||
class OtaCreate(OtaBase):
|
||||
pass
|
||||
|
||||
|
||||
class OtaUpdate(BaseModel):
|
||||
verCode: Optional[int] = Field(None, description="版本号(整数)")
|
||||
verName: Optional[str] = Field(None, min_length=1, max_length=20, description="版本名称")
|
||||
url: Optional[str] = Field(None, max_length=255, description="升级包 URL")
|
||||
md5: Optional[str] = Field(None, min_length=32, max_length=32, description="升级包 MD5")
|
||||
force: Optional[int] = Field(None, ge=0, le=1, description="是否强升;0-否")
|
||||
desc: Optional[str] = Field(None, max_length=255, description="描述")
|
||||
model: Optional[str] = Field(None, max_length=100, description="对应的设备型号")
|
||||
hw: Optional[int] = Field(None, description="硬件版本号")
|
||||
target: Optional[int] = Field(None, ge=0, le=1, description="是否定向,1-是,0-否")
|
||||
beta: Optional[int] = Field(None, ge=0, le=1, description="是否灰度,1-是,0-否")
|
||||
pawVerCode: Optional[int] = Field(None, description="配对版本号")
|
||||
pawVerName: Optional[str] = Field(None, max_length=20, description="配对版本名称")
|
||||
pawUrl: Optional[str] = Field(None, max_length=255, description="配对版本 URL")
|
||||
pawMd5: Optional[str] = Field(None, min_length=32, max_length=32, description="配对版本 MD5")
|
||||
startTime: Optional[datetime] = Field(None, description="升级开始时间")
|
||||
endTime: Optional[datetime] = Field(None, description="升级结束时间")
|
||||
status: Optional[int] = Field(None, ge=0, le=1, description="是否可用")
|
||||
|
||||
|
||||
class OtaResponse(OtaBase):
|
||||
id: int
|
||||
|
||||
class Config:
|
||||
from_attributes = True
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
import os
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import jwt
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import HTTPException, Security
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
|
||||
load_dotenv()
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
JWT_SECRET = os.getenv("JWT_SECRET", "dev-only-change-me-for-production")
|
||||
JWT_ALGORITHM = "HS256"
|
||||
TOKEN_TTL_HOURS = 12
|
||||
|
||||
ADMIN_USERNAME = os.getenv("DASHBOARD_ADMIN_USERNAME", "admin")
|
||||
ADMIN_PASSWORD = os.getenv("DASHBOARD_ADMIN_PASSWORD", "Eafon123")
|
||||
|
||||
security_bearer = HTTPBearer(auto_error=False)
|
||||
|
||||
|
||||
def create_access_token() -> str:
|
||||
now = datetime.now(timezone.utc)
|
||||
exp = now + timedelta(hours=TOKEN_TTL_HOURS)
|
||||
payload = {
|
||||
"sub": ADMIN_USERNAME,
|
||||
"iat": int(now.timestamp()),
|
||||
"exp": exp,
|
||||
}
|
||||
return jwt.encode(payload, JWT_SECRET, algorithm=JWT_ALGORITHM)
|
||||
|
||||
|
||||
def decode_token(token: str) -> dict:
|
||||
return jwt.decode(token, JWT_SECRET, algorithms=[JWT_ALGORITHM])
|
||||
|
||||
|
||||
def verify_credentials(username: str, password: str) -> bool:
|
||||
if username != ADMIN_USERNAME:
|
||||
return False
|
||||
try:
|
||||
return secrets.compare_digest(password, ADMIN_PASSWORD)
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def get_current_user(
|
||||
credentials: HTTPAuthorizationCredentials | None = Security(security_bearer),
|
||||
) -> str:
|
||||
if credentials is None or (credentials.scheme or "").lower() != "bearer":
|
||||
raise HTTPException(status_code=401, detail="未登录或缺少凭证")
|
||||
token = credentials.credentials
|
||||
try:
|
||||
payload = decode_token(token)
|
||||
sub = payload.get("sub")
|
||||
if not sub:
|
||||
raise HTTPException(status_code=401, detail="无效凭证")
|
||||
return str(sub)
|
||||
except jwt.ExpiredSignatureError:
|
||||
raise HTTPException(status_code=401, detail="登录已过期,请重新登录")
|
||||
except jwt.InvalidTokenError:
|
||||
raise HTTPException(status_code=401, detail="无效凭证")
|
||||
Reference in New Issue
Block a user