This commit is contained in:
gongwenxin
2025-05-28 15:55:46 +08:00
parent f0cc525141
commit 936714242f
313 changed files with 345685 additions and 7847 deletions
+2 -1
View File
@@ -62,7 +62,8 @@ class APICaller:
params=request_data.params,
json=json_payload,
data=request_data.data,
timeout=timeout
timeout=timeout,
verify=False
)
# 不立即引发异常,而是捕获状态码
@@ -1,531 +0,0 @@
"""规则执行引擎
该模块负责执行不同类型的规则,包括Python代码规则、断言规则等,
支持在API测试的不同生命周期阶段(请求准备、执行、响应验证等)执行规则。
"""
import logging
import importlib
import inspect
import time
import threading
from typing import Dict, List, Any, Optional, Union, Callable
from ..models.rule_models import (
AnyRule, BaseRule, RuleCategory, TargetType, RuleLifecycle, RuleScope,
PythonCodeRule, BusinessAssertionTemplate, PerformanceRule, SecurityRule,
RESTfulDesignRule, ErrorHandlingRule
)
from ..api_caller.caller import APIRequest, APIResponse
from ..rule_repository.repository import RuleRepository
class RuleExecutionError(Exception):
"""规则执行过程中发生的错误"""
pass
class RuleExecutionResult:
"""规则执行结果"""
def __init__(self,
rule: BaseRule,
is_valid: bool,
message: str = "",
details: Optional[Dict[str, Any]] = None,
error: Optional[Exception] = None):
"""
初始化规则执行结果
Args:
rule: 执行的规则
is_valid: 验证是否通过
message: 执行结果消息
details: 详细信息
error: 执行过程中发生的异常(如果有)
"""
self.rule = rule
self.rule_id = rule.id
self.rule_name = rule.name
self.rule_category = rule.category
self.is_valid = is_valid
self.message = message
self.details = details or {}
self.error = error
def __bool__(self):
"""允许直接使用结果对象作为布尔值,表示验证是否通过"""
return self.is_valid
def to_dict(self) -> Dict[str, Any]:
"""将结果转换为字典"""
return {
'rule_id': self.rule_id,
'rule_name': self.rule_name,
'rule_category': self.rule_category.value,
'is_valid': self.is_valid,
'message': self.message,
'details': self.details,
'error': str(self.error) if self.error else None
}
class RuleExecutor:
"""规则执行引擎"""
def __init__(self, rule_repository: RuleRepository):
"""
初始化规则执行引擎
Args:
rule_repository: 规则库实例
"""
self.rule_repository = rule_repository
self.logger = logging.getLogger(__name__)
def execute_rule(self, rule: BaseRule, context: Dict[str, Any]) -> RuleExecutionResult:
"""
执行单个规则
Args:
rule: 要执行的规则
context: 执行上下文,包含API请求、响应等信息
Returns:
执行结果
"""
if not rule.is_enabled:
return RuleExecutionResult(
rule=rule,
is_valid=True,
message=f"规则 {rule.id} 已禁用,跳过执行"
)
try:
# 根据规则类型选择适当的执行方法
if rule.category == RuleCategory.PYTHON_CODE or hasattr(rule, 'code') and rule.code:
return self._execute_python_code_rule(rule, context)
elif rule.category == RuleCategory.BUSINESS_LOGIC:
return self._execute_business_assertion_rule(rule, context)
elif rule.category == RuleCategory.PERFORMANCE:
return self._execute_performance_rule(rule, context)
elif rule.category == RuleCategory.SECURITY:
return self._execute_security_rule(rule, context)
elif rule.category == RuleCategory.API_DESIGN:
return self._execute_api_design_rule(rule, context)
elif rule.category == RuleCategory.ERROR_HANDLING:
return self._execute_error_handling_rule(rule, context)
else:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"不支持的规则类型: {rule.category}"
)
except Exception as e:
self.logger.error(f"执行规则 {rule.id} 失败: {e}", exc_info=True)
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"规则执行失败: {e}",
error=e
)
def _execute_python_code_rule(self, rule: BaseRule, context: Dict[str, Any]) -> RuleExecutionResult:
"""执行Python代码规则"""
# 获取规则中的Python代码
code = getattr(rule, 'code', None)
if not code:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="规则未提供Python代码"
)
# 准备执行环境
namespace = {'__builtins__': __builtins__}
namespace.update(context)
try:
# 编译并执行代码
compiled_code = compile(code, f"<rule:{rule.id}>", 'exec')
exec(compiled_code, namespace)
# 查找并执行验证函数
entry_function = 'validate'
if hasattr(rule, 'entry_function') and rule.entry_function:
entry_function = rule.entry_function
if entry_function not in namespace:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"找不到入口函数 '{entry_function}'"
)
validate_func = namespace[entry_function]
if not callable(validate_func):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"'{entry_function}' 不是可调用的函数"
)
# 调用验证函数
timeout = getattr(rule, 'timeout', 5) # 默认5秒超时
result = self._execute_with_timeout(validate_func, (context,), {}, timeout)
# 处理执行结果
if isinstance(result, dict):
return RuleExecutionResult(
rule=rule,
is_valid=bool(result.get('is_valid', False)),
message=result.get('message', ''),
details=result.get('details', {})
)
else:
return RuleExecutionResult(
rule=rule,
is_valid=bool(result),
message=str(result) if result is not True else "验证通过"
)
except Exception as e:
self.logger.error(f"执行Python代码规则 {rule.id} 失败: {e}", exc_info=True)
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"执行代码失败: {e}",
error=e
)
def _execute_with_timeout(self, func: Callable, args: tuple, kwargs: Dict[str, Any], timeout: int) -> Any:
"""使用超时执行函数"""
result = [None]
exception = [None]
def target():
try:
result[0] = func(*args, **kwargs)
except Exception as e:
exception[0] = e
thread = threading.Thread(target=target)
thread.daemon = True
thread.start()
thread.join(timeout)
if thread.is_alive():
raise RuleExecutionError(f"规则执行超时(超过{timeout}秒)")
if exception[0]:
raise exception[0]
return result[0]
def _execute_business_assertion_rule(self, rule: BusinessAssertionTemplate, context: Dict[str, Any]) -> RuleExecutionResult:
"""执行业务断言规则"""
# 验证是否提供了所有必需的参数
if rule.expected_parameters:
missing_params = [p for p in rule.expected_parameters if p not in context]
if missing_params:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"缺少必需的参数: {', '.join(missing_params)}"
)
if rule.template_language == "python_expression":
try:
# 使用eval执行Python表达式
result = eval(rule.template_expression, {'__builtins__': __builtins__}, context)
return RuleExecutionResult(
rule=rule,
is_valid=bool(result),
message="断言通过" if result else "断言失败"
)
except Exception as e:
self.logger.error(f"执行Python表达式断言失败: {e}", exc_info=True)
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"表达式执行失败: {e}",
error=e
)
else:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"不支持的模板语言: {rule.template_language}"
)
def _execute_performance_rule(self, rule: PerformanceRule, context: Dict[str, Any]) -> RuleExecutionResult:
"""执行性能规则"""
response = context.get('api_response')
if not response or not isinstance(response, APIResponse):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="缺少有效的API响应对象"
)
# 获取响应时间(毫秒)
elapsed_time = response.elapsed_time * 1000 # 转换为毫秒
# 检查是否超过阈值
if elapsed_time > rule.threshold:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"响应时间({elapsed_time:.2f}ms)超过阈值({rule.threshold}{rule.unit}",
details={
'actual_time': elapsed_time,
'threshold': rule.threshold,
'unit': rule.unit
}
)
return RuleExecutionResult(
rule=rule,
is_valid=True,
message=f"响应时间({elapsed_time:.2f}ms)在阈值范围内",
details={
'actual_time': elapsed_time,
'threshold': rule.threshold,
'unit': rule.unit
}
)
def _execute_security_rule(self, rule: SecurityRule, context: Dict[str, Any]) -> RuleExecutionResult:
"""执行安全规则"""
if rule.check_type == "transport_security":
request = context.get('api_request')
if not request or not isinstance(request, APIRequest):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="缺少有效的API请求对象"
)
url = str(request.url)
# 检查URL是否使用HTTPS
if not url.startswith('https://'):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="API请求必须使用HTTPS协议",
details={
'current_url': url,
'expected_protocol': 'https'
}
)
return RuleExecutionResult(
rule=rule,
is_valid=True,
message="API请求使用了HTTPS协议",
details={
'url': url
}
)
else:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"不支持的安全检查类型: {rule.check_type}"
)
def _execute_api_design_rule(self, rule: RESTfulDesignRule, context: Dict[str, Any]) -> RuleExecutionResult:
"""执行API设计规则"""
import re
request = context.get('api_request')
if not request or not isinstance(request, APIRequest):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="缺少有效的API请求对象"
)
url = str(request.url)
# 解析URL,获取路径部分
from urllib.parse import urlparse
parsed_url = urlparse(url)
path = parsed_url.path
# 使用正则表达式验证路径
if rule.pattern and not re.match(rule.pattern, path):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"API路径不符合{rule.design_aspect}规范",
details={
'current_path': path,
'expected_pattern': rule.pattern
}
)
return RuleExecutionResult(
rule=rule,
is_valid=True,
message=f"API路径符合{rule.design_aspect}规范",
details={
'path': path
}
)
def _execute_error_handling_rule(self, rule: ErrorHandlingRule, context: Dict[str, Any]) -> RuleExecutionResult:
"""执行错误处理规则"""
response = context.get('api_response')
if not response or not isinstance(response, APIResponse):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="缺少有效的API响应对象"
)
# 只验证4xx和5xx状态码
if response.status_code < 400:
return RuleExecutionResult(
rule=rule,
is_valid=True,
message="非错误响应,跳过验证",
details={
'status_code': response.status_code
}
)
# 验证状态码是否匹配
if rule.expected_status != -1 and response.status_code != rule.expected_status:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"响应状态码({response.status_code})与预期({rule.expected_status})不符",
details={
'actual_status': response.status_code,
'expected_status': rule.expected_status
}
)
# 验证JSON响应
if not response.json_content:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="错误响应不是有效的JSON格式",
details={
'status_code': response.status_code,
'content_type': response.headers.get('Content-Type', '未知')
}
)
# 检查错误码
if rule.error_code != "*" and str(response.json_content.get('code', '')) != rule.error_code:
return RuleExecutionResult(
rule=rule,
is_valid=False,
message=f"错误码({response.json_content.get('code')})与预期({rule.error_code})不符",
details={
'actual_code': response.json_content.get('code'),
'expected_code': rule.error_code
}
)
# 验证错误消息
if rule.expected_message and rule.expected_message not in str(response.json_content.get('message', '')):
return RuleExecutionResult(
rule=rule,
is_valid=False,
message="错误消息与预期不符",
details={
'actual_message': response.json_content.get('message'),
'expected_message': rule.expected_message
}
)
return RuleExecutionResult(
rule=rule,
is_valid=True,
message="错误响应符合预期",
details={
'status_code': response.status_code,
'error_code': response.json_content.get('code'),
'error_message': response.json_content.get('message')
}
)
def execute_rules_for_lifecycle(self, lifecycle: RuleLifecycle, context: Dict[str, Any]) -> List[RuleExecutionResult]:
"""
执行特定生命周期阶段的所有规则
Args:
lifecycle: 生命周期阶段
context: 执行上下文
Returns:
执行结果列表
"""
# 获取适用于该生命周期阶段的所有规则
rules = self.rule_repository.get_rules_by_lifecycle(lifecycle)
# 执行规则
results = []
for rule in rules:
result = self.execute_rule(rule, context)
results.append(result)
return results
def execute_rules_for_target(self, target_type: TargetType, target_id: str, context: Dict[str, Any]) -> List[RuleExecutionResult]:
"""
执行特定目标的所有规则
Args:
target_type: 目标类型
target_id: 目标ID
context: 执行上下文
Returns:
执行结果列表
"""
# 获取适用于该目标的所有规则
rules = self.rule_repository.get_rules_for_target(target_type, target_id)
# 执行规则
results = []
for rule in rules:
result = self.execute_rule(rule, context)
results.append(result)
return results
def execute_specific_rules(self, rules: List[AnyRule], context: Dict[str, Any], lifecycle_phase: Optional[RuleLifecycle] = None) -> List[RuleExecutionResult]:
"""
执行一个明确指定的规则列表。
Args:
rules: 要执行的规则对象的列表。
context: 执行上下文,包含API请求、响应等信息。
lifecycle_phase: 可选的,名义上的生命周期阶段,可能用于上下文或某些规则的内部逻辑。
注意:此参数目前主要用于信息传递,核心执行逻辑在 execute_rule 中
并不直接依赖它来选择执行路径。
Returns:
一个包含每个规则执行结果的列表。
"""
results = []
if not rules:
self.logger.info("execute_specific_rules_called_with_no_rules")
return results
# 如果需要,可以将 lifecycle_phase 添加到 context 中传递给每个规则
# updated_context = context.copy()
# if lifecycle_phase:
# updated_context['current_lifecycle_phase'] = lifecycle_phase
for rule in rules:
# 使用现有的 execute_rule 方法执行单个规则
# result = self.execute_rule(rule, updated_context if lifecycle_phase else context)
result = self.execute_rule(rule, context) # 简化:暂时不修改context传递
results.append(result)
return results
@@ -1 +0,0 @@
@@ -1 +0,0 @@
@@ -1,91 +0,0 @@
"""Base adapter interface for rule storage."""
from abc import ABC, abstractmethod
from typing import List, Optional, Dict, Any, Union
from ...models.rule_models import AnyRule, RuleQuery, BaseRule
class BaseRuleStorageAdapter(ABC):
"""Base class for rule storage adapters."""
@abstractmethod
def initialize(self) -> None:
"""初始化适配器(例如,连接到数据库,验证文件路径等)。"""
pass
@abstractmethod
def load_rule_by_id(self, rule_id: str, version: Optional[str] = None) -> Optional[AnyRule]:
"""
按ID加载单个规则。
Args:
rule_id: 规则的唯一标识符。
version: 可选的版本标识符。如果未提供,则使用配置的默认版本策略。
Returns:
找到的规则对象,或者如果未找到则返回None。
"""
pass
@abstractmethod
def query_rules(self, query: RuleQuery) -> List[AnyRule]:
"""
根据查询条件返回匹配的规则列表。
Args:
query: 规则查询条件。
Returns:
匹配规则的列表。
"""
pass
@abstractmethod
def save_rule(self, rule: BaseRule) -> bool:
"""
保存单个规则。
Args:
rule: 要保存的规则对象。
Returns:
操作是否成功。
"""
pass
@abstractmethod
def delete_rule(self, rule_id: str, version: Optional[str] = None) -> bool:
"""
删除单个规则。
Args:
rule_id: 规则的唯一标识符。
version: 可选的版本标识符。如果未提供,通常会删除所有版本。
Returns:
操作是否成功。
"""
pass
@abstractmethod
def list_all_rule_ids(self) -> List[str]:
"""
列出存储中的所有规则ID。
Returns:
规则ID的列表。
"""
pass
@abstractmethod
def get_rule_versions(self, rule_id: str) -> List[str]:
"""
获取特定规则ID的所有可用版本。
Args:
rule_id: 规则的唯一标识符。
Returns:
版本标识符的列表。
"""
pass
@@ -1,324 +0,0 @@
"""File system adapter for rule storage using JSON files."""
import os
import json
import glob
from typing import List, Dict, Optional, Any, Tuple, Type, Union, cast
import logging
from datetime import datetime
from .base_adapter import BaseRuleStorageAdapter
from ...models.rule_models import (
AnyRule, RuleQuery, BaseRule, RuleCategory, TargetType
)
from .rule_adapter_utils import parse_rule_data
class FilesystemAdapter(BaseRuleStorageAdapter):
"""
基于文件系统的规则存储适配器,使用JSON文件保存规则。
文件结构约定:
规则文件将按以下方式组织:
rules/
json_schemas/
rule_id1/
1.0.0.json
1.1.0.json
rule_id2/
1.0.0.json
business_logic/
rule_id3/
1.0.0.json
...
"""
def __init__(self, base_path: str = "./rules", file_pattern: str = "*.json"):
"""
初始化适配器。
Args:
base_path: 存储规则的基本目录路径。
file_pattern: 匹配规则文件的glob模式。
"""
self.base_path = os.path.abspath(base_path)
self.file_pattern = file_pattern
self.logger = logging.getLogger(__name__)
def initialize(self) -> None:
"""确保基本目录存在,并验证其可访问性。"""
if not os.path.exists(self.base_path):
try:
os.makedirs(self.base_path)
self.logger.info(f"Created rules directory at {self.base_path}")
except Exception as e:
self.logger.error(f"Failed to create rules directory at {self.base_path}: {e}")
raise ValueError(f"Rules directory {self.base_path} does not exist and could not be created")
if not os.access(self.base_path, os.R_OK | os.W_OK):
self.logger.error(f"Rules directory {self.base_path} is not readable and writable")
raise ValueError(f"Rules directory {self.base_path} is not readable and writable")
# 确保每个规则类别的子目录存在
for category in RuleCategory:
category_dir = self._get_category_dir(category)
if not os.path.exists(category_dir):
try:
os.makedirs(category_dir)
self.logger.debug(f"Created category directory at {category_dir}")
except Exception as e:
self.logger.warning(f"Failed to create category directory at {category_dir}: {e}")
def _get_category_dir(self, category: RuleCategory) -> str:
"""获取给定规则类别的目录路径。"""
return os.path.join(self.base_path, category.value.lower())
def _get_rule_dir(self, rule_id: str, category: Optional[RuleCategory] = None) -> str:
"""
获取给定规则ID的目录路径。
如果指定了类别,则直接使用该类别的目录;否则尝试在所有类别中查找。
"""
if category:
return os.path.join(self._get_category_dir(category), rule_id)
# 如果未指定类别,检查所有类别目录
for cat in RuleCategory:
rule_dir = os.path.join(self._get_category_dir(cat), rule_id)
if os.path.exists(rule_dir):
return rule_dir
# 如果未找到任何匹配项,默认返回通用类别目录
return os.path.join(self._get_category_dir(RuleCategory.GENERIC), rule_id)
def _get_rule_file_path(self, rule_id: str, version: str, category: Optional[RuleCategory] = None) -> str:
"""获取给定规则ID和版本的文件路径。"""
rule_dir = self._get_rule_dir(rule_id, category)
return os.path.join(rule_dir, f"{version}.json")
def _get_rule_from_file(self, file_path: str) -> Optional[AnyRule]:
"""从文件加载规则。"""
if not os.path.exists(file_path):
self.logger.debug(f"Rule file {file_path} does not exist")
return None
try:
with open(file_path, 'r', encoding='utf-8') as f:
raw_data = json.load(f)
# 使用工具函数解析规则数据
return parse_rule_data(raw_data)
except json.JSONDecodeError as e:
self.logger.error(f"Failed to parse JSON from {file_path}: {e}")
return None
except Exception as e:
self.logger.error(f"Error loading rule from {file_path}: {e}")
return None
def _save_rule_to_file(self, rule: BaseRule, file_path: str) -> bool:
"""将规则保存到文件。"""
try:
# 确保目录存在
dir_path = os.path.dirname(file_path)
if not os.path.exists(dir_path):
os.makedirs(dir_path)
# 序列化规则为JSON并写入文件
with open(file_path, 'w', encoding='utf-8') as f:
json.dump(rule.model_dump(), f, indent=2, ensure_ascii=False)
return True
except Exception as e:
self.logger.error(f"Failed to save rule to {file_path}: {e}")
return False
def _get_latest_version(self, rule_id: str, category: Optional[RuleCategory] = None) -> Optional[str]:
"""获取给定规则ID的最新版本。"""
versions = self.get_rule_versions(rule_id, category)
if not versions:
return None
# 简单地按字符串排序,假设版本格式是类似于"1.0.0"的语义化版本
# 如果需要更复杂的版本比较,可以使用packaging.version
versions.sort()
return versions[-1]
def get_rule_versions(self, rule_id: str, category: Optional[RuleCategory] = None) -> List[str]:
"""获取给定规则ID的所有版本。"""
rule_dir = self._get_rule_dir(rule_id, category)
if not os.path.exists(rule_dir):
return []
# 查找目录中所有匹配的JSON文件,并提取版本号(文件名)
pattern = os.path.join(rule_dir, self.file_pattern)
version_files = glob.glob(pattern)
versions = []
for vf in version_files:
version = os.path.splitext(os.path.basename(vf))[0] # 移除.json扩展名
versions.append(version)
return versions
def load_rule_by_id(self, rule_id: str, version: Optional[str] = None) -> Optional[AnyRule]:
"""
按ID加载单个规则。
Args:
rule_id: 规则的唯一标识符。
version: 可选的版本标识符。如果未提供,则加载最新版本。
Returns:
找到的规则对象,或者如果未找到则返回None。
"""
if not version:
version = self._get_latest_version(rule_id)
if not version:
self.logger.debug(f"No versions found for rule ID {rule_id}")
return None
file_path = self._get_rule_file_path(rule_id, version)
return self._get_rule_from_file(file_path)
def query_rules(self, query: RuleQuery) -> List[AnyRule]:
"""
根据查询条件查询规则。
Args:
query: 包含筛选条件的查询对象。
Returns:
匹配查询条件的规则列表。
"""
results = []
# 如果指定了规则ID,直接加载该规则
if query.rule_id:
rule = self.load_rule_by_id(query.rule_id, query.version if query.version != "latest" else None)
if rule and self._rule_matches_query(rule, query):
results.append(rule)
return results
# 否则,根据查询条件扫描规则文件
categories_to_search = [query.category] if query.category else list(RuleCategory)
for category in categories_to_search:
category_dir = self._get_category_dir(category)
if not os.path.exists(category_dir):
continue
# 获取该类别下的所有规则ID(子目录)
rule_dirs = [d for d in os.listdir(category_dir)
if os.path.isdir(os.path.join(category_dir, d))]
for rule_id in rule_dirs:
# 对于每个规则ID,加载指定版本或最新版本
if query.version and query.version != "latest":
file_path = self._get_rule_file_path(rule_id, query.version, category)
rule = self._get_rule_from_file(file_path)
if rule and self._rule_matches_query(rule, query):
results.append(rule)
else:
latest_version = self._get_latest_version(rule_id, category)
if latest_version:
file_path = self._get_rule_file_path(rule_id, latest_version, category)
rule = self._get_rule_from_file(file_path)
if rule and self._rule_matches_query(rule, query):
results.append(rule)
return results
def _rule_matches_query(self, rule: BaseRule, query: RuleQuery) -> bool:
"""检查规则是否匹配查询条件。"""
# 检查是否启用
if query.is_enabled is not None and rule.is_enabled != query.is_enabled:
return False
# 检查目标类型
if query.target_type and rule.target_type != query.target_type:
return False
# 检查目标标识符
if query.target_identifier and rule.target_identifier != query.target_identifier:
return False
# 检查标签
if query.tags:
if not rule.tags:
return False
# 检查是否所有查询标签都在规则标签中
if not all(tag in rule.tags for tag in query.tags):
return False
return True
def save_rule(self, rule: BaseRule) -> bool:
"""
保存规则到文件系统。
Args:
rule: 要保存的规则对象。
Returns:
操作是否成功。
"""
file_path = self._get_rule_file_path(rule.id, rule.version, rule.category)
return self._save_rule_to_file(rule, file_path)
def delete_rule(self, rule_id: str, version: Optional[str] = None) -> bool:
"""
删除规则。
Args:
rule_id: 规则的唯一标识符。
version: 如果提供,仅删除该版本;否则删除所有版本。
Returns:
操作是否成功。
"""
try:
if version:
# 删除特定版本
file_path = self._get_rule_file_path(rule_id, version)
if os.path.exists(file_path):
os.remove(file_path)
self.logger.info(f"Deleted rule file: {file_path}")
else:
self.logger.warning(f"Rule file not found for deletion: {file_path}")
return False
else:
# 删除所有版本(整个规则目录)
rule_dir = self._get_rule_dir(rule_id)
if os.path.exists(rule_dir):
# 递归删除目录及其内容
for root, dirs, files in os.walk(rule_dir, topdown=False):
for file in files:
os.remove(os.path.join(root, file))
for dir in dirs:
os.rmdir(os.path.join(root, dir))
os.rmdir(rule_dir)
self.logger.info(f"Deleted rule directory: {rule_dir}")
else:
self.logger.warning(f"Rule directory not found for deletion: {rule_dir}")
return False
return True
except Exception as e:
self.logger.error(f"Error deleting rule {rule_id} (version={version}): {e}")
return False
def list_all_rule_ids(self) -> List[str]:
"""列出所有规则ID。"""
all_rule_ids = set()
# 扫描所有类别目录
for category in RuleCategory:
category_dir = self._get_category_dir(category)
if not os.path.exists(category_dir):
continue
# 获取该类别下的所有规则ID(子目录)
rule_dirs = [d for d in os.listdir(category_dir)
if os.path.isdir(os.path.join(category_dir, d))]
all_rule_ids.update(rule_dirs)
return list(all_rule_ids)
@@ -1,60 +0,0 @@
"""规则适配器工具函数"""
import logging
from typing import Dict, Any, Optional, Type
from ...models.rule_models import (
AnyRule, BaseRule, RuleCategory,
JSONSchemaDefinition, APILintingRuleset, BusinessAssertionTemplate, DataQualityRule,
PythonCodeRule, PerformanceRule, SecurityRule, RESTfulDesignRule, ErrorHandlingRule
)
# 规则类别到Pydantic模型类的映射
RULE_CLASSES = {
RuleCategory.JSON_SCHEMA: JSONSchemaDefinition,
RuleCategory.API_LINTING: APILintingRuleset,
RuleCategory.BUSINESS_LOGIC: BusinessAssertionTemplate,
RuleCategory.DATA_QUALITY: DataQualityRule,
RuleCategory.PYTHON_CODE: PythonCodeRule,
RuleCategory.PERFORMANCE: PerformanceRule,
RuleCategory.SECURITY: SecurityRule,
RuleCategory.API_DESIGN: RESTfulDesignRule,
RuleCategory.ERROR_HANDLING: ErrorHandlingRule,
}
logger = logging.getLogger(__name__)
def get_rule_class_by_category(category: RuleCategory) -> Type[BaseRule]:
"""根据规则类别获取对应的规则类"""
return RULE_CLASSES.get(category, BaseRule)
def parse_rule_data(raw_data: Dict[str, Any]) -> Optional[AnyRule]:
"""
解析规则数据,返回对应类型的规则对象
Args:
raw_data: 从文件加载的原始规则数据
Returns:
规则对象,如果无法解析则返回None
"""
try:
# 确保数据包含类别信息
if 'category' not in raw_data:
logger.warning("Rule data missing 'category' field")
return None
# 尝试将字符串类别转换为枚举值
try:
category = RuleCategory(raw_data['category'])
except ValueError:
logger.warning(f"Rule data has invalid category: {raw_data['category']}")
return None
# 获取对应的Pydantic模型类
rule_class = get_rule_class_by_category(category)
# 使用Pydantic模型解析数据
return rule_class(**raw_data)
except Exception as e:
logger.error(f"Error parsing rule data: {e}")
return None
@@ -1,314 +0,0 @@
"""
Python代码规则执行器
负责安全地执行规则中包含的Python代码,确保代码在隔离的环境中运行,
并对可执行的操作进行限制,以防恶意代码。
"""
import ast
import logging
import importlib
import inspect
import time
import threading
import os
from typing import Dict, Any, Optional, List, Callable, Tuple
from concurrent.futures import ThreadPoolExecutor
from ..models.rule_models import PythonCodeRule
from ..models.rule_models import RuleCategory
logger = logging.getLogger(__name__)
class CodeExecutionError(Exception):
"""执行Python代码时发生的错误"""
pass
class TimeoutError(CodeExecutionError):
"""代码执行超时错误"""
pass
class ImportError(CodeExecutionError):
"""非法导入模块错误"""
pass
class ValidationResult:
"""Python代码验证的结果"""
def __init__(self,
is_valid: bool,
message: Optional[str] = None,
details: Optional[Dict[str, Any]] = None,
exception: Optional[Exception] = None):
self.is_valid = is_valid
self.message = message
self.details = details or {}
self.exception = exception
def __bool__(self):
return self.is_valid
class PythonRuleExecutor:
"""
Python代码规则执行器
负责安全地执行规则中的Python代码,并返回验证结果。
"""
def __init__(self, rules_base_path: str = "./rules"):
self.logger = logging.getLogger(__name__)
self.rules_base_path = os.path.abspath(rules_base_path)
def _load_code_from_file(self, rule: PythonCodeRule) -> str:
"""
从文件加载代码
Args:
rule: Python代码规则对象,包含code_file属性
Returns:
加载的代码内容
Raises:
CodeExecutionError: 如果加载代码失败
"""
if not rule.code_file:
raise CodeExecutionError("规则未指定code_file属性")
# 构建代码文件的绝对路径
# 如果规则ID和版本都存在,则在python_code目录下查找
if rule.id and rule.version:
# 代码文件路径格式: rules/python_code/{rule_id}/{version}.py
file_path = os.path.join(self.rules_base_path, "python_code", rule.id, f"{rule.version}.py")
else:
# 否则使用传入的相对路径
file_path = os.path.join(self.rules_base_path, rule.code_file)
# 检查文件是否存在
if not os.path.isfile(file_path):
raise CodeExecutionError(f"代码文件不存在: {file_path}")
try:
with open(file_path, 'r', encoding='utf-8') as f:
return f.read()
except Exception as e:
raise CodeExecutionError(f"读取代码文件失败: {e}")
def analyze_code(self, code: str) -> List[str]:
"""
分析代码,检查其中使用的导入模块
Args:
code: 要分析的Python代码
Returns:
代码中导入的模块列表
Raises:
SyntaxError: 如果代码存在语法错误
"""
try:
tree = ast.parse(code)
except SyntaxError as e:
raise CodeExecutionError(f"代码语法错误: {e}")
imports = []
for node in ast.walk(tree):
if isinstance(node, ast.Import):
for name in node.names:
imports.append(name.name)
elif isinstance(node, ast.ImportFrom):
imports.append(node.module)
return imports
def _execute_with_timeout(self,
func: Callable,
args: Tuple,
kwargs: Dict[str, Any],
timeout: int) -> Any:
"""
使用超时执行函数
Args:
func: 要执行的函数
args: 函数的位置参数
kwargs: 函数的关键字参数
timeout: 超时时间(秒)
Returns:
函数的返回值
Raises:
TimeoutError: 如果函数执行超过指定的超时时间
"""
result = [None]
exception = [None]
def target():
try:
result[0] = func(*args, **kwargs)
except Exception as e:
exception[0] = e
thread = threading.Thread(target=target)
thread.daemon = True
thread.start()
thread.join(timeout)
if thread.is_alive():
# 超时
raise TimeoutError(f"代码执行超时(超过{timeout}秒)")
if exception[0]:
raise exception[0]
return result[0]
def execute_rule(self, rule: PythonCodeRule, context: Dict[str, Any]) -> ValidationResult:
"""
执行Python代码规则
Args:
rule: Python代码规则定义
context: 验证上下文(包含规则执行所需的参数和数据)
Returns:
ValidationResult: 验证结果
"""
# 检查是否提供了所有必需的参数
if rule.expected_parameters:
missing_params = [p for p in rule.expected_parameters if p not in context]
if missing_params:
return ValidationResult(
is_valid=False,
message=f"缺少必需的参数: {', '.join(missing_params)}"
)
# 获取代码内容
try:
if rule.code:
code = rule.code
elif rule.code_file:
code = self._load_code_from_file(rule)
else:
return ValidationResult(
is_valid=False,
message="规则既没有code也没有code_file属性"
)
except CodeExecutionError as e:
return ValidationResult(
is_valid=False,
message=str(e),
exception=e
)
# 分析代码中的导入
try:
imports = self.analyze_code(code)
except CodeExecutionError as e:
return ValidationResult(
is_valid=False,
message=str(e),
exception=e
)
# 检查导入是否被允许
if imports and not rule.allow_imports:
return ValidationResult(
is_valid=False,
message=f"规则不允许导入模块,但代码尝试导入: {', '.join(imports)}"
)
# 如果允许导入,检查是否所有导入都在允许列表中
if imports and rule.allow_imports and rule.allowed_modules:
unauthorized_imports = [imp for imp in imports if imp not in rule.allowed_modules]
if unauthorized_imports:
return ValidationResult(
is_valid=False,
message=f"代码尝试导入未授权的模块: {', '.join(unauthorized_imports)}"
)
# 准备执行环境
# 创建一个隔离的命名空间
namespace = {'__builtins__': __builtins__}
# 添加上下文变量
namespace.update(context)
# 如果允许导入,预先导入允许的模块
if rule.allow_imports and imports:
for module_name in imports:
if not rule.allowed_modules or module_name in rule.allowed_modules:
try:
module = importlib.import_module(module_name)
namespace[module_name] = module
except Exception as e:
return ValidationResult(
is_valid=False,
message=f"导入模块 '{module_name}' 失败: {e}",
exception=e
)
# 执行代码
try:
# 编译代码
compiled_code = compile(code, f"<rule:{rule.id}>", 'exec')
# 执行代码
self._execute_with_timeout(
exec,
(compiled_code, namespace),
{},
rule.timeout
)
# 获取入口函数
if rule.entry_function not in namespace:
return ValidationResult(
is_valid=False,
message=f"代码未定义指定的入口函数 '{rule.entry_function}'"
)
entry_func = namespace[rule.entry_function]
if not callable(entry_func):
return ValidationResult(
is_valid=False,
message=f"'{rule.entry_function}' 不是一个可调用的函数"
)
# 调用入口函数
result = self._execute_with_timeout(
entry_func,
tuple(),
{},
rule.timeout
)
# 处理结果
if isinstance(result, dict):
# 如果函数返回一个字典,将它转换为ValidationResult
return ValidationResult(
is_valid=bool(result.get('is_valid', False)),
message=result.get('message'),
details=result.get('details', {})
)
else:
# 如果返回其他类型,将其解释为布尔值
return ValidationResult(
is_valid=bool(result),
message=str(result) if result is not True else "验证通过"
)
except TimeoutError as e:
return ValidationResult(
is_valid=False,
message=str(e),
exception=e
)
except Exception as e:
return ValidationResult(
is_valid=False,
message=f"执行代码时发生错误: {e}",
exception=e
)
@@ -1,366 +0,0 @@
"""规则库核心模块"""
import logging
from typing import Dict, List, Optional, Type, Union, Any
from ..models.rule_models import AnyRule, BaseRule, RuleQuery, RuleCategory, TargetType, RuleLifecycle, RuleScope
from ..models.config_models import RuleRepositoryConfig
from .adapters.base_adapter import BaseRuleStorageAdapter
from .adapters.filesystem_adapter import FilesystemAdapter
from .yaml_adapter import YAMLAdapter
# 未来可能添加的其他适配器
# from .adapters.db_adapter import DatabaseAdapter
# from .adapters.in_memory_adapter import InMemoryAdapter
class RuleRepository:
"""
规则库模块的核心类。
负责通过合适的存储适配器管理和提供规则。
"""
def __init__(self, config: RuleRepositoryConfig):
"""
初始化规则库。
Args:
config: 规则库配置
"""
self.config = config
self.logger = logging.getLogger(__name__)
# 创建适当的存储适配器
self.adapters = self._create_adapters()
# 用于在内存中缓存规则 (如果启用了preload_rules)
self.rule_cache: Dict[str, Dict[str, AnyRule]] = {} # {rule_id: {version: rule}}
# 初始化适配器
for adapter in self.adapters:
adapter.initialize()
# 如果配置了预加载规则,则加载所有规则到内存
if self.config.preload_rules:
self._preload_rules()
def _create_adapters(self) -> List[BaseRuleStorageAdapter]:
"""根据配置创建适当的存储适配器。"""
adapters = []
storage_type = self.config.storage.type.lower()
if storage_type == "filesystem":
# 添加JSON规则适配器
adapters.append(FilesystemAdapter(
base_path=self.config.storage.path or "./rules"
))
# 添加YAML规则适配器
adapters.append(YAMLAdapter(
base_path=self.config.storage.path or "./rules"
))
# 未来可能添加的其他适配器类型
# elif storage_type == "database":
# adapters.append(DatabaseAdapter(
# connection_string=self.config.storage.connection_string
# ))
# elif storage_type == "in_memory":
# adapters.append(InMemoryAdapter())
else:
raise ValueError(f"Unsupported rule storage type: {storage_type}")
return adapters
def _preload_rules(self) -> None:
"""预加载所有规则到内存缓存。"""
self.logger.info("Preloading rules from storage...")
all_rule_ids = set()
# 从所有适配器收集规则ID
for adapter in self.adapters:
rule_ids = adapter.list_all_rule_ids()
all_rule_ids.update(rule_ids)
loaded_count = 0
for rule_id in all_rule_ids:
versions_by_adapter = []
# 从所有适配器收集规则版本
for adapter in self.adapters:
versions = adapter.get_rule_versions(rule_id)
if versions:
versions_by_adapter.append((adapter, versions))
if not versions_by_adapter:
continue
if rule_id not in self.rule_cache:
self.rule_cache[rule_id] = {}
# 对于每个适配器的每个版本,尝试加载规则
for adapter, versions in versions_by_adapter:
for version in versions:
rule = adapter.load_rule_by_id(rule_id, version)
if rule:
self.rule_cache[rule_id][version] = rule
loaded_count += 1
self.logger.info(f"Preloaded {loaded_count} rules from {len(all_rule_ids)} rule IDs")
def get_rule(self, rule_id: str, version: Optional[str] = None) -> Optional[AnyRule]:
"""
获取指定ID和版本的规则。
Args:
rule_id: 规则ID
version: 规则版本(如果未指定,则使用配置的默认版本策略)
Returns:
规则对象,如果未找到则返回None
"""
# 优先从缓存中获取,如果启用了预加载
if self.config.preload_rules and rule_id in self.rule_cache:
if version and version in self.rule_cache[rule_id]:
return self.rule_cache[rule_id][version]
elif not version and self.rule_cache[rule_id]:
# 获取最新版本
latest_version = self._get_latest_version(list(self.rule_cache[rule_id].keys()))
return self.rule_cache[rule_id].get(latest_version)
# 从适配器加载
for adapter in self.adapters:
rule = adapter.load_rule_by_id(rule_id, version)
if rule:
return rule
return None
def _get_latest_version(self, versions: List[str]) -> str:
"""简单地按字符串排序获取最新版本。"""
if not versions:
return ""
versions.sort()
return versions[-1]
def query_rules(self, query: Optional[RuleQuery] = None) -> List[AnyRule]:
"""
根据查询条件查询规则。
Args:
query: 规则查询条件,如果为None则使用默认查询
Returns:
匹配规则的列表
"""
query = query or RuleQuery()
# 从所有适配器查询规则
results = []
for adapter in self.adapters:
adapter_results = adapter.query_rules(query)
if adapter_results:
results.extend(adapter_results)
# 去重(可能不同适配器返回相同ID和版本的规则)
deduplicated = {}
for rule in results:
key = f"{rule.id}:{rule.version}"
if key not in deduplicated:
deduplicated[key] = rule
return list(deduplicated.values())
def get_rules_by_tags(self, tags: List[str], match_all: bool = False) -> List[AnyRule]:
"""
根据标签查询规则。
Args:
tags: 要匹配的标签列表。
match_all: 如果为True,则规则必须包含所有指定的标签;
如果为False(默认),则规则包含任何一个指定标签即可匹配。
Returns:
匹配标签的规则列表。
"""
if not tags:
return [] # 如果没有提供标签,返回空列表
# 获取所有规则进行过滤。可以考虑优化,如果规则量非常大,
# 且适配器支持基于标签的查询,则直接调用适配器。
# 目前,我们先在查询所有规则后进行内存过滤。
all_rules = self.query_rules(RuleQuery(is_enabled=True)) # 通常只查询启用的规则
matched_rules = []
tag_set_query = set(tag.lower() for tag in tags) # 查询标签转换为小写集合以进行不区分大小写的比较
for rule in all_rules:
if not rule.tags: # 如果规则没有标签,则跳过
continue
rule_tags_set = set(t.lower() for t in rule.tags) # 规则的标签也转换为小写集合
if match_all:
# 需要匹配所有查询标签
if tag_set_query.issubset(rule_tags_set):
matched_rules.append(rule)
else:
# 只需要匹配任何一个查询标签
if not tag_set_query.isdisjoint(rule_tags_set): # 如果交集不为空
matched_rules.append(rule)
return matched_rules
def save_rule(self, rule: BaseRule) -> bool:
"""
保存规则到存储。
Args:
rule: 要保存的规则
Returns:
操作是否成功
"""
# 根据规则类别选择合适的适配器
adapter_to_use = self.adapters[0] # 默认使用第一个适配器
# 如果是YAML格式的规则,使用YAML适配器
if hasattr(rule, 'code') and rule.code:
for adapter in self.adapters:
if isinstance(adapter, YAMLAdapter):
adapter_to_use = adapter
break
result = adapter_to_use.save_rule(rule)
# 如果成功保存且启用了预加载,更新缓存
if result and self.config.preload_rules:
if rule.id not in self.rule_cache:
self.rule_cache[rule.id] = {}
self.rule_cache[rule.id][rule.version] = rule
return result
def delete_rule(self, rule_id: str, version: Optional[str] = None) -> bool:
"""
从存储中删除规则。
Args:
rule_id: 规则ID
version: 如果指定,仅删除该版本;否则删除所有版本
Returns:
操作是否成功
"""
# 从所有适配器删除规则
overall_result = True
for adapter in self.adapters:
try:
result = adapter.delete_rule(rule_id, version)
if not result:
overall_result = False
except Exception as e:
self.logger.error(f"Error deleting rule {rule_id} (version={version}) from adapter {adapter.__class__.__name__}: {e}")
overall_result = False
# 如果启用了预加载,更新缓存
if self.config.preload_rules:
if version and rule_id in self.rule_cache:
# 删除特定版本
if version in self.rule_cache[rule_id]:
del self.rule_cache[rule_id][version]
# 如果该规则没有更多版本,删除整个规则条目
if not self.rule_cache[rule_id]:
del self.rule_cache[rule_id]
elif rule_id in self.rule_cache:
# 删除所有版本
del self.rule_cache[rule_id]
return overall_result
def get_rules_for_target(self, target_type: TargetType, target_id: str) -> List[AnyRule]:
"""
获取适用于特定目标的规则。
这是一个便捷方法,用于当TestExecutor或JSONSchemaValidator需要找到适用于特定API操作或数据对象的规则。
Args:
target_type: 目标类型(如APIRequest, APIResponse, DataObject
target_id: 目标标识符(如API操作ID, 数据对象名称)
Returns:
适用于该目标的规则列表
"""
query = RuleQuery(
target_type=target_type,
target_identifier=target_id,
is_enabled=True
)
return self.query_rules(query)
def get_rules_by_lifecycle(self, lifecycle: RuleLifecycle, target_type: Optional[TargetType] = None) -> List[AnyRule]:
"""
获取适用于特定生命周期阶段的规则。
Args:
lifecycle: 规则适用的生命周期阶段
target_type: 可选的目标类型过滤
Returns:
适用于该生命周期阶段的规则列表
"""
query = RuleQuery(
lifecycle=lifecycle,
target_type=target_type,
is_enabled=True
)
return self.query_rules(query)
def get_rules_by_scope(self, scope: RuleScope, target_type: Optional[TargetType] = None) -> List[AnyRule]:
"""
获取适用于特定作用域的规则。
Args:
scope: 规则的作用域
target_type: 可选的目标类型过滤
Returns:
适用于该作用域的规则列表
"""
query = RuleQuery(
scope=scope,
target_type=target_type,
is_enabled=True
)
return self.query_rules(query)
def get_schema_for_target(self, target_type: TargetType, target_id: str) -> Optional[Dict[str, Any]]:
"""
获取适用于特定目标的JSON Schema。
这是一个便捷方法,用于当JSONSchemaValidator需要找到适用于特定API操作或数据对象的JSON Schema。
Args:
target_type: 目标类型(如APIRequest, APIResponse, DataObject
target_id: 目标标识符(如API操作ID, 数据对象名称)
Returns:
JSON Schema字典,如果未找到则返回None
"""
query = RuleQuery(
category=RuleCategory.JSON_SCHEMA,
target_type=target_type,
target_identifier=target_id,
is_enabled=True
)
schemas = self.query_rules(query)
if not schemas:
return None
# 如果有多个匹配的Schema规则,使用最新版本的
# 注意:这里可以根据需要实现更复杂的选择逻辑
schemas.sort(key=lambda x: x.version)
latest_schema = schemas[-1]
# 假设是JSONSchemaDefinition类型
if hasattr(latest_schema, 'schema_content'):
return latest_schema.schema_content
return None
@@ -1,333 +0,0 @@
"""YAML adapter for rule storage using YAML files."""
import os
import yaml
import glob
from typing import List, Dict, Optional, Any, Tuple, Type, Union, cast
import logging
from datetime import datetime
from .adapters.base_adapter import BaseRuleStorageAdapter
from ..models.rule_models import (
AnyRule, RuleQuery, BaseRule, RuleCategory, TargetType, RuleLifecycle, RuleScope
)
from .adapters.rule_adapter_utils import parse_rule_data
class YAMLAdapter(BaseRuleStorageAdapter):
"""
基于YAML文件的规则存储适配器,使用YAML文件保存规则。
文件结构约定:
规则文件将按以下方式组织:
rules/
yaml_rules/
json_schemas/
rule_id1/
1.0.0.yaml
1.1.0.yaml
rule_id2/
1.0.0.yaml
business_logic/
rule_id3/
1.0.0.yaml
...
"""
def __init__(self, base_path: str = "./rules", file_pattern: str = "*.yaml"):
"""
初始化适配器。
Args:
base_path: 存储规则的基本目录路径。
file_pattern: 匹配规则文件的glob模式。
"""
self.base_path = os.path.abspath(base_path)
self.file_pattern = file_pattern
self.logger = logging.getLogger(__name__)
self.yaml_dir = os.path.join(self.base_path, "yaml_rules")
def initialize(self) -> None:
"""确保基本目录存在,并验证其可访问性。"""
if not os.path.exists(self.yaml_dir):
try:
os.makedirs(self.yaml_dir)
self.logger.info(f"Created YAML rules directory at {self.yaml_dir}")
except Exception as e:
self.logger.error(f"Failed to create YAML rules directory at {self.yaml_dir}: {e}")
raise ValueError(f"YAML rules directory {self.yaml_dir} does not exist and could not be created")
if not os.access(self.yaml_dir, os.R_OK | os.W_OK):
self.logger.error(f"YAML rules directory {self.yaml_dir} is not readable and writable")
raise ValueError(f"YAML rules directory {self.yaml_dir} is not readable and writable")
# 确保每个规则类别的子目录存在
for category in RuleCategory:
category_dir = self._get_category_dir(category)
if not os.path.exists(category_dir):
try:
os.makedirs(category_dir)
self.logger.debug(f"Created category directory at {category_dir}")
except Exception as e:
self.logger.warning(f"Failed to create category directory at {category_dir}: {e}")
def _get_category_dir(self, category: RuleCategory) -> str:
"""获取给定规则类别的目录路径。"""
return os.path.join(self.yaml_dir, category.value.lower())
def _get_rule_dir(self, rule_id: str, category: Optional[RuleCategory] = None) -> str:
"""
获取给定规则ID的目录路径。
如果指定了类别,则直接使用该类别的目录;否则尝试在所有类别中查找。
"""
if category:
return os.path.join(self._get_category_dir(category), rule_id)
# 如果未指定类别,检查所有类别目录
for cat in RuleCategory:
rule_dir = os.path.join(self._get_category_dir(cat), rule_id)
if os.path.exists(rule_dir):
return rule_dir
# 如果未找到任何匹配项,默认返回通用类别目录
return os.path.join(self._get_category_dir(RuleCategory.GENERIC), rule_id)
def _get_rule_file_path(self, rule_id: str, version: str, category: Optional[RuleCategory] = None) -> str:
"""获取给定规则ID和版本的文件路径。"""
rule_dir = self._get_rule_dir(rule_id, category)
return os.path.join(rule_dir, f"{version}.yaml")
def _get_rule_from_file(self, file_path: str) -> Optional[AnyRule]:
"""从YAML文件加载规则。"""
if not os.path.exists(file_path):
self.logger.debug(f"Rule file {file_path} does not exist")
return None
try:
with open(file_path, 'r', encoding='utf-8') as f:
raw_data = yaml.safe_load(f)
# 使用工具函数解析规则数据
return parse_rule_data(raw_data)
except yaml.YAMLError as e:
self.logger.error(f"Failed to parse YAML from {file_path}: {e}")
return None
except Exception as e:
self.logger.error(f"Error loading rule from {file_path}: {e}")
return None
def _save_rule_to_file(self, rule: BaseRule, file_path: str) -> bool:
"""将规则保存到YAML文件。"""
try:
# 确保目录存在
dir_path = os.path.dirname(file_path)
if not os.path.exists(dir_path):
os.makedirs(dir_path)
# 将规则序列化为YAML并写入文件
with open(file_path, 'w', encoding='utf-8') as f:
yaml.dump(rule.model_dump(), f, default_flow_style=False, sort_keys=False)
return True
except Exception as e:
self.logger.error(f"Failed to save rule to {file_path}: {e}")
return False
def _get_latest_version(self, rule_id: str, category: Optional[RuleCategory] = None) -> Optional[str]:
"""获取给定规则ID的最新版本。"""
versions = self.get_rule_versions(rule_id, category)
if not versions:
return None
# 简单地按字符串排序,假设版本格式是类似于"1.0.0"的语义化版本
versions.sort()
return versions[-1]
def get_rule_versions(self, rule_id: str, category: Optional[RuleCategory] = None) -> List[str]:
"""获取给定规则ID的所有版本。"""
rule_dir = self._get_rule_dir(rule_id, category)
if not os.path.exists(rule_dir):
return []
# 查找目录中所有匹配的YAML文件,并提取版本号(文件名)
pattern = os.path.join(rule_dir, self.file_pattern)
version_files = glob.glob(pattern)
versions = []
for vf in version_files:
version = os.path.splitext(os.path.basename(vf))[0] # 移除.yaml扩展名
versions.append(version)
return versions
def load_rule_by_id(self, rule_id: str, version: Optional[str] = None) -> Optional[AnyRule]:
"""
按ID加载单个规则。
Args:
rule_id: 规则的唯一标识符。
version: 可选的版本标识符。如果未提供,则加载最新版本。
Returns:
找到的规则对象,或者如果未找到则返回None。
"""
if not version:
version = self._get_latest_version(rule_id)
if not version:
self.logger.debug(f"No versions found for rule ID {rule_id}")
return None
file_path = self._get_rule_file_path(rule_id, version)
return self._get_rule_from_file(file_path)
def query_rules(self, query: RuleQuery) -> List[AnyRule]:
"""
根据查询条件查询规则。
Args:
query: 包含筛选条件的查询对象。
Returns:
匹配查询条件的规则列表。
"""
results = []
# 如果指定了规则ID,直接加载该规则
if query.rule_id:
rule = self.load_rule_by_id(query.rule_id, query.version if query.version != "latest" else None)
if rule and self._rule_matches_query(rule, query):
results.append(rule)
return results
# 否则,根据查询条件扫描规则文件
categories_to_search = [query.category] if query.category else list(RuleCategory)
for category in categories_to_search:
category_dir = self._get_category_dir(category)
if not os.path.exists(category_dir):
continue
# 获取该类别下的所有规则ID(子目录)
rule_dirs = [d for d in os.listdir(category_dir)
if os.path.isdir(os.path.join(category_dir, d))]
for rule_id in rule_dirs:
# 对于每个规则ID,加载指定版本或最新版本
if query.version and query.version != "latest":
file_path = self._get_rule_file_path(rule_id, query.version, category)
rule = self._get_rule_from_file(file_path)
if rule and self._rule_matches_query(rule, query):
results.append(rule)
else:
latest_version = self._get_latest_version(rule_id, category)
if latest_version:
file_path = self._get_rule_file_path(rule_id, latest_version, category)
rule = self._get_rule_from_file(file_path)
if rule and self._rule_matches_query(rule, query):
results.append(rule)
return results
def _rule_matches_query(self, rule: BaseRule, query: RuleQuery) -> bool:
"""检查规则是否匹配查询条件。"""
# 检查是否启用
if query.is_enabled is not None and rule.is_enabled != query.is_enabled:
return False
# 检查目标类型
if query.target_type and rule.target_type != query.target_type:
return False
# 检查目标标识符
if query.target_identifier and rule.target_identifier != query.target_identifier:
return False
# 检查标签
if query.tags:
if not rule.tags:
return False
# 检查是否所有查询标签都在规则标签中
if not all(tag in rule.tags for tag in query.tags):
return False
# 检查生命周期
if query.lifecycle and rule.lifecycle != query.lifecycle:
return False
# 检查作用域
if query.scope and rule.scope != query.scope:
return False
return True
def save_rule(self, rule: BaseRule) -> bool:
"""
保存规则到文件系统。
Args:
rule: 要保存的规则对象。
Returns:
操作是否成功。
"""
file_path = self._get_rule_file_path(rule.id, rule.version, rule.category)
return self._save_rule_to_file(rule, file_path)
def delete_rule(self, rule_id: str, version: Optional[str] = None) -> bool:
"""
删除规则。
Args:
rule_id: 规则的唯一标识符。
version: 如果提供,仅删除该版本;否则删除所有版本。
Returns:
操作是否成功。
"""
try:
if version:
# 删除特定版本
file_path = self._get_rule_file_path(rule_id, version)
if os.path.exists(file_path):
os.remove(file_path)
self.logger.info(f"Deleted rule file: {file_path}")
else:
self.logger.warning(f"Rule file not found for deletion: {file_path}")
return False
else:
# 删除所有版本(整个规则目录)
rule_dir = self._get_rule_dir(rule_id)
if os.path.exists(rule_dir):
# 递归删除目录及其内容
for root, dirs, files in os.walk(rule_dir, topdown=False):
for file in files:
os.remove(os.path.join(root, file))
for dir in dirs:
os.rmdir(os.path.join(root, dir))
os.rmdir(rule_dir)
self.logger.info(f"Deleted rule directory: {rule_dir}")
else:
self.logger.warning(f"Rule directory not found for deletion: {rule_dir}")
return False
return True
except Exception as e:
self.logger.error(f"Error deleting rule {rule_id} (version={version}): {e}")
return False
def list_all_rule_ids(self) -> List[str]:
"""列出所有规则ID。"""
all_rule_ids = set()
# 扫描所有类别目录
for category in RuleCategory:
category_dir = self._get_category_dir(category)
if not os.path.exists(category_dir):
continue
# 获取该类别下的所有规则ID(子目录)
rule_dirs = [d for d in os.listdir(category_dir)
if os.path.isdir(os.path.join(category_dir, d))]
all_rule_ids.update(rule_dirs)
return list(all_rule_ids)
-100
View File
@@ -1,100 +0,0 @@
# ddms_compliance_suite/test_loader/loader.py
from abc import ABC, abstractmethod
from pathlib import Path
from typing import List, Union, Dict, Any
import yaml # PyYAML
import json
from ..models.test_models import TestSuite # Ensure correct path
# from ..models.test_models import TestCase # If TestCase can be loaded standalone
class LoadError(Exception):
"""Custom exception for errors during test case loading."""
pass
class BaseTestCaseLoader(ABC):
@abstractmethod
def load_suites_from_file(self, file_path: Union[str, Path]) -> List[TestSuite]:
"""Loads a list of TestSuite objects from a single file."""
pass
def load_suites_from_directory(self, directory_path: Union[str, Path], recursive: bool = False, pattern: str = "*.yaml") -> List[TestSuite]:
"""Loads TestSuite objects from all matching files in a directory."""
suites: List[TestSuite] = []
p = Path(directory_path)
if not p.is_dir():
raise LoadError(f"Directory not found: {directory_path}")
file_paths = list(p.rglob(pattern)) if recursive else list(p.glob(pattern))
for file_path in file_paths:
if file_path.is_file():
try:
# Ensure that the loader instance calls its own method
# This part might need adjustment if called directly on BaseTestCaseLoader
# For now, assuming it's called from a concrete instance like YAMLTestCaseLoader
suites.extend(self.load_suites_from_file(file_path))
except LoadError as e:
# Consider using a logger for warnings/errors
print(f"Warning: Could not load test suites from {file_path}: {e}")
except Exception as e:
print(f"Warning: An unexpected error occurred loading {file_path}: {e}")
return suites
class YAMLTestCaseLoader(BaseTestCaseLoader):
def load_suites_from_file(self, file_path: Union[str, Path]) -> List[TestSuite]:
file_p = Path(file_path)
if not file_p.exists() or not file_p.is_file():
raise LoadError(f"Test case file not found or is not a file: {file_path}")
try:
with open(file_p, 'r', encoding='utf-8') as f:
data = yaml.safe_load(f)
except yaml.YAMLError as e:
raise LoadError(f"Error parsing YAML file {file_p}: {e}")
except Exception as e:
raise LoadError(f"An unexpected error occurred while reading {file_p}: {e}")
if data is None: # Handle empty YAML file
return []
if not isinstance(data, dict) or "test_suites" not in data:
raise LoadError(f"Invalid format in {file_p}: Missing 'test_suites' top-level key or not a dictionary.")
suite_data_list = data.get("test_suites") # Use .get for safer access
if not isinstance(suite_data_list, list):
raise LoadError(f"Invalid format in {file_p}: 'test_suites' must be a list.")
loaded_suites: List[TestSuite] = []
for i, suite_data in enumerate(suite_data_list):
if not isinstance(suite_data, dict):
print(f"Warning: Suite data #{i+1} in {file_p} is not a dictionary, skipping.")
continue
try:
suite = TestSuite.model_validate(suite_data) # For Pydantic v2+
# For Pydantic v1, use: suite = TestSuite.parse_obj(suite_data)
loaded_suites.append(suite)
except Exception as e: # Catch Pydantic validation errors and others
# It's often better to log this and continue, or collect all errors
raise LoadError(f"Error validating test suite #{i+1} (ID: {suite_data.get('id', 'N/A')}) in {file_p}: {e}")
return loaded_suites
# Example Usage (conceptual):
# if __name__ == '__main__':
# yaml_loader = YAMLTestCaseLoader()
# try:
# # suites_from_single_file = yaml_loader.load_suites_from_file('path/to/your/single_test_suite_file.yaml')
# # for suite in suites_from_single_file:
# # print(f"Loaded Suite: {suite.name}")
# all_suites_in_dir = yaml_loader.load_suites_from_directory('path/to/test_suites_directory', recursive=True, pattern="*.test_suite.yaml")
# for suite in all_suites_in_dir:
# print(f"Loaded Suite from Dir: {suite.name}")
# for tc in suite.test_cases:
# print(f" - TestCase: {tc.name}")
# except LoadError as e:
# print(f"Loading Error: {e}")