diff --git a/scripts/generate_psa_wrappers.py b/scripts/generate_psa_wrappers.py index 29cb4e3fd..01baad1d6 100755 --- a/scripts/generate_psa_wrappers.py +++ b/scripts/generate_psa_wrappers.py @@ -7,32 +7,40 @@ import argparse from mbedtls_framework.code_wrapper.psa_test_wrapper import PSATestWrapper, PSALoggingTestWrapper - -DEFAULT_C_OUTPUT_FILE_NAME = 'tests/src/psa_test_wrappers.c' -DEFAULT_H_OUTPUT_FILE_NAME = 'tests/include/test/psa_test_wrappers.h' +from mbedtls_framework import build_tree def main() -> None: + default_c_output_file_name = 'tests/src/psa_test_wrappers.c' + default_h_output_file_name = 'tests/include/test/psa_test_wrappers.h' + + project_root = build_tree.guess_project_root() + if build_tree.looks_like_mbedtls_root(project_root) and \ + not build_tree.is_mbedtls_3_6(): + default_c_output_file_name = 'tf-psa-crypto/' + default_c_output_file_name + default_h_output_file_name = 'tf-psa-crypto/' + default_h_output_file_name + parser = argparse.ArgumentParser(description=globals()['__doc__']) parser.add_argument('--log', help='Stream to log to (default: no logging code)') parser.add_argument('--output-c', metavar='FILENAME', - default=DEFAULT_C_OUTPUT_FILE_NAME, + default=default_c_output_file_name, help=('Output .c file path (default: {}; skip .c output if empty)' - .format(DEFAULT_C_OUTPUT_FILE_NAME))) + .format(default_c_output_file_name))) parser.add_argument('--output-h', metavar='FILENAME', - default=DEFAULT_H_OUTPUT_FILE_NAME, + default=default_h_output_file_name, help=('Output .h file path (default: {}; skip .h output if empty)' - .format(DEFAULT_H_OUTPUT_FILE_NAME))) + .format(default_h_output_file_name))) options = parser.parse_args() + if options.log: - generator = PSALoggingTestWrapper(DEFAULT_H_OUTPUT_FILE_NAME, - DEFAULT_C_OUTPUT_FILE_NAME, + generator = PSALoggingTestWrapper(default_h_output_file_name, + default_c_output_file_name, options.log) #type: PSATestWrapper else: - generator = PSATestWrapper(DEFAULT_H_OUTPUT_FILE_NAME, - DEFAULT_C_OUTPUT_FILE_NAME) + generator = PSATestWrapper(default_h_output_file_name, + default_c_output_file_name) if options.output_h: generator.write_h_file(options.output_h) diff --git a/scripts/generate_test_keys.py b/scripts/generate_test_keys.py index ace328c88..f5d69019e 100755 --- a/scripts/generate_test_keys.py +++ b/scripts/generate_test_keys.py @@ -168,7 +168,7 @@ def collect_keys() -> Tuple[str, str]: return ''.join(arrays), '\n'.join(look_up_table) def main() -> None: - default_output_path = guess_project_root() + "/framework/tests/src/test_keys.h" + default_output_path = guess_project_root() + "/framework/tests/include/test/test_keys.h" argparser = argparse.ArgumentParser() argparser.add_argument("--output", help="Output file", default=default_output_path)