Files
AEmotionStudio-ComfyUI-Disc…/tests/test_image_node_sanitization.py
T
AEmotionStudio 87112b8b2a Fix CI tests: Improve test isolation and mock handling
- Add numpy to requirements-nodes.txt
- Create conftest.py for test configuration
- Fix test_numpy_subprocess.py to import real numpy before any mocking
- Fix test_media.py to import real numpy early
- Fix test_image_node_sanitization.py to use MockTensor instead of torch.zeros
2026-01-21 00:46:02 -08:00

153 lines
5.3 KiB
Python

import unittest
import json
import sys
import os
from unittest.mock import MagicMock, patch
import importlib
# Add project root to sys.path
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
# Clean up any potential pollution from other tests before we start
if 'PIL' in sys.modules:
# Check if it's a mock
if isinstance(sys.modules['PIL'], MagicMock):
del sys.modules['PIL']
if 'PIL.PngImagePlugin' in sys.modules:
del sys.modules['PIL.PngImagePlugin']
if 'PIL.Image' in sys.modules:
del sys.modules['PIL.Image']
# Now we can import real modules or mock them as we see fit LOCALLY
# But wait, discord_image_node imports them at module level.
# So we need to ensure environment is set up before importing it.
# Mock comfy modules
sys.modules['comfy'] = MagicMock()
sys.modules['comfy.cli_args'] = MagicMock()
sys.modules['comfy.cli_args'].args = MagicMock()
sys.modules['comfy.cli_args'].args.disable_metadata = False
sys.modules['comfy.utils'] = MagicMock()
sys.modules['folder_paths'] = MagicMock()
sys.modules['folder_paths'].get_output_directory = MagicMock(return_value="/tmp")
sys.modules['folder_paths'].get_temp_directory = MagicMock(return_value="/tmp")
sys.modules['folder_paths'].get_save_image_path = MagicMock(return_value=("/tmp", "test", 0, "", "test"))
sys.modules['server'] = MagicMock()
# We need real PIL for this test to verify PngInfo
try:
import PIL.PngImagePlugin
except ImportError:
# If it failed because it was mocked out and we deleted it, reload
pass
from nodes.image_node import DiscordSendSaveImage
class TestDiscordImageNodeOptimization(unittest.TestCase):
def setUp(self):
self.node = DiscordSendSaveImage()
self.webhook_url = "https://discord.com/api/webhooks/12345/abcdef"
self.github_token = "ghp_sensitive12345"
def test_save_images_sanitization(self):
# Create a mock image tensor using numpy (torch not available in CI)
# The image_node iterates over images and accesses shape, so we need
# an object that supports iteration and has proper shape
import numpy as np
# Create a simple class that mimics torch.Tensor behavior for the node
class MockTensor:
def __init__(self, data):
self._data = data
self.shape = data.shape
def __len__(self):
return len(self._data)
def __getitem__(self, idx):
return self._data[idx]
def __iter__(self):
return iter(self._data)
# Create a 1x64x64x3 "image batch" using numpy
image_data = np.zeros((1, 64, 64, 3), dtype=np.float32)
mock_image = MockTensor(image_data)
# Create prompt and extra_pnginfo with sensitive data
prompt = {
"3": {
"inputs": {
"webhook_url": self.webhook_url,
"github_token": self.github_token,
"seed": 123
},
"class_type": "DiscordSendSaveImage"
}
}
extra_pnginfo = {
"workflow": {
"nodes": [
{
"id": 3,
"type": "DiscordSendSaveImage",
"widgets_values": [self.webhook_url, "message", self.github_token]
}
]
}
}
# Mock Image.save to check metadata
with patch('PIL.Image.Image.save') as mock_save:
self.node.save_images(
images=mock_image,
prompt=prompt,
extra_pnginfo=extra_pnginfo,
save_output=True,
send_to_discord=False # Disable discord sending to focus on save/metadata
)
# Check if save was called
self.assertTrue(mock_save.called)
# Get the pnginfo passed to save
args, kwargs = mock_save.call_args
pnginfo = kwargs.get('pnginfo')
self.assertIsNotNone(pnginfo)
found_prompt = False
found_workflow = False
# Check chunks - PIL PngInfo internal structure
for tag_type, data, after_idat in pnginfo.chunks:
try:
# decode data
decoded = data.decode('latin-1')
except Exception:
# Skip chunks that can't be decoded
continue
if '\0' in decoded:
try:
k, v = decoded.split('\0', 1)
except ValueError:
continue
if k == "prompt":
found_prompt = True
# Verify sensitive data is gone
self.assertNotIn("discord.com/api/webhooks", v)
self.assertNotIn("ghp_", v)
if k == "workflow":
found_workflow = True
self.assertNotIn("discord.com/api/webhooks", v)
self.assertNotIn("ghp_", v)
self.assertTrue(found_prompt, "Prompt metadata not found")
self.assertTrue(found_workflow, "Workflow metadata not found")
if __name__ == "__main__":
unittest.main()