Files
CapacityReport/app/services/api_tokens.py
T

330 lines
9.8 KiB
Python

import base64
import hashlib
import hmac
import json
import secrets
from dataclasses import asdict, dataclass
from datetime import date, datetime, time, timedelta, timezone
from threading import RLock
from typing import Any
from app.config import BASE_DIR
API_TOKENS_PATH = BASE_DIR / "api_tokens.json"
API_TOKEN_PREFIX = "cap_"
API_TOKEN_SECRET = "CapaReportApiTokenSecret2026"
_STORE_LOCK = RLock()
@dataclass
class ApiTokenRecord:
id: str
name: str
token_hash: str
prefix: str
suffix: str
created_at: str
expires_at: str | None
enabled: bool
last_used_at: str | None = None
last_used_from: str | None = None
token: str | None = None
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def ensure_store() -> None:
if API_TOKENS_PATH.exists():
return
with _STORE_LOCK:
if API_TOKENS_PATH.exists():
return
API_TOKENS_PATH.write_text(json.dumps({"tokens": []}, ensure_ascii=False, indent=2), encoding="utf-8")
def list_tokens() -> list[dict[str, Any]]:
with _STORE_LOCK:
return [record_to_public_dict(record) for record in _load_records()]
def export_tokens() -> list[dict[str, Any]]:
with _STORE_LOCK:
return [record.to_dict() for record in _load_records()]
def import_tokens(items: Any) -> int:
if not isinstance(items, list):
return 0
records: list[ApiTokenRecord] = []
for item in items:
if not isinstance(item, dict):
continue
record = _record_from_dict(item)
if record.token_hash:
records.append(record)
with _STORE_LOCK:
_save_records(records)
return len(records)
def create_token(
name: str,
expires_in_days: int | None = None,
enabled: bool = True,
expires_at: str | None = None,
) -> tuple[str, dict[str, Any]]:
raw_token = generate_raw_token()
resolved_expires_at = normalize_expiration(expires_at)
if resolved_expires_at is None and expires_in_days is not None:
resolved_expires_at = expires_at_from_days(expires_in_days)
record = ApiTokenRecord(
id=secrets.token_hex(8),
name=name.strip() or "未命名 Token",
token_hash=hash_token(raw_token),
prefix=raw_token[:12],
suffix=raw_token[-12:],
created_at=utc_now(),
expires_at=resolved_expires_at,
enabled=bool(enabled),
token=raw_token,
)
with _STORE_LOCK:
records = _load_records()
records.append(record)
_save_records(records)
return raw_token, record_to_public_dict(record)
def update_token(token_id: str, **changes: Any) -> dict[str, Any]:
with _STORE_LOCK:
records = _load_records()
for index, record in enumerate(records):
if record.id != token_id:
continue
if "name" in changes and isinstance(changes["name"], str):
record.name = changes["name"].strip() or record.name
if "enabled" in changes:
record.enabled = bool(changes["enabled"])
if "expires_at" in changes:
record.expires_at = normalize_expiration(changes["expires_at"])
records[index] = record
_save_records(records)
return record_to_public_dict(record)
raise KeyError(f"Token not found: {token_id}")
def delete_token(token_id: str) -> None:
with _STORE_LOCK:
records = [record for record in _load_records() if record.id != token_id]
_save_records(records)
def delete_tokens(token_ids: list[str]) -> int:
token_id_set = {str(token_id).strip() for token_id in token_ids if str(token_id).strip()}
if not token_id_set:
return 0
with _STORE_LOCK:
records = _load_records()
kept_records = [record for record in records if record.id not in token_id_set]
_save_records(kept_records)
return len(records) - len(kept_records)
def regenerate_token(token_id: str) -> tuple[str, dict[str, Any]]:
with _STORE_LOCK:
records = _load_records()
for index, record in enumerate(records):
if record.id != token_id:
continue
raw_token = generate_raw_token()
record.token_hash = hash_token(raw_token)
record.prefix = raw_token[:12]
record.suffix = raw_token[-12:]
record.token = raw_token
record.created_at = utc_now()
record.last_used_at = None
record.last_used_from = None
records[index] = record
_save_records(records)
return raw_token, record_to_public_dict(record)
raise KeyError(f"Token not found: {token_id}")
def verify_api_token(raw_token: str) -> dict[str, Any] | None:
ensure_store()
token_hash = hash_token(raw_token)
now = datetime.now(timezone.utc)
with _STORE_LOCK:
for record in _load_records():
if not record.enabled or record.token_hash != token_hash:
continue
expires_at = parse_datetime(record.expires_at)
if expires_at and expires_at < now:
continue
return record_to_context(record)
return None
def touch_token_usage(raw_token: str, source: str | None = None) -> None:
token_hash = hash_token(raw_token)
now = utc_now()
with _STORE_LOCK:
records = _load_records()
updated = False
for index, record in enumerate(records):
if record.token_hash != token_hash:
continue
record.last_used_at = now
record.last_used_from = source or record.last_used_from
records[index] = record
updated = True
break
if updated:
_save_records(records)
def generate_raw_token() -> str:
return API_TOKEN_PREFIX + secrets.token_urlsafe(36)
def hash_token(raw_token: str) -> str:
digest = hmac.new(API_TOKEN_SECRET.encode(), raw_token.encode(), hashlib.sha256).digest()
return base64.urlsafe_b64encode(digest).decode().rstrip("=")
def normalize_expiration(value: str | None) -> str | None:
if value is None:
return None
trimmed = str(value).strip()
if not trimmed:
return None
parsed = parse_datetime(trimmed)
if parsed is None:
raise ValueError("Invalid expiration date")
return parsed.isoformat(timespec="seconds")
def expires_at_from_days(days: int) -> str:
safe_days = max(int(days), 1)
expires_at = datetime.now(timezone.utc) + timedelta(days=safe_days)
return expires_at.isoformat(timespec="seconds")
def parse_datetime(value: str | None) -> datetime | None:
if not value:
return None
try:
parsed = datetime.fromisoformat(value)
except ValueError:
try:
parsed_date = date.fromisoformat(value)
except ValueError:
return None
parsed = datetime.combine(parsed_date, time.max)
if parsed.tzinfo is None:
return parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc)
def utc_now() -> str:
return datetime.now(timezone.utc).isoformat(timespec="seconds")
def record_to_public_dict(record: ApiTokenRecord, include_hash: bool = False) -> dict[str, Any]:
expires_at = parse_datetime(record.expires_at)
data = {
"id": record.id,
"name": record.name,
"prefix": record.prefix,
"suffix": record.suffix,
"created_at": record.created_at,
"expires_at": record.expires_at,
"enabled": record.enabled,
"last_used_at": record.last_used_at,
"last_used_from": record.last_used_from,
"expired": bool(expires_at and expires_at < datetime.now(timezone.utc)),
"token": record.token,
"token_available": bool(record.token),
}
if include_hash:
data["token_hash"] = record.token_hash
return data
def record_to_context(record: ApiTokenRecord) -> dict[str, Any]:
return {
"token_id": record.id,
"name": record.name,
"created_at": record.created_at,
"expires_at": record.expires_at,
"token_type": "api_token",
}
def _load_records() -> list[ApiTokenRecord]:
ensure_store()
try:
raw = json.loads(API_TOKENS_PATH.read_text(encoding="utf-8"))
except json.JSONDecodeError:
raw = {"tokens": []}
tokens = raw.get("tokens", []) if isinstance(raw, dict) else []
records: list[ApiTokenRecord] = []
for item in tokens:
if not isinstance(item, dict):
continue
record = _record_from_dict(item)
if record.token_hash:
records.append(record)
records.sort(key=lambda record: record.created_at, reverse=True)
return records
def _record_from_dict(item: dict[str, Any]) -> ApiTokenRecord:
raw_token = str(item.get("token") or "").strip() or None
token_hash = str(item.get("token_hash") or "").strip()
if raw_token and not token_hash:
token_hash = hash_token(raw_token)
prefix = str(item.get("prefix") or "")
suffix = str(item.get("suffix") or "")
if raw_token:
prefix = prefix or raw_token[:12]
suffix = suffix or raw_token[-12:]
return ApiTokenRecord(
id=str(item.get("id", "")) or secrets.token_hex(8),
name=str(item.get("name", "未命名 Token")),
token_hash=token_hash,
prefix=prefix,
suffix=suffix,
created_at=str(item.get("created_at", utc_now())),
expires_at=item.get("expires_at"),
enabled=bool(item.get("enabled", True)),
last_used_at=item.get("last_used_at"),
last_used_from=item.get("last_used_from"),
token=raw_token,
)
def _save_records(records: list[ApiTokenRecord]) -> None:
payload = {"tokens": [record.to_dict() for record in records]}
API_TOKENS_PATH.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")