Team Ai
Apppublic

cjovs/codex-console

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
utils.py571 linesDownload Raw Back to core
1"""2通用工具函数3"""4 5import os6import sys7import json8import time9import random10import string11import secrets12import hashlib13import logging14import base6415import re16import uuid17from datetime import datetime, timedelta18from typing import Any, Dict, List, Optional, Union, Callable19from pathlib import Path20 21from ..config.constants import PASSWORD_CHARSET, DEFAULT_PASSWORD_LENGTH22from ..config.settings import get_settings23 24 25def setup_logging(26    log_level: str = "INFO",27    log_file: Optional[str] = None,28    log_format: str = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"29) -> logging.Logger:30    """31    配置日志系统32 33    Args:34        log_level: 日志级别 (DEBUG, INFO, WARNING, ERROR, CRITICAL)35        log_file: 日志文件路径,如果不指定则只输出到控制台36        log_format: 日志格式37 38    Returns:39        根日志记录器40    """41    # 设置日志级别42    numeric_level = getattr(logging, log_level.upper(), None)43    if not isinstance(numeric_level, int):44        numeric_level = logging.INFO45 46    # 配置根日志记录器47    root_logger = logging.getLogger()48    root_logger.setLevel(numeric_level)49 50    # 清除现有的处理器51    root_logger.handlers.clear()52 53    # 创建格式化器54    formatter = logging.Formatter(log_format)55 56    # 控制台处理器57    console_handler = logging.StreamHandler(sys.stdout)58    console_handler.setFormatter(formatter)59    console_handler.setLevel(numeric_level)60    root_logger.addHandler(console_handler)61 62    # 文件处理器(如果指定了日志文件)63    if log_file:64        # 确保日志目录存在65        log_dir = os.path.dirname(log_file)66        if log_dir:67            os.makedirs(log_dir, exist_ok=True)68 69        file_handler = logging.FileHandler(log_file, encoding="utf-8")70        file_handler.setFormatter(formatter)71        file_handler.setLevel(numeric_level)72        root_logger.addHandler(file_handler)73 74    return root_logger75 76 77def generate_password(length: int = DEFAULT_PASSWORD_LENGTH) -> str:78    """79    生成随机密码80 81    Args:82        length: 密码长度83 84    Returns:85        随机密码字符串86    """87    if length < 4:88        length = 489 90    # 确保密码包含至少一个大写字母、一个小写字母和一个数字91    password = [92        secrets.choice(string.ascii_lowercase),93        secrets.choice(string.ascii_uppercase),94        secrets.choice(string.digits),95    ]96 97    # 添加剩余字符98    password.extend(secrets.choice(PASSWORD_CHARSET) for _ in range(length - 3))99 100    # 随机打乱101    secrets.SystemRandom().shuffle(password)102 103    return ''.join(password)104 105 106def generate_random_string(length: int = 8) -> str:107    """108    生成随机字符串(仅字母)109 110    Args:111        length: 字符串长度112 113    Returns:114        随机字符串115    """116    chars = string.ascii_letters117    return ''.join(secrets.choice(chars) for _ in range(length))118 119 120def generate_uuid() -> str:121    """生成 UUID 字符串"""122    return str(uuid.uuid4())123 124 125def get_timestamp() -> int:126    """获取当前时间戳(秒)"""127    return int(time.time())128 129 130def format_datetime(dt: Optional[datetime] = None, fmt: str = "%Y-%m-%d %H:%M:%S") -> str:131    """132    格式化日期时间133 134    Args:135        dt: 日期时间对象,如果为 None 则使用当前时间136        fmt: 格式字符串137 138    Returns:139        格式化后的字符串140    """141    if dt is None:142        dt = datetime.now()143    return dt.strftime(fmt)144 145 146def parse_datetime(dt_str: str, fmt: str = "%Y-%m-%d %H:%M:%S") -> Optional[datetime]:147    """148    解析日期时间字符串149 150    Args:151        dt_str: 日期时间字符串152        fmt: 格式字符串153 154    Returns:155        日期时间对象,如果解析失败返回 None156    """157    try:158        return datetime.strptime(dt_str, fmt)159    except (ValueError, TypeError):160        return None161 162 163def human_readable_size(size_bytes: int) -> str:164    """165    将字节大小转换为人类可读的格式166 167    Args:168        size_bytes: 字节大小169 170    Returns:171        人类可读的字符串172    """173    if size_bytes < 0:174        return "0 B"175 176    units = ["B", "KB", "MB", "GB", "TB", "PB"]177    unit_index = 0178 179    while size_bytes >= 1024 and unit_index < len(units) - 1:180        size_bytes /= 1024181        unit_index += 1182 183    return f"{size_bytes:.2f} {units[unit_index]}"184 185 186def retry_with_backoff(187    func: Callable,188    max_retries: int = 3,189    base_delay: float = 1.0,190    max_delay: float = 30.0,191    backoff_factor: float = 2.0,192    exceptions: tuple = (Exception,)193) -> Any:194    """195    带有指数退避的重试装饰器/函数196 197    Args:198        func: 要重试的函数199        max_retries: 最大重试次数200        base_delay: 基础延迟(秒)201        max_delay: 最大延迟(秒)202        backoff_factor: 退避因子203        exceptions: 要捕获的异常类型204 205    Returns:206        函数的返回值207 208    Raises:209        最后一次尝试的异常210    """211    last_exception = None212 213    for attempt in range(max_retries + 1):214        try:215            return func()216        except exceptions as e:217            last_exception = e218 219            # 如果是最后一次尝试,直接抛出异常220            if attempt == max_retries:221                break222 223            # 计算延迟时间224            delay = min(base_delay * (backoff_factor ** attempt), max_delay)225 226            # 添加随机抖动227            delay *= (0.5 + random.random())228 229            # 记录日志230            logger = logging.getLogger(__name__)231            logger.warning(232                f"尝试 {func.__name__} 失败 (attempt {attempt + 1}/{max_retries + 1}): {e}. "233                f"等待 {delay:.2f} 秒后重试..."234            )235 236            time.sleep(delay)237 238    # 所有重试都失败,抛出最后一个异常239    raise last_exception240 241 242class RetryDecorator:243    """重试装饰器类"""244 245    def __init__(246        self,247        max_retries: int = 3,248        base_delay: float = 1.0,249        max_delay: float = 30.0,250        backoff_factor: float = 2.0,251        exceptions: tuple = (Exception,)252    ):253        self.max_retries = max_retries254        self.base_delay = base_delay255        self.max_delay = max_delay256        self.backoff_factor = backoff_factor257        self.exceptions = exceptions258 259    def __call__(self, func: Callable) -> Callable:260        """装饰器调用"""261        def wrapper(*args, **kwargs):262            def func_to_retry():263                return func(*args, **kwargs)264 265            return retry_with_backoff(266                func_to_retry,267                max_retries=self.max_retries,268                base_delay=self.base_delay,269                max_delay=self.max_delay,270                backoff_factor=self.backoff_factor,271                exceptions=self.exceptions272            )273 274        return wrapper275 276 277def validate_email(email: str) -> bool:278    """279    验证邮箱地址格式280 281    Args:282        email: 邮箱地址283 284    Returns:285        是否有效286    """287    pattern = r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$"288    return bool(re.match(pattern, email))289 290 291def validate_url(url: str) -> bool:292    """293    验证 URL 格式294 295    Args:296        url: URL297 298    Returns:299        是否有效300    """301    pattern = r"^https?://[^\s/$.?#].[^\s]*$"302    return bool(re.match(pattern, url))303 304 305def sanitize_filename(filename: str) -> str:306    """307    清理文件名,移除不安全的字符308 309    Args:310        filename: 原始文件名311 312    Returns:313        清理后的文件名314    """315    # 移除危险字符316    filename = re.sub(r'[<>:"/\\|?*]', '_', filename)317    # 移除控制字符318    filename = ''.join(char for char in filename if ord(char) >= 32)319    # 限制长度320    if len(filename) > 255:321        name, ext = os.path.splitext(filename)322        filename = name[:255 - len(ext)] + ext323    return filename324 325 326def read_json_file(filepath: str) -> Optional[Dict[str, Any]]:327    """328    读取 JSON 文件329 330    Args:331        filepath: 文件路径332 333    Returns:334        JSON 数据,如果读取失败返回 None335    """336    try:337        with open(filepath, 'r', encoding='utf-8') as f:338            return json.load(f)339    except (FileNotFoundError, json.JSONDecodeError, IOError) as e:340        logging.getLogger(__name__).warning(f"读取 JSON 文件失败: {filepath} - {e}")341        return None342 343 344def write_json_file(filepath: str, data: Dict[str, Any], indent: int = 2) -> bool:345    """346    写入 JSON 文件347 348    Args:349        filepath: 文件路径350        data: 要写入的数据351        indent: 缩进空格数352 353    Returns:354        是否成功355    """356    try:357        # 确保目录存在358        os.makedirs(os.path.dirname(filepath), exist_ok=True)359 360        with open(filepath, 'w', encoding='utf-8') as f:361            json.dump(data, f, ensure_ascii=False, indent=indent)362 363        return True364    except (IOError, TypeError) as e:365        logging.getLogger(__name__).error(f"写入 JSON 文件失败: {filepath} - {e}")366        return False367 368 369def get_project_root() -> Path:370    """371    获取项目根目录372 373    Returns:374        项目根目录 Path 对象375    """376    # 当前文件所在目录377    current_dir = Path(__file__).parent378 379    # 向上查找直到找到项目根目录(包含 pyproject.toml 或 setup.py)380    for parent in [current_dir] + list(current_dir.parents):381        if (parent / "pyproject.toml").exists() or (parent / "setup.py").exists():382            return parent383 384    # 如果找不到,返回当前目录的父目录385    return current_dir.parent386 387 388def get_data_dir() -> Path:389    """390    获取数据目录391 392    Returns:393        数据目录 Path 对象394    """395    settings = get_settings()396    if not settings.database_url.startswith("sqlite"):397        data_dir = Path(os.environ.get("APP_DATA_DIR", "data"))398        data_dir.mkdir(parents=True, exist_ok=True)399        return data_dir400    data_dir = Path(settings.database_url).parent401 402    # 如果 database_url 是 SQLite URL,提取路径403    if settings.database_url.startswith("sqlite:///"):404        db_path = settings.database_url[10:]  # 移除 "sqlite:///"405        data_dir = Path(db_path).parent406 407    # 确保目录存在408    data_dir.mkdir(parents=True, exist_ok=True)409 410    return data_dir411 412 413def get_logs_dir() -> Path:414    """415    获取日志目录416 417    Returns:418        日志目录 Path 对象419    """420    settings = get_settings()421    log_file = Path(settings.log_file)422    log_dir = log_file.parent423 424    # 确保目录存在425    log_dir.mkdir(parents=True, exist_ok=True)426 427    return log_dir428 429 430def format_duration(seconds: int) -> str:431    """432    格式化持续时间433 434    Args:435        seconds: 秒数436 437    Returns:438        格式化的持续时间字符串439    """440    if seconds < 60:441        return f"{seconds}秒"442 443    minutes, seconds = divmod(seconds, 60)444    if minutes < 60:445        return f"{minutes}分{seconds}秒"446 447    hours, minutes = divmod(minutes, 60)448    if hours < 24:449        return f"{hours}小时{minutes}分"450 451    days, hours = divmod(hours, 24)452    return f"{days}天{hours}小时"453 454 455def mask_sensitive_data(data: Union[str, Dict, List], mask_char: str = "*") -> Union[str, Dict, List]:456    """457    掩码敏感数据458 459    Args:460        data: 要掩码的数据461        mask_char: 掩码字符462 463    Returns:464        掩码后的数据465    """466    if isinstance(data, str):467        # 如果是邮箱,掩码中间部分468        if "@" in data:469            local, domain = data.split("@", 1)470            if len(local) > 2:471                masked_local = local[0] + mask_char * (len(local) - 2) + local[-1]472            else:473                masked_local = mask_char * len(local)474            return f"{masked_local}@{domain}"475 476        # 如果是 token 或密钥,掩码大部分内容477        if len(data) > 10:478            return data[:4] + mask_char * (len(data) - 8) + data[-4:]479        return mask_char * len(data)480 481    elif isinstance(data, dict):482        masked_dict = {}483        for key, value in data.items():484            # 敏感字段名485            sensitive_keys = ["password", "token", "secret", "key", "auth", "credential"]486            if any(sensitive in key.lower() for sensitive in sensitive_keys):487                masked_dict[key] = mask_sensitive_data(value, mask_char)488            else:489                masked_dict[key] = value490        return masked_dict491 492    elif isinstance(data, list):493        return [mask_sensitive_data(item, mask_char) for item in data]494 495    return data496 497 498def calculate_md5(data: Union[str, bytes]) -> str:499    """500    计算 MD5 哈希501 502    Args:503        data: 要哈希的数据504 505    Returns:506        MD5 哈希字符串507    """508    if isinstance(data, str):509        data = data.encode('utf-8')510 511    return hashlib.md5(data).hexdigest()512 513 514def calculate_sha256(data: Union[str, bytes]) -> str:515    """516    计算 SHA256 哈希517 518    Args:519        data: 要哈希的数据520 521    Returns:522        SHA256 哈希字符串523    """524    if isinstance(data, str):525        data = data.encode('utf-8')526 527    return hashlib.sha256(data).hexdigest()528 529 530def base64_encode(data: Union[str, bytes]) -> str:531    """Base64 编码"""532    if isinstance(data, str):533        data = data.encode('utf-8')534 535    return base64.b64encode(data).decode('utf-8')536 537 538def base64_decode(data: str) -> str:539    """Base64 解码"""540    try:541        decoded = base64.b64decode(data)542        return decoded.decode('utf-8')543    except (base64.binascii.Error, UnicodeDecodeError):544        return ""545 546 547class Timer:548    """计时器上下文管理器"""549 550    def __init__(self, name: str = "操作"):551        self.name = name552        self.start_time = None553        self.elapsed = None554 555    def __enter__(self):556        self.start_time = time.time()557        return self558 559    def __exit__(self, exc_type, exc_val, exc_tb):560        self.elapsed = time.time() - self.start_time561        logger = logging.getLogger(__name__)562        logger.debug(f"{self.name} 耗时: {self.elapsed:.2f} 秒")563 564    def get_elapsed(self) -> float:565        """获取经过的时间(秒)"""566        if self.elapsed is not None:567            return self.elapsed568        if self.start_time is not None:569            return time.time() - self.start_time570        return 0.0571