#!/usr/bin/env python3 # -*- coding: UTF-8 -*- # ----------------------------------------------------------------------------------------------------------- # Copyright (c) 2025 Huawei Technologies Co., Ltd. # This program is free software, you can redistribute it and/or modify it under the terms and conditions of # CANN Open Software License Agreement Version 2.0 (the "License"). # Please refer to the License for details. You may not use this file except in compliance with the License. # THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, # INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. # See LICENSE in the root of the software repository for the full text of the License. # ----------------------------------------------------------------------------------------------------------- import os import xml.etree.ElementTree as ET from functools import total_ordering from pathlib import Path from typing import NamedTuple import regex as re class VersionInfoError(Exception): """版本信息异常基类。""" class VersionFormatNotMatch(VersionInfoError): """版本格式未匹配。""" class IntervalFormatNotMatch(VersionInfoError): """区间格式未匹配。""" class DuplicatedPkgConfig(VersionInfoError): """解析版本配置失败。重复的包配置。""" def __init__(self, pkg_name): super().__init__(pkg_name) self.pkg_name = pkg_name class ParseVersionFailed(VersionInfoError): """解析版本失败。""" class CollectRequiresFailed(VersionInfoError): """收集包需求失败。""" def __init__(self, pkg_name, version_str, msg): super().__init__(pkg_name, version_str, msg) self.pkg_name = pkg_name self.version_str = version_str self.msg = msg @total_ordering class Version: """版本号。""" def __init__(self, version): self.version = version @classmethod def match(cls, input_str): """输入字符串是否匹配版本号模式。""" m = re.match(r"[.a-zA-Z0-9]+$", input_str) return bool(m) @classmethod def parse(cls, input_str): """解析版本号。""" if not cls.match(input_str): raise VersionFormatNotMatch() return cls(input_str) @classmethod def try_convert_to_int_list(cls, str_list): """尝试转换为int数组。""" for idx, item in enumerate(str_list): try: int_item = int(item) str_list[idx] = int_item except ValueError: pass def to_required_list(self): """转换为版本需求字符串列表。""" return [self.version] def __eq__(self, other): """等于。""" if not isinstance(other, self.__class__): return False return self.version == other.version def __lt__(self, other): """小于。""" if not isinstance(other, self.__class__): return True self_list = self.version.split(".") other_list = other.version.split(".") self.try_convert_to_int_list(self_list) self.try_convert_to_int_list(other_list) self_tuple = tuple(self_list) other_tuple = tuple(other_list) return self_tuple < other_tuple def __str__(self): return self.version def __repr__(self): return repr(self.version) class Point(NamedTuple): """区间端点。""" type_: int # 类型,0为闭区间,1为开区间 value: Version class Interval(NamedTuple): """版本号区间。""" low: Point high: Point @classmethod def match(cls, input_str: str) -> bool: """输入字符串是否匹配区间模式。""" if not input_str.startswith("(") and not input_str.startswith("["): return False if not input_str.endswith(")") and not input_str.endswith("]"): return False input_str = input_str[1:-1] return input_str.count(",") <= 1 @classmethod def parse(cls, input_str): """解析版本号区间。""" if not cls.match(input_str): raise IntervalFormatNotMatch() if input_str[0] == "[": low_type = 0 elif input_str[0] == "(": low_type = 1 else: raise AssertionError("should not go here.") if input_str[-1] == "]": high_type = 0 elif input_str[-1] == ")": high_type = 1 else: raise AssertionError("should not go here.") input_str = input_str[1:-1] input_list = input_str.split(",") low = input_list[0].strip() if len(input_list) > 1: high = input_list[1].strip() else: high = None if low: low_version = Point(low_type, Version(low)) else: low_version = None if high: high_version = Point(high_type, Version(high)) else: high_version = None return cls(low=low_version, high=high_version) def to_required_list(self): """转换为版本需求字符串列表。""" result = [] if self.low: if self.low.type_ == 0: operator = ">=" else: operator = ">" required_str = f"{operator}{self.low.value.version}" result.append(required_str) if self.high: if self.high.type_ == 0: operator = "<=" else: operator = "<" required_str = f"{operator}{self.high.value.version}" result.append(required_str) return result class Require(NamedTuple): """包需求。""" pkg_name: str versions: list @classmethod def _sort_key(cls, item) -> tuple: """排序键。""" if isinstance(item, Interval): # 如果存在区间左值,则左值参与排序。 if item.low: return item.low.value, item.low.type_ # 否则使用区间右值,由于开区间更小,所以type_取负。 return item.high.value, -item.high.type_ return item, 0 @classmethod def _sort_versions(cls, versions: list) -> bool: """排序版本序列。""" versions.sort(key=cls._sort_key) return True @classmethod def _to_required_list(cls, versions: list) -> list[str]: """转换为版本需求字符串列表。""" result = [] for version in versions: requires = version.to_required_list() result.extend(requires) return result @classmethod def _to_required_str(cls, versions: list) -> str: """转换为版本需求字符串。""" requires = cls._to_required_list(versions) required_str = ", ".join(requires) return required_str def sort_versions(self) -> bool: """排序版本序列。""" return self._sort_versions(self.versions) def to_required_full_str(self) -> str: """转换为版本需求字符串。""" required_str = self._to_required_str(self.versions) required_full_str = f'required_package_{self.pkg_name}_version="{required_str}"' return required_full_str class ItemElement(NamedTuple): """item元素。""" name: str version: str @classmethod def parse(cls, item_ele: ET.Element, cur_ver: str): """解析item元素。""" name = item_ele.attrib["name"] version = item_ele.attrib["version"].replace("$(CUR_VER)", cur_ver) return cls(name=name, version=version) @classmethod def skip(cls, item_ele: ET.Element): """是否跳过item元素。""" version = item_ele.attrib["version"] return version.strip() == "" class CompatibleElement(NamedTuple): """compatible元素。""" items: list @classmethod def parse(cls, compatible_ele: ET.Element, cur_ver: str): """解析compatible元素。""" items = [] for item_ele in compatible_ele.findall("./item"): if ItemElement.skip(item_ele): continue item = ItemElement.parse(item_ele, cur_ver) items.append(item) return cls(items=items) def is_version_number(version: str) -> bool: """字符串是否为版本号。""" has_slash = "/" in version return not has_slash and len(version.split(".")) >= 3 class VersionXml(NamedTuple): """版本配置。""" release_version: str version_dir: str packages: dict @classmethod def match(cls, filepath: Path | str) -> bool: """文件路径是否匹配版本信息文件。""" return str(filepath).endswith(".xml") @classmethod def parse_version(cls, version_str: str): """解析版本配置。""" ret = Interval.match(version_str) if ret: result = Interval.parse(version_str) return result ret = Version.match(version_str) if ret: result = Version.parse(version_str) return result raise ParseVersionFailed() def get_release_version(self): """获取发布版本号。""" return self.release_version def get_version_dir(self): """获取多版本目录。""" return self.version_dir def collect_requires(self, package: str) -> list[Require]: """收集对应包的包需求列表。""" requires = {} if package not in self.packages: return [] compatible = self.packages[package] for item in compatible.items: pkg_name = item.name if pkg_name not in requires: requires[pkg_name] = Require(pkg_name=pkg_name, versions=[]) version_str = item.version try: version = self.parse_version(version_str) except ParseVersionFailed as ex: msg = f"parse pkg {pkg_name} version {version_str} failed" raise CollectRequiresFailed(pkg_name, version_str, msg) from ex requires[pkg_name].versions.append(version) result = [] for pkg_name in sorted(requires.keys()): requires[pkg_name].sort_versions() result.append(requires[pkg_name]) return result def get_version_dir(version_xml: VersionXml | None, disable_multi_version: bool, version_dir: str | None) -> str | None: """获取版本目录名。""" if disable_multi_version: return None if version_dir: return version_dir # 支持从version.xml中获取version_dir if version_xml and version_xml.get_version_dir(): return version_xml.get_version_dir() return None def is_multi_version(version_dir: str) -> bool: """是否多版本。""" return bool(version_dir) class VersionInfo(NamedTuple): """版本信息。""" install_version_info: bool install_version_info_attrib: dict[str, str] | None itf_versions: list[str] version: str version_xml: VersionXml | None timestamp: str | None class VersionInfoFile(NamedTuple): """生成的版本配置。""" version: str itf_version_info: str | None = None requires: list[Require] | None = None version_dir: str | None = None timestamp: str | None = None def _get_content(self) -> str: """获取版本配置内容。""" lines = [f"Version={self.version}"] if self.version_dir: lines.append(f"version_dir={self.version_dir}") if self.timestamp: lines.append(f"timestamp={self.timestamp}") if self.itf_version_info: lines.append(self.itf_version_info) if self.requires: requires_str = [require.to_required_full_str() for require in self.requires] lines.extend(requires_str) lines.append("") return "\n".join(lines) def save(self, target_path: Path | str): """保存版本配置。""" content = self._get_content() target_dir = os.path.dirname(target_path) if not os.path.exists(target_dir): os.makedirs(target_dir) with open(target_path, "w") as file: file.write(content)