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:
yangy
2026-05-14 17:54:36 +08:00
parent b1d088a755
commit 661b85ce62
23 changed files with 2104 additions and 119 deletions
+9 -5
View File
@@ -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("/")
+54
View File
@@ -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
}
+1
View File
@@ -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
+2 -1
View File
@@ -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"]
+36
View File
@@ -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
View File
@@ -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}")
+195
View File
@@ -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
View File
@@ -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
+64
View File
@@ -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="无效凭证")