[UPD] script search class model: remove astor and update code, add parameters

This commit is contained in:
Mathieu Benoit 2025-11-11 09:34:14 -05:00
parent 01d92d3c8e
commit 32fe737bd8
2 changed files with 30 additions and 21 deletions

View file

@ -67,6 +67,3 @@ git+https://github.com/prauscher/pysaml2.git@replace-pyopenssl
# Fix compilation, until pymssql==2.3.9 is release # Fix compilation, until pymssql==2.3.9 is release
git+https://github.com/pymssql/pymssql.git git+https://github.com/pymssql/pymssql.git
# Force last version, because too old into pip
git+https://github.com/berkerpeksag/astor.git

View file

@ -10,8 +10,6 @@ import os
import sys import sys
from pathlib import Path from pathlib import Path
import astor
logging.basicConfig(level=logging.DEBUG) logging.basicConfig(level=logging.DEBUG)
_logger = logging.getLogger(__name__) _logger = logging.getLogger(__name__)
@ -57,6 +55,16 @@ def get_config():
action="store_true", action="store_true",
help="Return result in json", help="Return result in json",
) )
parser.add_argument(
"--format_json",
action="store_true",
help="If --json is enabled, will format json for human reader.",
)
parser.add_argument(
"--show_error_json",
action="store_true",
help="Will show error when --json is enabled.",
)
parser.add_argument( parser.add_argument(
"--extract_field", "--extract_field",
action="store_true", action="store_true",
@ -124,7 +132,7 @@ def search_and_replace(
def extract_lambda(node): def extract_lambda(node):
result = astor.to_source(node).strip().replace("\n", "") result = ast.unparse(node)
if result[0] == "(" and result[-1] == ")": if result[0] == "(" and result[-1] == ")":
result = result[1:-1] result = result[1:-1]
return result return result
@ -134,13 +142,9 @@ def fill_search_field(ast_obj, var_name="", py_filename=""):
ast_obj_type = type(ast_obj) ast_obj_type = type(ast_obj)
result = None result = None
if ast_obj_type is ast.Constant: if ast_obj_type is ast.Constant:
result = ast_obj.s result = ast_obj.value
elif ast_obj_type is ast.Lambda: elif ast_obj_type is ast.Lambda:
result = extract_lambda(ast_obj) result = extract_lambda(ast_obj)
elif ast_obj_type is ast.NameConstant:
result = ast_obj.value
elif ast_obj_type is ast.Num:
result = ast_obj.n
elif ast_obj_type is ast.UnaryOp: elif ast_obj_type is ast.UnaryOp:
if type(ast_obj.op) is ast.USub: if type(ast_obj.op) is ast.USub:
# value is negative # value is negative
@ -232,15 +236,15 @@ def main():
lst_search_target lst_search_target
and node.targets[0].id in lst_search_target and node.targets[0].id in lst_search_target
): ):
if node.value.s in lst_model_name: if node.value.value in lst_model_name:
is_duplicated = True is_duplicated = True
_logger.warning( _logger.warning(
"Duplicated model name" "Duplicated model name"
f" {node.value.s} from file {py_file}" f" {node.value.value} from file {py_file}"
) )
else: else:
model_name = node.value.s model_name = node.value.value
lst_model_name.append(node.value.s) lst_model_name.append(node.value.value)
if ( if (
lst_search_inherit_target lst_search_inherit_target
@ -248,14 +252,16 @@ def main():
in lst_search_inherit_target in lst_search_inherit_target
): ):
is_inherit = True is_inherit = True
if node.value.s in lst_model_inherit_name: if node.value.value in lst_model_inherit_name:
_logger.warning( _logger.warning(
"Duplicated model inherit name" "Duplicated model inherit name"
f" {node.value.s} from file {py_file}" f" {node.value.value} from file {py_file}"
) )
else: else:
model_name = node.value.s model_name = node.value.value
lst_model_inherit_name.append(node.value.s) lst_model_inherit_name.append(
node.value.value
)
dct_fields = {} dct_fields = {}
if model_name: if model_name:
dct_model[model_name] = { dct_model[model_name] = {
@ -267,7 +273,7 @@ def main():
# TODO do it! # TODO do it!
if model_name and ( if model_name and (
type(node.value) is ast.Constant type(node.value) is ast.Constant
and node.value.s == model_name and node.value.value == model_name
or type(node.value) is ast.List or type(node.value) is ast.List
and model_name and model_name
in [a.s for a in node.value.elts] in [a.s for a in node.value.elts]
@ -312,7 +318,13 @@ def main():
print(models_name) print(models_name)
print(models_inherit_name) print(models_inherit_name)
else: else:
output = json.dumps(dct_model) if config.show_error_json and not models_name:
_logger.warning(f"Missing models class in {config.directory}")
if config.format_json:
output = json.dumps(dct_model, indent=4, sort_keys=True)
else:
output = json.dumps(dct_model)
print(output) print(output)
if config.template_dir: if config.template_dir: