fix(security): complete security and delivery compliance remediation

This commit is contained in:
2026-07-28 17:56:28 +08:00
parent 5b07ec6df0
commit 81ab82e9ba
32 changed files with 798 additions and 94 deletions
+42 -19
View File
@@ -1,33 +1,48 @@
"""认证 API"""
import base64
import json
from fastapi import APIRouter, Depends, HTTPException, Header
import hmac
import secrets
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
from fastapi import APIRouter, Depends, HTTPException, Header, Request, Response
from pydantic import BaseModel
from sqlalchemy.orm import Session
from app.core.database import get_db
from app.core.casdoor import casdoor_sdk
from app.core.security import create_access_token, verify_token
from app.core.errors import internal_error
from app.core.security import create_access_token, verify_casdoor_token, verify_token
from app.core.config import settings
from app.models.user import User
from app.schemas.auth import Token, UserInfo
from datetime import datetime
router = APIRouter(prefix="/api/auth", tags=["认证"])
OAUTH_STATE_COOKIE = "h3c_oauth_state"
OAUTH_STATE_TTL_SECONDS = 300
def decode_jwt_payload(token: str) -> dict:
"""直接解码 JWT payload,不验签(Casdoor 已完成认证)"""
payload_b64 = token.split(".")[1]
rem = len(payload_b64) % 4
if rem:
payload_b64 += "=" * (4 - rem)
return json.loads(base64.urlsafe_b64decode(payload_b64))
def _with_oauth_state(url: str, state: str) -> str:
"""Replace the SDK-generated state with the browser-bound state value."""
parts = urlsplit(url)
query = [(key, value) for key, value in parse_qsl(parts.query, keep_blank_values=True) if key != "state"]
query.append(("state", state))
return urlunsplit((parts.scheme, parts.netloc, parts.path, urlencode(query), parts.fragment))
@router.get("/login")
def login():
def login(response: Response):
"""获取 Casdoor 登录 URL"""
return {"url": casdoor_sdk.get_auth_link(settings.CASDOOR_REDIRECT_URL)}
state = secrets.token_urlsafe(32)
response.set_cookie(
key=OAUTH_STATE_COOKIE,
value=state,
max_age=OAUTH_STATE_TTL_SECONDS,
httponly=True,
secure=not settings.DEBUG,
samesite="lax",
path="/api/auth",
)
login_url = casdoor_sdk.get_auth_link(settings.CASDOOR_REDIRECT_URL)
return {"url": _with_oauth_state(login_url, state)}
class CallbackRequest(BaseModel):
@@ -36,18 +51,29 @@ class CallbackRequest(BaseModel):
@router.post("/callback", response_model=Token)
def callback(body: CallbackRequest, db: Session = Depends(get_db)):
def callback(
body: CallbackRequest,
request: Request,
response: Response,
db: Session = Depends(get_db),
):
"""Casdoor 登录回调"""
try:
expected_state = request.cookies.get(OAUTH_STATE_COOKIE)
response.delete_cookie(OAUTH_STATE_COOKIE, path="/api/auth")
if not expected_state or not hmac.compare_digest(body.state, expected_state):
raise HTTPException(status_code=400, detail="登录状态校验失败,请重新登录")
token_response = casdoor_sdk.get_oauth_token(code=body.code)
if isinstance(token_response, dict) and "error" in token_response:
raise HTTPException(status_code=400, detail=token_response.get("error_description", token_response["error"]))
access_token = token_response.get("access_token") if isinstance(token_response, dict) else token_response
identity_token = token_response.get("id_token") if isinstance(token_response, dict) else None
if not access_token:
raise HTTPException(status_code=400, detail="Casdoor 未返回 access_token")
casdoor_user = decode_jwt_payload(access_token)
casdoor_user = verify_casdoor_token(identity_token or access_token)
user = db.query(User).filter(User.casdoor_id == casdoor_user["sub"]).first()
if not user:
@@ -76,10 +102,7 @@ def callback(body: CallbackRequest, db: Session = Depends(get_db)):
except HTTPException:
raise
except Exception as e:
import traceback
import logging
logging.getLogger(__name__).error("callback error: %s\n%s", e, traceback.format_exc())
raise HTTPException(status_code=500, detail=str(e))
raise internal_error("Casdoor login callback", e)
@router.get("/permissions")
+4 -4
View File
@@ -11,6 +11,7 @@ from slowapi.util import get_remote_address
from app.tasks.check_tasks import check_all_devices
from app.core.celery_app import celery_app
from app.core.database import get_db
from app.core.errors import internal_error
from app.services.check_service import CheckService
from app.middleware.permission_middleware import require_permission
@@ -41,8 +42,7 @@ def trigger_check(request: Request, _: dict = Depends(require_permission('device
task = check_all_devices.delay()
return {"task_id": task.id, "status": "started"}
except Exception as e:
logger.error(f"触发状态检查失败: {str(e)}")
raise HTTPException(status_code=500, detail=f"触发状态检查失败: {str(e)}")
raise internal_error("Trigger device status check", e)
@router.get("/status/{task_id}")
@@ -84,7 +84,7 @@ def scan_olt(
result = asyncio.run(service.scan_olt(olt_id))
return result
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
raise internal_error("Scan OLT", e)
@router.post("/discover/{olt_id}")
@@ -99,4 +99,4 @@ def discover_olt(
result = service.scan_and_discover(olt_id)
return result
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
raise internal_error("Discover OLT devices", e)
+5 -4
View File
@@ -8,6 +8,7 @@ from sqlalchemy import asc, desc, distinct, or_
from pydantic import BaseModel
from typing import Optional
from app.core.database import get_db
from app.core.errors import internal_error
from app.middleware.permission_middleware import require_permission
from app.models.device import ONUDevice, DeviceStatusHistory, OLTDevice, DeviceReplacement
from app.schemas.device import DeviceListResponse, ONUDeviceResponse, RebootResponse, OpticalPowerResponse
@@ -439,7 +440,7 @@ def refresh_device_status(
result = service.check_single_device(device_id)
return result
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
raise internal_error("Refresh device status", e)
@@ -641,7 +642,7 @@ def reboot_device(
result = IMCService().reboot_onu(device.mac_address)
return RebootResponse(**result)
except Exception as e:
raise HTTPException(status_code=500, detail=f"重启失败: {str(e)}")
raise internal_error("Reboot ONU", e)
@router.get("/{device_id}/optical-power", response_model=OpticalPowerResponse)
@@ -694,7 +695,7 @@ def get_device_optical_power(
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=f"获取光功率失败: {str(e)}")
raise internal_error("Get optical power", e)
@router.get("/{device_id}/onu-events")
@@ -726,7 +727,7 @@ def get_onu_events(
"events": events,
}
except Exception as e:
raise HTTPException(status_code=500, detail=f"查询失败: {str(e)}")
raise internal_error("Get ONU events", e)
@router.get("/{device_id}/optical-power-history")
+8 -6
View File
@@ -4,9 +4,11 @@ from sqlalchemy.orm import Session
from sqlalchemy import distinct
from pydantic import BaseModel
from app.core.database import get_db
from app.core.errors import internal_error
from app.core.config import settings
from app.middleware.permission_middleware import require_permission
from app.models.device import OLTDevice
from app.schemas.olt import serialize_olt
import pandas as pd
import io
@@ -64,7 +66,7 @@ def get_devices(
q = q.filter(OLTDevice.region.in_(areas))
else:
return []
return q.all()
return [serialize_olt(device) for device in q.all()]
@router.post("/devices")
@@ -148,8 +150,8 @@ async def import_devices(
df = pd.read_excel(io.BytesIO(content))
# 标准化列名
df.columns = [str(c).strip() for c in df.columns]
except Exception as e:
raise HTTPException(status_code=400, detail=f"文件解析失败: {str(e)}")
except Exception:
raise HTTPException(status_code=400, detail="文件解析失败,请确认文件格式")
required_cols = ['IP地址', '用户名', '密码']
missing = [c for c in required_cols if c not in df.columns]
@@ -267,7 +269,7 @@ def clear_onu_port(
with SSHService(olt.ip_address, olt.username, olt.password) as ssh:
ssh.clear_onu_port(body.port_id)
except Exception as e:
raise HTTPException(status_code=500, detail=f"清除失败: {str(e)}")
raise internal_error("Clear ONU port", e)
# 从 ports 列表移除已清除的端口
remaining = [p for p in record.ports if p["port_id"] != body.port_id]
@@ -571,7 +573,7 @@ def get_olt_ports(
ports = ssh.get_olt_ports()
return {"ports": ports}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
raise internal_error("Get OLT ports", e)
@router.post("/devices/{olt_id}/ports/toggle")
@@ -588,5 +590,5 @@ def toggle_olt_port(olt_id: int, body: TogglePortRequest, port_name: str, db: Se
ssh.toggle_olt_port(port_name, body.action)
return {"message": f"端口 {port_name}{'关闭' if body.action == 'shutdown' else '开启'}"}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
raise internal_error("Toggle OLT port", e)