feat: 增加使用期限激活限制

This commit is contained in:
2026-05-20 11:27:07 +08:00
parent fcac708814
commit 90023663fd
12 changed files with 508 additions and 22 deletions
+26
View File
@@ -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(),
}
+21 -2
View File
@@ -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 文件,已跳过授权日期比对")
+25 -1
View File
@@ -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
View File
@@ -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,
+190
View File
@@ -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))