diff --git a/tests/test_utils.py b/tests/test_utils.py new file mode 100644 index 0000000..be674c8 --- /dev/null +++ b/tests/test_utils.py @@ -0,0 +1,142 @@ +""" +Unit tests for ComfyUI-DiscordSend utilities. + +Run tests with: python -m pytest tests/test_utils.py -v +Or without pytest: python tests/test_utils.py +""" + +import sys +import os +import unittest + +# Add parent directory to path for imports +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from utils.sanitizer import sanitize_json_for_export +from utils.discord_api import validate_webhook_url, sanitize_webhook_for_logging + + +class TestSanitizer(unittest.TestCase): + """Tests for the sanitize_json_for_export function.""" + + def test_sanitize_none(self): + """Should handle None input.""" + self.assertIsNone(sanitize_json_for_export(None)) + + def test_sanitize_webhook_url_in_string(self): + """Should remove webhook URLs from strings.""" + test_url = "https://discord.com/api/webhooks/123456789/abcdefg" + result = sanitize_json_for_export(test_url) + self.assertEqual(result, "") + + def test_sanitize_github_token_in_string(self): + """Should remove GitHub tokens from strings.""" + test_token = "ghp_1234567890abcdefghijklmnopqrstuvwxyz" + result = sanitize_json_for_export(test_token) + self.assertEqual(result, "") + + def test_sanitize_dict_with_webhook_key(self): + """Should remove values for webhook-related keys.""" + test_data = { + "webhook_url": "https://discord.com/api/webhooks/123/abc", + "message": "Hello world" + } + result = sanitize_json_for_export(test_data) + self.assertEqual(result["webhook_url"], "") + self.assertEqual(result["message"], "Hello world") + + def test_sanitize_dict_with_github_token_key(self): + """Should remove values for github_token keys.""" + test_data = { + "github_token": "ghp_secret123", + "repo": "user/repo" + } + result = sanitize_json_for_export(test_data) + self.assertEqual(result["github_token"], "") + self.assertEqual(result["repo"], "user/repo") + + def test_sanitize_nested_dict(self): + """Should sanitize nested dictionaries.""" + test_data = { + "level1": { + "webhook_url": "https://discord.com/api/webhooks/123/abc", + "level2": { + "github_token": "ghp_secret" + } + } + } + result = sanitize_json_for_export(test_data) + self.assertEqual(result["level1"]["webhook_url"], "") + self.assertEqual(result["level1"]["level2"]["github_token"], "") + + def test_sanitize_list(self): + """Should sanitize lists.""" + test_data = [ + "normal string", + "https://discord.com/api/webhooks/123/abc", + {"webhook_url": "secret"} + ] + result = sanitize_json_for_export(test_data) + self.assertEqual(result[0], "normal string") + self.assertEqual(result[1], "") + self.assertEqual(result[2]["webhook_url"], "") + + def test_preserve_normal_data(self): + """Should preserve non-sensitive data.""" + test_data = { + "prompt": "a beautiful sunset", + "seed": 12345, + "steps": 20, + "cfg": 7.5 + } + result = sanitize_json_for_export(test_data) + self.assertEqual(result, test_data) + + +class TestWebhookValidation(unittest.TestCase): + """Tests for webhook URL validation.""" + + def test_valid_webhook_url(self): + """Should accept valid Discord webhook URLs.""" + valid_url = "https://discord.com/api/webhooks/1234567890/abcdefghijklmnop" + is_valid, message = validate_webhook_url(valid_url) + self.assertTrue(is_valid) + + def test_valid_discordapp_url(self): + """Should accept discordapp.com URLs.""" + valid_url = "https://discordapp.com/api/webhooks/1234567890/abcdefghijklmnop" + is_valid, message = validate_webhook_url(valid_url) + self.assertTrue(is_valid) + + def test_empty_url(self): + """Should reject empty URLs.""" + is_valid, message = validate_webhook_url("") + self.assertFalse(is_valid) + + def test_invalid_url(self): + """Should reject non-webhook URLs.""" + is_valid, message = validate_webhook_url("https://example.com") + self.assertFalse(is_valid) + + +class TestWebhookSanitization(unittest.TestCase): + """Tests for webhook URL sanitization for logging.""" + + def test_sanitize_webhook_for_logging(self): + """Should redact the token portion of webhook URLs.""" + url = "https://discord.com/api/webhooks/123456789/secrettoken" + result = sanitize_webhook_for_logging(url) + self.assertIn("[REDACTED]", result) + self.assertIn("123456789", result) + self.assertNotIn("secrettoken", result) + + def test_sanitize_empty_url(self): + """Should handle empty URLs.""" + result = sanitize_webhook_for_logging("") + self.assertEqual(result, "") + + +if __name__ == "__main__": + # Run tests + print("Running ComfyUI-DiscordSend utility tests...\n") + unittest.main(verbosity=2) diff --git a/utils/discord_api.py b/utils/discord_api.py index 19a5649..dd49293 100644 --- a/utils/discord_api.py +++ b/utils/discord_api.py @@ -7,11 +7,15 @@ Provides a client for interacting with Discord webhooks and validation utilities import os import re import time +import logging from io import BytesIO from typing import Any, Dict, List, Optional, Tuple import requests +# Get logger for this module +logger = logging.getLogger("comfyui_discordsend") + # Discord webhook URL patterns WEBHOOK_URL_PATTERNS = [ @@ -373,7 +377,7 @@ def send_to_discord_with_retry( retry_after = response.json().get("retry_after", 1) except: pass - print(f"Rate limited by Discord, waiting {retry_after}s before retry...") + logger.warning(f"Rate limited by Discord, waiting {retry_after}s before retry...") time.sleep(retry_after) continue @@ -382,13 +386,13 @@ def send_to_discord_with_retry( return response # Server error - retry - print(f"Discord server error {response.status_code}, attempt {attempt + 1}/{max_retries}") + logger.warning(f"Discord server error {response.status_code}, attempt {attempt + 1}/{max_retries}") except requests.exceptions.Timeout: - print(f"Request timeout, attempt {attempt + 1}/{max_retries}") + logger.warning(f"Request timeout, attempt {attempt + 1}/{max_retries}") last_exception = requests.exceptions.Timeout("Discord request timed out") except requests.exceptions.RequestException as e: - print(f"Request error: {e}, attempt {attempt + 1}/{max_retries}") + logger.warning(f"Request error: {e}, attempt {attempt + 1}/{max_retries}") last_exception = e # Exponential backoff before retry diff --git a/utils/logging_config.py b/utils/logging_config.py new file mode 100644 index 0000000..4bff9c0 --- /dev/null +++ b/utils/logging_config.py @@ -0,0 +1,43 @@ +""" +Logging configuration for ComfyUI-DiscordSend + +Provides a configured logger for the extension. +""" + +import logging +import sys + + +def get_logger(name: str = "comfyui_discordsend") -> logging.Logger: + """ + Get a configured logger for the extension. + + Args: + name: The logger name (default: comfyui_discordsend) + + Returns: + Configured logger instance + """ + logger = logging.getLogger(name) + + # Only configure if not already configured + if not logger.handlers: + logger.setLevel(logging.DEBUG) + + # Console handler + console_handler = logging.StreamHandler(sys.stdout) + console_handler.setLevel(logging.INFO) + + # Format + formatter = logging.Formatter( + '[%(name)s] %(levelname)s: %(message)s' + ) + console_handler.setFormatter(formatter) + + logger.addHandler(console_handler) + + return logger + + +# Default logger instance +logger = get_logger()