Files
module1_3/final/scripts/utils.py
2026-04-23 12:48:59 +08:00

110 lines
4.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import numpy as np
from typing import Dict, List, Any
from scipy.spatial.distance import euclidean
class NetworkUtils:
"""无需 aux_data全部现场从节点属性读取"""
""" 角色匹配矩阵暂时没加 """
# ----------- 通用兜底 -----------
_DEFAULT = {
"protocol": 0.80,
"format": 0.85,
"security": 0.90,
"role": 0.70,
"area": 0.50,
"org": 0.70,
"area_1": 1,
"area_2": 0.7,
"history": 0.8
}
# ----------- 空间 -----------
@staticmethod
def calculate_distance(pos1, pos2):
return euclidean(pos1, pos2)
@staticmethod
def calculate_range_match(R_i: float, R_j: float, d_ij: float) -> float:
r_min = min(R_i, R_j)
if d_ij <= r_min:
return 1.0
return r_min / d_ij if d_ij > 0 else 0.0
@staticmethod
def calculate_time_window_overlap(node_i, node_j):
w_i = node_i["temporal"].get("time_window", 24)
w_j = node_j["temporal"].get("time_window", 24)
overlap = min(w_i, w_j)
return overlap / max(w_i, w_j, 1)
# ----------- 协议兼容性Jaccard -----------
@staticmethod
def get_protocol_compatibility(node_i, node_j, _dummy=None):
"""节点属性里放 'protocol_list'"""
p_i = set(node_i.get("protocol_list", []))
p_j = set(node_j.get("protocol_list", []))
if not p_i or not p_j:
return NetworkUtils._DEFAULT["protocol"]
intersection = len(p_i & p_j)
union = len(p_i | p_j)
return intersection / union if union else NetworkUtils._DEFAULT["protocol"]
# ----------- 安全等级1-5 映射 0-1 -----------
@staticmethod
def get_security_compatibility(node_i, node_j, _dummy=None):
lv_i = node_i.get("security_level", 3)
lv_j = node_j.get("security_level", 3)
gap = abs(lv_i - lv_j)
return max(0, 1 - gap / 5)
# ----------- 组织隶属:同单位给高分 -----------
@staticmethod
def get_organization_relation(node_i, node_j, _dummy=None):
unit_i = node_i.get("org_unit", "")
unit_j = node_j.get("org_unit", "")
return 1.0 if unit_i == unit_j and unit_i else NetworkUtils._DEFAULT["org"]
# ----------- 历史交互:直接读属性 -----------
@staticmethod
def get_interaction_history(node_i, node_j, _dummy=None):
his_i = node_i.get("history_success", NetworkUtils._DEFAULT["history"])
# print("111111111111111111111",node_i.get("history_success"))
his_j = node_j.get("history_success", NetworkUtils._DEFAULT["history"])
return (his_i + his_j) / 2
# ----------- 区域关联:同 code 高分 -----------
@staticmethod
def get_area_relation(node_i, node_j, _dummy=None):
code_i = node_i.get("nation")
code_j = node_j.get("nation")
if code_i == code_j:
return NetworkUtils._DEFAULT["area_1"]
else:
return NetworkUtils._DEFAULT["area_2"]
# ----------- 接口/角色/格式:若节点属性扩展了再读,否则给默认 -----------
@staticmethod
def get_role_match(node_i, node_j, _dummy=None):
return node_i.get("role_match", NetworkUtils._DEFAULT["role"])
@staticmethod
def get_interface_standard(node_i, node_j, _dummy=None) -> float:
"""接口标准兼容性 = Jaccard(i_list, j_list)"""
i_set = set(node_i.get("interface_list", []))
j_set = set(node_j.get("interface_list", []))
if not i_set or not j_set:
return NetworkUtils._DEFAULT["protocol"] # 默认0.8
inter = len(i_set & j_set)
union = len(i_set | j_set)
return inter / union if union else NetworkUtils._DEFAULT["protocol"]
@staticmethod
def get_format_compatibility(node_i, node_j, _dummy=None) -> float:
"""数据格式兼容性 = Jaccard(i_format, j_format)"""
i_set = set(node_i.get("format_list", []))
j_set = set(node_j.get("format_list", []))
if not i_set or not j_set:
return NetworkUtils._DEFAULT["format"] # 默认0.85
inter = len(i_set & j_set)
union = len(i_set | j_set)
return inter / union if union else NetworkUtils._DEFAULT["format"]