diff --git a/kikotools/tools/gemini_prompt/__init__.py b/kikotools/tools/gemini_prompt/__init__.py index 7d3e475..dd2baf8 100644 --- a/kikotools/tools/gemini_prompt/__init__.py +++ b/kikotools/tools/gemini_prompt/__init__.py @@ -2,4 +2,4 @@ from .node import GeminiPromptNode -__all__ = ["GeminiPromptNode"] \ No newline at end of file +__all__ = ["GeminiPromptNode"] diff --git a/kikotools/tools/gemini_prompt/logic.py b/kikotools/tools/gemini_prompt/logic.py index 00469a1..14a74fc 100644 --- a/kikotools/tools/gemini_prompt/logic.py +++ b/kikotools/tools/gemini_prompt/logic.py @@ -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 \ No newline at end of file + return prompt_type in PROMPT_TEMPLATES diff --git a/kikotools/tools/gemini_prompt/node.py b/kikotools/tools/gemini_prompt/node.py index 86174b4..7d88de9 100644 --- a/kikotools/tools/gemini_prompt/node.py +++ b/kikotools/tools/gemini_prompt/node.py @@ -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" \ No newline at end of file +NODE_DISPLAY_NAME = "Gemini Prompt Engineer" diff --git a/kikotools/tools/gemini_prompt/prompts.py b/kikotools/tools/gemini_prompt/prompts.py index 0ffcfb8..f4e919c 100644 --- a/kikotools/tools/gemini_prompt/prompts.py +++ b/kikotools/tools/gemini_prompt/prompts.py @@ -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", -} \ No newline at end of file +} diff --git a/tests/unit/tools/test_gemini_prompt.py b/tests/unit/tools/test_gemini_prompt.py index e12078d..e26af2f 100644 --- a/tests/unit/tools/test_gemini_prompt.py +++ b/tests/unit/tools/test_gemini_prompt.py @@ -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() \ No newline at end of file + assert "temporal" in PROMPT_TEMPLATES["video"].lower()