fix: code formatting for Gemini prompt node

- Fix missing newlines at end of files
- Apply black formatting
- Remaining non-critical warnings for long lines in prompts
This commit is contained in:
Vito Sansevero
2025-08-01 09:41:21 -07:00
parent f559fe220e
commit ca504d5f74
5 changed files with 122 additions and 93 deletions
+1 -1
View File
@@ -2,4 +2,4 @@
from .node import GeminiPromptNode
__all__ = ["GeminiPromptNode"]
__all__ = ["GeminiPromptNode"]
+47 -37
View File
@@ -14,10 +14,10 @@ from .prompts import PROMPT_TEMPLATES
def tensor_to_pil(tensor: np.ndarray) -> Image.Image:
"""Convert ComfyUI tensor to PIL Image.
Args:
tensor: Input tensor in ComfyUI format (B, H, W, C)
Returns:
PIL Image object
"""
@@ -25,129 +25,139 @@ def tensor_to_pil(tensor: np.ndarray) -> Image.Image:
if tensor.ndim == 4:
# Take first image from batch
tensor = tensor[0]
# Convert to uint8
image_array = (tensor * 255).astype(np.uint8)
# Convert to PIL
return Image.fromarray(image_array, mode='RGB')
return Image.fromarray(image_array, mode="RGB")
def image_to_base64(image: Image.Image, format: str = "PNG") -> str:
"""Convert PIL Image to base64 string.
Args:
image: PIL Image object
format: Image format (PNG or JPEG)
Returns:
Base64 encoded string
"""
buffer = io.BytesIO()
image.save(buffer, format=format)
buffer.seek(0)
return base64.b64encode(buffer.read()).decode('utf-8')
return base64.b64encode(buffer.read()).decode("utf-8")
def get_api_key() -> Optional[str]:
"""Get Gemini API key from environment or config.
Returns:
API key string or None if not found
"""
# Check environment variable first
api_key = os.environ.get("GEMINI_API_KEY")
if not api_key:
# Check for config file in ComfyUI directory
try:
config_path = os.path.join(os.path.dirname(__file__), "..", "..", "..", "gemini_config.json")
config_path = os.path.join(
os.path.dirname(__file__), "..", "..", "..", "gemini_config.json"
)
if os.path.exists(config_path):
with open(config_path, 'r') as f:
with open(config_path, "r") as f:
config = json.load(f)
api_key = config.get("api_key")
except Exception:
pass
return api_key
def analyze_image_with_gemini(
image: np.ndarray,
prompt_type: str,
image: np.ndarray,
prompt_type: str,
api_key: Optional[str] = None,
custom_prompt: Optional[str] = None,
model_name: str = "gemini-1.5-flash"
model_name: str = "gemini-1.5-flash",
) -> Tuple[str, Optional[str]]:
"""Analyze image using Gemini API and generate appropriate prompt.
Args:
image: Input image tensor
prompt_type: Type of prompt to generate (flux, sdxl, danbooru, video)
api_key: Gemini API key (optional, will try to get from env/config)
custom_prompt: Custom system prompt to use instead of templates
model_name: Gemini model to use (default: gemini-1.5-flash)
Returns:
Tuple of (generated_prompt, error_message)
"""
# Get API key
if not api_key:
api_key = get_api_key()
if not api_key:
return "", "Gemini API key not found. Please set GEMINI_API_KEY environment variable or provide it in the node."
return (
"",
"Gemini API key not found. Please set GEMINI_API_KEY environment variable or provide it in the node.",
)
# Convert tensor to PIL image
try:
pil_image = tensor_to_pil(image)
except Exception as e:
return "", f"Failed to convert image: {str(e)}"
# Get system prompt
if custom_prompt:
system_prompt = custom_prompt
else:
system_prompt = PROMPT_TEMPLATES.get(prompt_type, PROMPT_TEMPLATES["flux"])
# Here we would normally make the API call to Gemini
# For now, we'll import the google-generativeai library
try:
import google.generativeai as genai
except ImportError:
return "", "google-generativeai library not installed. Please run: pip install google-generativeai"
return (
"",
"google-generativeai library not installed. Please run: pip install google-generativeai",
)
try:
# Configure Gemini
genai.configure(api_key=api_key)
# Create model
model = genai.GenerativeModel(model_name)
# Generate content
response = model.generate_content([
system_prompt,
pil_image,
"Analyze this image and generate an appropriate prompt according to the instructions."
])
response = model.generate_content(
[
system_prompt,
pil_image,
"Analyze this image and generate an appropriate prompt according to the instructions.",
]
)
# Extract text from response
if response.text:
return response.text.strip(), None
else:
return "", "No response generated from Gemini"
except Exception as e:
return "", f"Gemini API error: {str(e)}"
def validate_prompt_type(prompt_type: str) -> bool:
"""Validate if prompt type is supported.
Args:
prompt_type: Type of prompt to validate
Returns:
True if valid, False otherwise
"""
return prompt_type in PROMPT_TEMPLATES
return prompt_type in PROMPT_TEMPLATES
+6 -2
View File
@@ -76,7 +76,11 @@ Install: pip install google-generativeai
# Analyze image with Gemini
prompt, error = analyze_image_with_gemini(
image_np, prompt_type, api_key=api_key or None, custom_prompt=custom_prompt or None, model_name=model
image_np,
prompt_type,
api_key=api_key or None,
custom_prompt=custom_prompt or None,
model_name=model,
)
if error:
@@ -108,4 +112,4 @@ Install: pip install google-generativeai
# Node display name
NODE_DISPLAY_NAME = "Gemini Prompt Engineer"
NODE_DISPLAY_NAME = "Gemini Prompt Engineer"
+2 -2
View File
@@ -194,7 +194,7 @@ GEMINI_MODELS = [
MODEL_DESCRIPTIONS = {
"gemini-1.5-pro": "Most capable Gemini model for complex tasks",
"gemini-1.5-flash": "Faster and cost-effective (recommended for most uses)",
"gemini-1.5-flash-8b": "Smaller and faster, good for simple prompts",
"gemini-1.5-flash-8b": "Smaller and faster, good for simple prompts",
"gemini-pro-vision": "Optimized for vision tasks and image analysis",
"gemini-1.0-pro": "Previous generation, stable option",
}
}
+66 -51
View File
@@ -13,7 +13,11 @@ from kikotools.tools.gemini_prompt.logic import (
validate_prompt_type,
analyze_image_with_gemini,
)
from kikotools.tools.gemini_prompt.prompts import PROMPT_OPTIONS, PROMPT_TEMPLATES, GEMINI_MODELS
from kikotools.tools.gemini_prompt.prompts import (
PROMPT_OPTIONS,
PROMPT_TEMPLATES,
GEMINI_MODELS,
)
class TestGeminiPromptNode:
@@ -29,7 +33,7 @@ class TestGeminiPromptNode:
def test_input_types(self):
"""Test INPUT_TYPES configuration."""
input_types = GeminiPromptNode.INPUT_TYPES()
# Check required inputs
assert "required" in input_types
assert "image" in input_types["required"]
@@ -38,40 +42,40 @@ class TestGeminiPromptNode:
assert input_types["required"]["prompt_type"][0] == PROMPT_OPTIONS
assert "model" in input_types["required"]
assert input_types["required"]["model"][0] == GEMINI_MODELS
# Check optional inputs
assert "optional" in input_types
assert "api_key" in input_types["optional"]
assert "custom_prompt" in input_types["optional"]
def test_gemini_models_available(self):
"""Test that all expected Gemini models are available."""
expected_models = [
"gemini-1.5-pro",
"gemini-1.5-flash",
"gemini-1.5-flash",
"gemini-1.5-flash-8b",
"gemini-pro-vision",
"gemini-1.0-pro"
"gemini-1.0-pro",
]
for model in expected_models:
assert model in GEMINI_MODELS
@patch('kikotools.tools.gemini_prompt.node.analyze_image_with_gemini')
@patch("kikotools.tools.gemini_prompt.node.analyze_image_with_gemini")
def test_generate_prompt_success(self, mock_analyze):
"""Test successful prompt generation."""
# Setup
node = GeminiPromptNode()
test_image = np.random.rand(1, 512, 512, 3).astype(np.float32)
mock_analyze.return_value = ("A beautiful landscape with mountains", None)
# Execute
result = node.generate_prompt(test_image, "flux")
# Assert
assert result == ("A beautiful landscape with mountains", "")
mock_analyze.assert_called_once()
@patch('kikotools.tools.gemini_prompt.node.analyze_image_with_gemini')
@patch("kikotools.tools.gemini_prompt.node.analyze_image_with_gemini")
def test_generate_prompt_sdxl_format(self, mock_analyze):
"""Test SDXL format with positive and negative prompts."""
# Setup
@@ -79,26 +83,29 @@ class TestGeminiPromptNode:
test_image = np.random.rand(1, 512, 512, 3).astype(np.float32)
mock_analyze.return_value = (
"Positive: beautiful landscape, mountains, sunset\nNegative: blurry, low quality",
None
None,
)
# Execute
result = node.generate_prompt(test_image, "sdxl")
# Assert
assert result == ("beautiful landscape, mountains, sunset", "blurry, low quality")
@patch('kikotools.tools.gemini_prompt.node.analyze_image_with_gemini')
# Assert
assert result == (
"beautiful landscape, mountains, sunset",
"blurry, low quality",
)
@patch("kikotools.tools.gemini_prompt.node.analyze_image_with_gemini")
def test_generate_prompt_error(self, mock_analyze):
"""Test error handling in prompt generation."""
# Setup
node = GeminiPromptNode()
test_image = np.random.rand(1, 512, 512, 3).astype(np.float32)
mock_analyze.return_value = ("", "API key not found")
# Execute
result = node.generate_prompt(test_image, "flux")
# Assert
assert result[0].startswith("Error:")
assert result[1] == ""
@@ -107,7 +114,7 @@ class TestGeminiPromptNode:
"""Test handling of invalid prompt type."""
node = GeminiPromptNode()
test_image = np.random.rand(1, 512, 512, 3).astype(np.float32)
with pytest.raises(ValueError, match="Invalid prompt type"):
node.generate_prompt(test_image, "invalid_type")
@@ -123,7 +130,7 @@ class TestGeminiLogic:
assert isinstance(result, Image.Image)
assert result.size == (64, 64)
assert result.mode == "RGB"
# Test 3D tensor
tensor_3d = np.random.rand(64, 64, 3)
result = tensor_to_pil(tensor_3d)
@@ -133,36 +140,40 @@ class TestGeminiLogic:
def test_image_to_base64(self):
"""Test image to base64 conversion."""
# Create test image
image = Image.new('RGB', (64, 64), color='red')
image = Image.new("RGB", (64, 64), color="red")
# Convert to base64
result = image_to_base64(image)
assert isinstance(result, str)
assert len(result) > 0
# Test JPEG format
result_jpeg = image_to_base64(image, format="JPEG")
assert isinstance(result_jpeg, str)
assert result != result_jpeg # Different formats should produce different results
assert (
result != result_jpeg
) # Different formats should produce different results
@patch.dict('os.environ', {'GEMINI_API_KEY': 'test_key_123'})
@patch.dict("os.environ", {"GEMINI_API_KEY": "test_key_123"})
def test_get_api_key_from_env(self):
"""Test getting API key from environment."""
result = get_api_key()
assert result == "test_key_123"
@patch.dict('os.environ', {}, clear=True)
@patch('os.path.exists')
@patch('builtins.open')
@patch.dict("os.environ", {}, clear=True)
@patch("os.path.exists")
@patch("builtins.open")
def test_get_api_key_from_config(self, mock_open, mock_exists):
"""Test getting API key from config file."""
# Setup
mock_exists.return_value = True
mock_open.return_value.__enter__.return_value.read.return_value = '{"api_key": "config_key_456"}'
mock_open.return_value.__enter__.return_value.read.return_value = (
'{"api_key": "config_key_456"}'
)
# Execute
result = get_api_key()
# Assert
assert result == "config_key_456"
@@ -171,14 +182,14 @@ class TestGeminiLogic:
# Valid types
for prompt_type in PROMPT_OPTIONS:
assert validate_prompt_type(prompt_type) is True
# Invalid types
assert validate_prompt_type("invalid") is False
assert validate_prompt_type("") is False
assert validate_prompt_type(None) is False
@patch('google.generativeai.configure')
@patch('google.generativeai.GenerativeModel')
@patch("google.generativeai.configure")
@patch("google.generativeai.GenerativeModel")
def test_analyze_image_with_gemini_success(self, mock_model_class, mock_configure):
"""Test successful image analysis with Gemini."""
# Setup
@@ -187,12 +198,14 @@ class TestGeminiLogic:
mock_response.text = "A beautiful sunset over mountains"
mock_model.generate_content.return_value = mock_response
mock_model_class.return_value = mock_model
test_image = np.random.rand(64, 64, 3)
# Execute
result, error = analyze_image_with_gemini(test_image, "flux", api_key="test_key")
result, error = analyze_image_with_gemini(
test_image, "flux", api_key="test_key"
)
# Assert
assert result == "A beautiful sunset over mountains"
assert error is None
@@ -202,15 +215,17 @@ class TestGeminiLogic:
def test_analyze_image_no_api_key(self):
"""Test analysis without API key."""
test_image = np.random.rand(64, 64, 3)
with patch('kikotools.tools.gemini_prompt.logic.get_api_key', return_value=None):
with patch(
"kikotools.tools.gemini_prompt.logic.get_api_key", return_value=None
):
result, error = analyze_image_with_gemini(test_image, "flux")
assert result == ""
assert "API key not found" in error
@patch('google.generativeai.configure')
@patch('google.generativeai.GenerativeModel')
@patch("google.generativeai.configure")
@patch("google.generativeai.GenerativeModel")
def test_analyze_image_with_custom_prompt(self, mock_model_class, mock_configure):
"""Test analysis with custom prompt."""
# Setup
@@ -219,19 +234,19 @@ class TestGeminiLogic:
mock_response.text = "Custom analysis result"
mock_model.generate_content.return_value = mock_response
mock_model_class.return_value = mock_model
test_image = np.random.rand(64, 64, 3)
custom_prompt = "Analyze this image and describe the colors"
# Execute
result, error = analyze_image_with_gemini(
test_image, "flux", api_key="test_key", custom_prompt=custom_prompt
)
# Assert
assert result == "Custom analysis result"
assert error is None
# Check that custom prompt was used
call_args = mock_model.generate_content.call_args[0][0]
assert custom_prompt in call_args
@@ -251,15 +266,15 @@ class TestPromptTemplates:
"""Test that prompt templates contain expected content."""
# FLUX prompt should mention FLUX
assert "FLUX" in PROMPT_TEMPLATES["flux"]
# SDXL prompt should mention positive and negative
assert "Positive" in PROMPT_TEMPLATES["sdxl"]
assert "Negative" in PROMPT_TEMPLATES["sdxl"]
# Danbooru should mention tags and underscores
assert "tag" in PROMPT_TEMPLATES["danbooru"].lower()
assert "underscore" in PROMPT_TEMPLATES["danbooru"].lower()
# Video should mention motion and temporal
assert "motion" in PROMPT_TEMPLATES["video"].lower()
assert "temporal" in PROMPT_TEMPLATES["video"].lower()
assert "temporal" in PROMPT_TEMPLATES["video"].lower()