feat: 增加使用期限激活限制
This commit is contained in:
@@ -0,0 +1,26 @@
|
||||
from fastapi import APIRouter, Body, HTTPException
|
||||
|
||||
from app.services.license import InvalidActivationCodeError, activate, get_license_info
|
||||
|
||||
|
||||
router = APIRouter(tags=["license"])
|
||||
|
||||
|
||||
@router.get("/api/license/status")
|
||||
async def get_license_status():
|
||||
info = get_license_info()
|
||||
return {"success": True, **info.to_dict()}
|
||||
|
||||
|
||||
@router.post("/api/license/activate")
|
||||
async def activate_license(code: str = Body(..., embed=True)):
|
||||
try:
|
||||
info = activate(code)
|
||||
except InvalidActivationCodeError as exc:
|
||||
raise HTTPException(status_code=400, detail=exc.to_detail()) from exc
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": "激活成功,到期日期已延长 30 天",
|
||||
**info.to_dict(),
|
||||
}
|
||||
@@ -8,6 +8,7 @@ from fastapi import APIRouter, Body, HTTPException
|
||||
from app import state
|
||||
from app.config import AppConfig, CACHE_DIR, RemoteDataConfig
|
||||
from app.processor import DataProcessor, ProcessLogger
|
||||
from app.services.license import LicenseError, LicenseInfo, check_processing_allowed
|
||||
from app.services.remote_download import RemoteDataDownloader
|
||||
|
||||
|
||||
@@ -115,11 +116,13 @@ def _run_remote_processing(
|
||||
raise RuntimeError("远程目录中未下载到任何文件")
|
||||
|
||||
state.history_manager.update(task_id, file_count=download_result.file_count)
|
||||
logger.set_stage("extracting")
|
||||
logger.set_stage("license")
|
||||
_log_license_check(logger, check_processing_allowed(work_dir))
|
||||
|
||||
processor = DataProcessor(app_config, work_dir, logger)
|
||||
result = processor.process()
|
||||
status = "completed" if result.get("success") else "failed"
|
||||
error = result.get("error")
|
||||
if status == "completed" and remote_config.auto_delete_source:
|
||||
try:
|
||||
deleted_count = downloader.delete_source_files(download_result.remote_files)
|
||||
@@ -131,21 +134,25 @@ def _run_remote_processing(
|
||||
task_id,
|
||||
status=status,
|
||||
elapsed_time=result.get("elapsed_time", 0),
|
||||
error=result.get("error"),
|
||||
error=error,
|
||||
result_tables=["4G_结果表", "5G_结果表"],
|
||||
)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": status,
|
||||
"stage": status,
|
||||
"error": error,
|
||||
}
|
||||
except Exception as exc:
|
||||
error_detail = exc.to_detail() if isinstance(exc, LicenseError) else None
|
||||
logger.error(f"远程自动化任务失败: {exc}")
|
||||
state.history_manager.update(task_id, status="failed", error=str(exc))
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": "failed",
|
||||
"stage": "failed",
|
||||
"error": str(exc),
|
||||
"error_detail": error_detail,
|
||||
}
|
||||
finally:
|
||||
try:
|
||||
@@ -162,3 +169,15 @@ def _format_bytes(size: int) -> str:
|
||||
return f"{value:.1f} {unit}"
|
||||
value /= 1024
|
||||
return f"{value:.1f} TB"
|
||||
|
||||
|
||||
def _log_license_check(logger: ProcessLogger, info: LicenseInfo) -> None:
|
||||
if info.current_date:
|
||||
logger.info(
|
||||
f"授权校验通过,数据日期: {info.current_date.isoformat()},"
|
||||
f"到期日期: {info.expires_on.isoformat()}"
|
||||
)
|
||||
elif info.zip_count:
|
||||
logger.warning("未从 ZIP 文件名识别到日期,已跳过授权日期比对")
|
||||
else:
|
||||
logger.warning("未找到 ZIP 文件,已跳过授权日期比对")
|
||||
|
||||
@@ -8,6 +8,7 @@ from fastapi import APIRouter, Body, HTTPException
|
||||
from app import state
|
||||
from app.config import AppConfig
|
||||
from app.processor import DataProcessor, ProcessLogger
|
||||
from app.services.license import LicenseError, LicenseInfo, check_processing_allowed
|
||||
|
||||
|
||||
router = APIRouter(tags=["tasks"])
|
||||
@@ -133,6 +134,8 @@ async def get_processing_status(task_id: str = Body(..., embed=True)):
|
||||
"status": task_info["status"],
|
||||
"stage": task_info.get("stage"),
|
||||
"logs": logs,
|
||||
"error": task_info.get("error"),
|
||||
"error_detail": task_info.get("error_detail"),
|
||||
}
|
||||
|
||||
record = state.history_manager.get(task_id)
|
||||
@@ -170,27 +173,36 @@ def _set_task_stage(task_id: str, stage: str, logs: list[str], status: str = "pr
|
||||
|
||||
def _run_processing(task_id: str, work_dir: Path, logger: ProcessLogger, app_config: AppConfig) -> None:
|
||||
try:
|
||||
logger.set_stage("license")
|
||||
_log_license_check(logger, check_processing_allowed(work_dir))
|
||||
|
||||
processor = DataProcessor(app_config, work_dir, logger)
|
||||
result = processor.process()
|
||||
status = "completed" if result.get("success") else "failed"
|
||||
error = result.get("error")
|
||||
state.history_manager.update(
|
||||
task_id,
|
||||
status=status,
|
||||
elapsed_time=result.get("elapsed_time", 0),
|
||||
error=result.get("error"),
|
||||
error=error,
|
||||
result_tables=["4G_结果表", "5G_结果表"],
|
||||
)
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": status,
|
||||
"stage": status,
|
||||
"error": error,
|
||||
}
|
||||
except Exception as exc:
|
||||
error_detail = exc.to_detail() if isinstance(exc, LicenseError) else None
|
||||
logger.error(str(exc))
|
||||
state.history_manager.update(task_id, status="failed", error=str(exc))
|
||||
state.processing_tasks[task_id] = {
|
||||
"logs": state.history_manager.get_logs(task_id),
|
||||
"status": "failed",
|
||||
"stage": "failed",
|
||||
"error": str(exc),
|
||||
"error_detail": error_detail,
|
||||
}
|
||||
finally:
|
||||
try:
|
||||
@@ -198,3 +210,15 @@ def _run_processing(task_id: str, work_dir: Path, logger: ProcessLogger, app_con
|
||||
except Exception as exc:
|
||||
print(f"自动清理处理历史失败: {exc}")
|
||||
state.reset_task_lock()
|
||||
|
||||
|
||||
def _log_license_check(logger: ProcessLogger, info: LicenseInfo) -> None:
|
||||
if info.current_date:
|
||||
logger.info(
|
||||
f"授权校验通过,数据日期: {info.current_date.isoformat()},"
|
||||
f"到期日期: {info.expires_on.isoformat()}"
|
||||
)
|
||||
elif info.zip_count:
|
||||
logger.warning("未从 ZIP 文件名识别到日期,已跳过授权日期比对")
|
||||
else:
|
||||
logger.warning("未找到 ZIP 文件,已跳过授权日期比对")
|
||||
|
||||
+2
-1
@@ -7,7 +7,7 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse, JSONResponse
|
||||
|
||||
from app import state
|
||||
from app.api.routers import auth, cache, config, database, health, history, remote, script, service, tasks, upload
|
||||
from app.api.routers import auth, cache, config, database, health, history, license, remote, script, service, tasks, upload
|
||||
from app.auth import verify_jwt_token
|
||||
from app.config import BASE_DIR
|
||||
|
||||
@@ -64,6 +64,7 @@ def register_routes(app: FastAPI) -> None:
|
||||
remote.router,
|
||||
tasks.router,
|
||||
history.router,
|
||||
license.router,
|
||||
database.router,
|
||||
config.router,
|
||||
cache.router,
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
"""本地授权期限校验。"""
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from app.config import BASE_DIR
|
||||
|
||||
|
||||
DEFAULT_EXPIRES_ON = date(2026, 6, 20)
|
||||
EXTEND_DAYS = 30
|
||||
LICENSE_FILE = BASE_DIR / "license.dat"
|
||||
_SECRET = b"CapacityReport local license v1"
|
||||
_ZIP_DATE_RE = re.compile(r"(?<!\d)(20\d{10}(?:\d{2})?)(?!\d)")
|
||||
|
||||
|
||||
class LicenseError(Exception):
|
||||
"""授权校验错误。"""
|
||||
|
||||
code = "LICENSE_ERROR"
|
||||
|
||||
def to_detail(self) -> dict[str, Any]:
|
||||
return {"code": self.code, "message": str(self)}
|
||||
|
||||
|
||||
class LicenseExpiredError(LicenseError):
|
||||
"""数据日期超过授权到期日期。"""
|
||||
|
||||
code = "LICENSE_EXPIRED"
|
||||
|
||||
def __init__(self, expires_on: date, current_date: date):
|
||||
self.expires_on = expires_on
|
||||
self.current_date = current_date
|
||||
super().__init__(
|
||||
f"授权已过期:数据日期 {current_date.isoformat()} 已超过到期日期 {expires_on.isoformat()}"
|
||||
)
|
||||
|
||||
def to_detail(self) -> dict[str, Any]:
|
||||
return {
|
||||
"code": self.code,
|
||||
"message": str(self),
|
||||
"expires_on": self.expires_on.isoformat(),
|
||||
"current_date": self.current_date.isoformat(),
|
||||
"key_label": format_key_label(self.expires_on),
|
||||
}
|
||||
|
||||
|
||||
class InvalidActivationCodeError(LicenseError):
|
||||
"""激活码错误。"""
|
||||
|
||||
code = "LICENSE_INVALID"
|
||||
|
||||
def __init__(self, expires_on: date):
|
||||
self.expires_on = expires_on
|
||||
super().__init__("激活码无效,请按当前 key 重新计算后输入")
|
||||
|
||||
def to_detail(self) -> dict[str, Any]:
|
||||
return {
|
||||
"code": self.code,
|
||||
"message": str(self),
|
||||
"expires_on": self.expires_on.isoformat(),
|
||||
"key_label": format_key_label(self.expires_on),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LicenseInfo:
|
||||
expires_on: date
|
||||
current_date: date | None = None
|
||||
zip_count: int = 0
|
||||
|
||||
@property
|
||||
def key_label(self) -> str:
|
||||
return format_key_label(self.expires_on)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"expires_on": self.expires_on.isoformat(),
|
||||
"key_label": self.key_label,
|
||||
"current_date": self.current_date.isoformat() if self.current_date else None,
|
||||
"zip_count": self.zip_count,
|
||||
}
|
||||
|
||||
|
||||
def get_license_info() -> LicenseInfo:
|
||||
return LicenseInfo(expires_on=read_expires_on())
|
||||
|
||||
|
||||
def activate(code: str) -> LicenseInfo:
|
||||
expires_on = read_expires_on()
|
||||
expected = activation_hash(expires_on)
|
||||
normalized_code = (code or "").strip().lower()
|
||||
if not hmac.compare_digest(normalized_code, expected):
|
||||
raise InvalidActivationCodeError(expires_on)
|
||||
|
||||
new_expires_on = expires_on + timedelta(days=EXTEND_DAYS)
|
||||
write_expires_on(new_expires_on)
|
||||
return LicenseInfo(expires_on=new_expires_on)
|
||||
|
||||
|
||||
def check_processing_allowed(work_dir: Path) -> LicenseInfo:
|
||||
expires_on = read_expires_on()
|
||||
zip_count, current_date = extract_max_zip_date(work_dir)
|
||||
info = LicenseInfo(expires_on=expires_on, current_date=current_date, zip_count=zip_count)
|
||||
|
||||
if current_date and current_date > expires_on:
|
||||
raise LicenseExpiredError(expires_on, current_date)
|
||||
|
||||
return info
|
||||
|
||||
|
||||
def extract_max_zip_date(work_dir: Path) -> tuple[int, date | None]:
|
||||
max_date: date | None = None
|
||||
zip_count = 0
|
||||
for zip_file in work_dir.rglob("*.zip"):
|
||||
zip_count += 1
|
||||
for raw_value in _ZIP_DATE_RE.findall(zip_file.name):
|
||||
parsed_date = _parse_zip_timestamp(raw_value)
|
||||
if parsed_date and (max_date is None or parsed_date > max_date):
|
||||
max_date = parsed_date
|
||||
|
||||
return zip_count, max_date
|
||||
|
||||
|
||||
def activation_hash(expires_on: date) -> str:
|
||||
return hashlib.sha256(format_key_label(expires_on).encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def format_key_label(value: date) -> str:
|
||||
return value.strftime("%Y/%m/%d")
|
||||
|
||||
|
||||
def read_expires_on() -> date:
|
||||
if not LICENSE_FILE.exists():
|
||||
write_expires_on(DEFAULT_EXPIRES_ON)
|
||||
return DEFAULT_EXPIRES_ON
|
||||
|
||||
try:
|
||||
encrypted = base64.urlsafe_b64decode(LICENSE_FILE.read_text(encoding="utf-8").encode("ascii"))
|
||||
raw = _xor_bytes(encrypted)
|
||||
data = json.loads(raw.decode("utf-8"))
|
||||
payload = data["payload"]
|
||||
signature = data["signature"]
|
||||
payload_raw = _dump_json(payload)
|
||||
expected_signature = hmac.new(_SECRET, payload_raw, hashlib.sha256).hexdigest()
|
||||
if not hmac.compare_digest(signature, expected_signature):
|
||||
raise ValueError("signature mismatch")
|
||||
|
||||
return date.fromisoformat(str(payload["expires_on"]))
|
||||
except Exception:
|
||||
write_expires_on(DEFAULT_EXPIRES_ON)
|
||||
return DEFAULT_EXPIRES_ON
|
||||
|
||||
|
||||
def write_expires_on(expires_on: date) -> None:
|
||||
payload = {"expires_on": expires_on.isoformat()}
|
||||
payload_raw = _dump_json(payload)
|
||||
data = {
|
||||
"payload": payload,
|
||||
"signature": hmac.new(_SECRET, payload_raw, hashlib.sha256).hexdigest(),
|
||||
}
|
||||
encrypted = _xor_bytes(_dump_json(data))
|
||||
LICENSE_FILE.write_text(base64.urlsafe_b64encode(encrypted).decode("ascii"), encoding="utf-8")
|
||||
|
||||
|
||||
def _parse_zip_timestamp(value: str) -> date | None:
|
||||
fmt = "%Y%m%d%H%M%S" if len(value) == 14 else "%Y%m%d%H%M"
|
||||
try:
|
||||
return datetime.strptime(value, fmt).date()
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def _dump_json(data: dict[str, Any]) -> bytes:
|
||||
return json.dumps(data, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
||||
|
||||
|
||||
def _xor_bytes(data: bytes) -> bytes:
|
||||
output = bytearray()
|
||||
counter = 0
|
||||
while len(output) < len(data):
|
||||
block = hashlib.sha256(_SECRET + counter.to_bytes(4, "big")).digest()
|
||||
output.extend(block)
|
||||
counter += 1
|
||||
return bytes(value ^ key for value, key in zip(data, output))
|
||||
Reference in New Issue
Block a user