mirror of
https://github.com/volatilityfoundation/volatility3.git
synced 2026-08-17 20:35:40 +02:00
Testing: Adds validation of vol3 imports in check script
This checks `ast.ImportFrom` statements to see if anything other than modules are being imported in this way. It enumerates all instances of this and suggests a fix.
This commit is contained in:
@@ -22,6 +22,7 @@ import inspect
|
||||
import logging
|
||||
import pkgutil
|
||||
import sys
|
||||
import traceback
|
||||
import types
|
||||
from typing import Any, Iterator, List, Optional, Tuple, Type, Union
|
||||
|
||||
@@ -105,6 +106,36 @@ class UnrequiredVersionableUsage(CodeViolation):
|
||||
)
|
||||
|
||||
|
||||
class DirectVolatilityImportUsage(CodeViolation):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
module: types.ModuleType,
|
||||
node: ast.AST,
|
||||
importing_module: str,
|
||||
imported_item: object,
|
||||
imported_name: str,
|
||||
) -> None:
|
||||
self.imported_item = imported_item
|
||||
self.imported_name = imported_name
|
||||
self.importing_module = importing_module
|
||||
super().__init__(module, node)
|
||||
|
||||
def __str__(self) -> str:
|
||||
components = self.importing_module.split(".")
|
||||
return (
|
||||
super().__str__()
|
||||
+ ": "
|
||||
+ (
|
||||
f"Direct import of {self.imported_name} "
|
||||
f"({type(self.imported_item)}) "
|
||||
f"from module {self.importing_module} - "
|
||||
"change to "
|
||||
f"'from {'.'.join(components[:-1])} import {components[-1]} and using {components[-1]}.{self.imported_name}"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def is_versionable(var):
|
||||
try:
|
||||
return (
|
||||
@@ -134,6 +165,39 @@ class ModuleVisitor(NodeVisitor):
|
||||
def violations(self):
|
||||
return self._violations
|
||||
|
||||
def enter_ImportFrom(self, node: ast.ImportFrom):
|
||||
if not node.module:
|
||||
return
|
||||
|
||||
if (
|
||||
node.module
|
||||
and node.module.startswith("volatility3")
|
||||
and node.module != "volatility3.framework.constants._version" # make an exception for this
|
||||
):
|
||||
for name in node.names:
|
||||
try:
|
||||
item = vars(self._module)[
|
||||
name.asname if name.asname is not None else name.name
|
||||
]
|
||||
except KeyError:
|
||||
logger.debug(
|
||||
"Couldn't find imported name %s in module %s",
|
||||
name.asname or name.name,
|
||||
self._module.__name__,
|
||||
)
|
||||
continue
|
||||
|
||||
if not (isinstance(item, types.ModuleType) or inspect.isfunction(item)):
|
||||
self._violations.append(
|
||||
DirectVolatilityImportUsage(
|
||||
self._module,
|
||||
node,
|
||||
node.module,
|
||||
item,
|
||||
name.asname or name.name,
|
||||
)
|
||||
)
|
||||
|
||||
def enter_ClassDef(self, node: ast.ClassDef) -> Any:
|
||||
logger.debug("Entering class %s", node.name)
|
||||
clazz = None
|
||||
@@ -304,6 +368,7 @@ def report_missing_requirements() -> Iterator[Tuple[str, UnrequiredVersionableUs
|
||||
modname,
|
||||
str(exc),
|
||||
)
|
||||
traceback.print_exc()
|
||||
continue
|
||||
|
||||
logger.info("Checking module %s", plugin_module.__name__)
|
||||
@@ -344,9 +409,7 @@ def perform_review():
|
||||
print(str(usage))
|
||||
|
||||
if found:
|
||||
print(
|
||||
f"Found {found} issues"
|
||||
)
|
||||
print(f"Found {found} issues")
|
||||
sys.exit(1)
|
||||
|
||||
print("All configurable classes passed validation!")
|
||||
|
||||
Reference in New Issue
Block a user