diff --git a/__init__.py b/__init__.py index 4b47670..32f647b 100644 --- a/__init__.py +++ b/__init__.py @@ -1,3 +1,11 @@ -from .apply_sdxl_style import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +# Copyright (c) 2023, Manfred Moitzi +# License: MIT License +from __future__ import annotations +from . import apply_sdxl_style __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] + +NODE_CLASS_MAPPINGS: dict[str, type] = {} +NODE_DISPLAY_NAME_MAPPINGS: dict[str, str] = {} + +apply_sdxl_style.setup_nodes(NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS) \ No newline at end of file diff --git a/apply_sdxl_style.py b/apply_sdxl_style.py index 75ddd84..0040138 100644 --- a/apply_sdxl_style.py +++ b/apply_sdxl_style.py @@ -5,6 +5,7 @@ import pathlib import dataclasses import json +__all__ = ["setup_nodes"] BYPASS = "bypass" BYPASS_PROMPT = "{prompt}" @@ -118,8 +119,6 @@ class ApplyStyle: return output_positive, output_negative -NODE_CLASS_MAPPINGS = {} -NODE_DISPLAY_NAME_MAPPINGS = {} NODE_DEFINITIONS = [ ("sdxl_styles_sai.json", "Apply SDXL Style SAI"), ("sdxl_styles_twri.json", "Apply SDXL Style TWRI"), @@ -131,17 +130,16 @@ NODE_DEFINITIONS = [ ] -def _setup_classes(): +def setup_nodes( + class_mapping: dict[str, type], display_name_mapping: dict[str, str] +) -> None: cwd = pathlib.Path(__file__).parent for file_name, display_name in NODE_DEFINITIONS: templates = load_style_templates(cwd / "styles" / file_name) class_name = display_name.replace(" ", "") # create classes dynamically: - NODE_CLASS_MAPPINGS[class_name] = type( + class_mapping[class_name] = type( class_name, (ApplyStyle,), {"TEMPLATES": templates} ) - NODE_DISPLAY_NAME_MAPPINGS[class_name] = display_name - - -_setup_classes() + display_name_mapping[class_name] = display_name diff --git a/test_apply_sdxl_styles.py b/test_apply_sdxl_styles.py index a13ba3d..c488236 100644 --- a/test_apply_sdxl_styles.py +++ b/test_apply_sdxl_styles.py @@ -3,11 +3,8 @@ import pytest import re -from .apply_sdxl_style import ( - NODE_CLASS_MAPPINGS, - NODE_DISPLAY_NAME_MAPPINGS, - NODE_DEFINITIONS, -) +from . import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .apply_sdxl_style import NODE_DEFINITIONS def test_all_classes_created(): @@ -25,6 +22,7 @@ def test_classes_have_templates(): for cls in NODE_CLASS_MAPPINGS.values(): assert len(cls.TEMPLATES) > 1, "no styles loaded - correct filename?" + def test_templates_have_bypass_style(): """The style "bypass" will be added automatically to every style file.""" for cls in NODE_CLASS_MAPPINGS.values():