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:
@@ -2,4 +2,4 @@
|
||||
|
||||
from .node import GeminiPromptNode
|
||||
|
||||
__all__ = ["GeminiPromptNode"]
|
||||
__all__ = ["GeminiPromptNode"]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user