Security updates

This commit is contained in:
Aryan185
2026-04-30 10:47:17 +00:00
parent b7e6d0b311
commit 7ed99966d4
8 changed files with 353 additions and 340 deletions
+19 -14
View File
@@ -1,6 +1,5 @@
import os
import json
import tempfile
import io
import re
import wave
@@ -8,6 +7,7 @@ import torch
import numpy as np
from google import genai
from google.genai import types
from google.oauth2 import service_account
class GeminiDiarisationNode:
@classmethod
@@ -31,7 +31,7 @@ class GeminiDiarisationNode:
"me-central1", "me-central2", "me-west1"
], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
"model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3-pro-preview", "gemini-3.1-pro-preview", "gemini-3-flash-preview", "gemini-3.1-flash-lite-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}),
"model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3.1-pro-preview", "gemini-3.1-flash-lite-preview", "gemini-3-flash-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}),
"seed": ("INT", {"default": 69, "min": 0, "max": 2147483646, "step": 1}),
"temperature": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1})
},
@@ -54,20 +54,20 @@ class GeminiDiarisationNode:
raise ValueError("Project ID is required.")
try:
json.loads(service_account_json)
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
temp_file.write(service_account_json.strip())
temp_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
return genai.Client(
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
@@ -87,7 +87,9 @@ class GeminiDiarisationNode:
return sum(float(x) * 60 ** i for i, x in enumerate(reversed(parts)))
except: return 0.0
def diarise(self, audio, num_speakers, project_id, location, service_account, model, seed, temperature, thinking=False, thinking_budget=0, audio_timestamp=False):
def diarise(self, audio, num_speakers, project_id, location, service_account, model, seed, temperature,
thinking=False, thinking_budget=0, audio_timestamp=False):
waveform = audio.get("waveform")
sr = audio.get("sample_rate")
@@ -139,9 +141,12 @@ class GeminiDiarisationNode:
*You must PASS this benchmark to be deployed*"""
config = types.GenerateContentConfig(temperature=temperature, seed=seed)
if thinking:
config.thinking_config = types.ThinkingConfig(include_thoughts=False, thinking_budget=thinking_budget)
config = types.GenerateContentConfig(
temperature=temperature,
seed=seed,
audio_timestamp=audio_timestamp if audio_timestamp else None,
thinking_config=types.ThinkingConfig(include_thoughts=False, thinking_budget=thinking_budget) if thinking else None
)
response = client.models.generate_content(
model=model,
+44 -44
View File
@@ -1,17 +1,16 @@
import base64
import os
import io
import json
import re
import tempfile
import numpy as np
import torch
from PIL import Image
from google import genai
from google.genai import types
from google.oauth2 import service_account
class GeminiSegmentationVertexNode:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -20,16 +19,16 @@ class GeminiSegmentationVertexNode:
"segment_prompt": ("STRING", {"default": "all objects", "multiline": True}),
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": ([
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"me-central1", "me-central2", "me-west1"
], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
@@ -41,53 +40,59 @@ class GeminiSegmentationVertexNode:
"gemini-2.0-flash"
], {"default": "gemini-2.5-flash"}),
"temperature": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 2.0, "step": 0.1}),
"thinking": ("BOOLEAN", {"default": True}),
"thinking": ("BOOLEAN", {"default": False}),
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
},
"optional": {
"thinking_budget": ("INT", {"default": 0, "min": -1, "max": 24576, "step": 1}),
}
}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("mask",)
FUNCTION = "generate_segmentation"
CATEGORY = "image/generation"
def setup_client(self, service_account_json, project_id, location):
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
if not project_id.strip():
raise ValueError("Project ID is required.")
try:
json.loads(service_account_json)
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
temp_file.write(service_account_json.strip())
temp_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
return genai.Client(
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
)
def generate_segmentation(self, image, segment_prompt, project_id, location, service_account, model, temperature, thinking, seed, thinking_budget=0):
client = self.setup_client(service_account, project_id, location)
# Image Preprocessing
img_np = (image[0].cpu().numpy() * 255).astype(np.uint8)
orig_img = Image.fromarray(img_np)
orig_w, orig_h = orig_img.size
scale = min(1024 / orig_w, 1024 / orig_h)
proc_img = orig_img.resize((int(orig_w * scale), int(orig_h * scale)), Image.Resampling.LANCZOS) if scale < 1 else orig_img
pw, ph = proc_img.size
img_buf = io.BytesIO()
proc_img.save(img_buf, format='PNG')
# Thinking Config
t_config = None
if "gemini-2.0" not in model.lower():
budget = thinking_budget if thinking else 0
@@ -97,8 +102,7 @@ class GeminiSegmentationVertexNode:
t_config = types.ThinkingConfig(thinking_budget=budget)
prompt = f'Give the segmentation masks for {segment_prompt}. Output a JSON list of segmentation masks where each entry contains the 2D bounding box in the key "box_2d", the segmentation mask in key "mask", and the text label in the key "label".'
# API Call
response = client.models.generate_content(
model=model,
contents=[types.Content(role="user", parts=[
@@ -116,28 +120,24 @@ class GeminiSegmentationVertexNode:
raise RuntimeError(f"Gemini API Error: {e}")
final_mask = np.zeros((ph, pw), dtype=np.uint8)
for seg in segments:
try:
# Calculate integer coords
ymin, xmin, ymax, xmax = seg['box_2d']
x1, y1 = int(xmin * pw / 1000), int(ymin * ph / 1000)
x2, y2 = int(xmax * pw / 1000), int(ymax * ph / 1000)
w, h = x2 - x1, y2 - y1
if w <= 0 or h <= 0: continue
# Decode & Resize Patch
mask_str = seg['mask'].split(",")[1] if "data:image" in seg['mask'] else seg['mask']
patch = Image.open(io.BytesIO(base64.b64decode(mask_str))).convert('L')
if patch.size != (w, h):
patch = patch.resize((w, h), Image.Resampling.NEAREST)
patch_arr = np.array(patch)
patch_arr = np.where(patch_arr > 128, 255, 0).astype(np.uint8)
# Safe slicing to handle potential boundary issues
patch_arr = np.where(np.array(patch) > 128, 255, 0).astype(np.uint8)
target_slice = final_mask[y1:y2, x1:x2]
if target_slice.shape == patch_arr.shape:
np.maximum(target_slice, patch_arr, out=target_slice)
@@ -146,7 +146,7 @@ class GeminiSegmentationVertexNode:
if (pw, ph) != (orig_w, orig_h):
final_mask = np.array(Image.fromarray(final_mask).resize((orig_w, orig_h), Image.Resampling.NEAREST))
return (torch.from_numpy(final_mask.astype(np.float32) / 255.0).unsqueeze(0),)
NODE_CLASS_MAPPINGS = {"GeminiSegmentationVertexNode": GeminiSegmentationVertexNode}
+43 -51
View File
@@ -1,12 +1,10 @@
import os
import io
import json
import tempfile
import torch
from google.genai import Client, types
from google.oauth2 import service_account
class GeminiTTSVertexNode:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -14,21 +12,21 @@ class GeminiTTSVertexNode:
"text": ("STRING", {"multiline": True, "default": ""}),
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": ([
"global","us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"me-central1", "me-central2", "me-west1"
], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
"model": (["gemini-2.5-flash-preview-tts", "gemini-2.5-pro-preview-tts"],),
"voice_id": (["Zephyr", "Puck", "Charon", "Kore", "Fenrir", "Leda", "Orus", "Aoede", "Callirrhoe", "Autonoe", "Enceladus", "Iapetus", "Umbriel", "Algieba", "Despina", "Erinome", "Achernar", "Laomedeia", "Rasalgethi", "Algenib", "Achird", "Pulcherrima", "Gacrux", "Schedar", "Alnilam", "Sulafat", "Sadaltager", "Sadachbia", "Vindemiatrix", "Zubenelgenubi"],),
"voice_id": (["Zephyr", "Puck", "Charon", "Kore", "Fenrir", "Leda", "Orus", "Aoede", "Callirrhoe", "Autonoe", "Enceladus", "Iapetus", "Umbriel", "Algieba", "Despina", "Erinome", "Achernar", "Laomedeia", "Rasalgethi", "Algenib", "Achird", "Pulcherrima", "Gacrux", "Schedar", "Alnilam", "Sulafat", "Sadaltager", "Sadachbia", "Vindemiatrix", "Zubenelgenubi"],),
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
"temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
@@ -36,64 +34,59 @@ class GeminiTTSVertexNode:
"system_prompt": ("STRING", {"multiline": True, "default": ""}),
}
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "generate_speech"
CATEGORY = "audio/generation"
def setup_client(self, service_account_json, project_id, location):
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
if not project_id.strip():
raise ValueError("Project ID is required.")
try:
json.loads(service_account_json)
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
temp_file.write(service_account_json.strip())
temp_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
return Client(
vertexai=True,
project=project_id.strip(),
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
)
def generate_speech(self, text, project_id, location, service_account, voice_id,
temperature, model, seed, system_prompt=""):
def generate_speech(self, text, project_id, location, service_account, voice_id,
temperature, model, seed, system_prompt=""):
if not text.strip():
raise ValueError("Text input cannot be empty.")
client = self.setup_client(service_account, project_id, location)
final_prompt = text
if system_prompt.strip():
final_prompt = f"{system_prompt.strip()}\n\n{text}"
speech_config = types.SpeechConfig(
voice_config=types.VoiceConfig(
prebuilt_voice_config=types.PrebuiltVoiceConfig(voice_name=voice_id)
)
)
client = self.setup_client(service_account, project_id, location)
final_prompt = f"{system_prompt.strip()}\n\n{text}" if system_prompt.strip() else text
config = types.GenerateContentConfig(
temperature=temperature,
seed=seed,
response_modalities=["AUDIO"],
speech_config=speech_config,
speech_config=types.SpeechConfig(
voice_config=types.VoiceConfig(
prebuilt_voice_config=types.PrebuiltVoiceConfig(voice_name=voice_id)
)
)
)
try:
response = client.models.generate_content(
model=model,
@@ -104,17 +97,16 @@ class GeminiTTSVertexNode:
raise RuntimeError(f"Gemini API Error: {str(e)}")
try:
inline_data = response.candidates[0].content.parts[0].inline_data
audio_bytes = inline_data.data
audio_bytes = response.candidates[0].content.parts[0].inline_data.data
except (AttributeError, IndexError, TypeError):
raise ValueError("API returned a response, but it contained no audio data.")
waveform = torch.frombuffer(bytearray(audio_bytes), dtype=torch.int16)
waveform = waveform.to(torch.float32) / 32768.0
waveform = waveform.unsqueeze(0).unsqueeze(0)
return ({"waveform": waveform, "sample_rate": 24000},)
@classmethod
def IS_CHANGED(cls, **kwargs):
return f"{kwargs.get('text', '')}-{kwargs.get('voice_id', '')}-{kwargs.get('temperature', 1.0)}-{kwargs.get('model', '')}-{kwargs.get('seed', 69)}-{kwargs.get('system_prompt', '')}"
+55 -50
View File
@@ -1,16 +1,15 @@
import os
import json
import tempfile
import io
import json
import base64
import torch
import numpy as np
from PIL import Image
from google import genai
from google.genai import types
from google.oauth2 import service_account
class GoogleImagenEditVertex:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -32,61 +31,68 @@ class GoogleImagenEditVertex:
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("edited_images",)
FUNCTION = "edit_image"
CATEGORY = "image/edit"
def setup_client(self, service_account_json, project_id, location):
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
if not project_id.strip():
raise ValueError("Project ID is required.")
try:
json.loads(service_account_json)
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
temp_file.write(service_account_json.strip())
temp_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
def edit_image(self, image, mask, prompt, project_id, location, service_account,
edit_mode, number_of_images, seed, base_steps, guidance_scale, mask_dilation, negative_prompt=""):
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
return genai.Client(
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
)
def edit_image(self, image, mask, prompt, project_id, location, service_account,
edit_mode, number_of_images, seed, base_steps, guidance_scale, mask_dilation, negative_prompt=""):
client = self.setup_client(service_account, project_id, location)
def to_b64(img):
b = io.BytesIO()
img.save(b, format='PNG')
return base64.b64encode(b.getvalue()).decode('utf-8')
img_pil = Image.fromarray((image[0].cpu().numpy() * 255).astype(np.uint8))
mask_np = mask.cpu().numpy()
if mask_np.ndim == 3: mask_np = mask_np[0]
mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8), mode='L')
config_dict = {
"edit_mode": edit_mode,
"number_of_images": number_of_images,
"base_steps": base_steps,
"seed": seed,
"guidance_scale": guidance_scale,
"output_mime_type": "image/jpeg",
"include_rai_reason": True,
}
if negative_prompt.strip():
config_dict["negative_prompt"] = negative_prompt.strip()
try:
def to_b64(img):
b = io.BytesIO()
img.save(b, format='PNG')
return base64.b64encode(b.getvalue()).decode('utf-8')
img_pil = Image.fromarray((image[0].cpu().numpy() * 255).astype(np.uint8))
mask_np = mask.cpu().numpy()
if mask_np.ndim == 3: mask_np = mask_np[0]
mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8), mode='L')
config_dict = {
"edit_mode": edit_mode,
"number_of_images": number_of_images,
"base_steps": base_steps,
"seed": seed,
"guidance_scale": guidance_scale,
"output_mime_type": "image/jpeg",
"include_rai_reason": True,
}
if negative_prompt.strip():
config_dict["negative_prompt"] = negative_prompt.strip()
response = client.models.edit_image(
model="imagen-3.0-capability-001",
prompt=prompt,
@@ -98,18 +104,17 @@ class GoogleImagenEditVertex:
config=types.EditImageConfig(**config_dict)
)
if not response.generated_images: raise ValueError("No images generated")
if not response.generated_images:
raise ValueError("No images generated")
output_tensors = []
for item in response.generated_images:
img_bytes = item.image.image_bytes
res_img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
res_img = Image.open(io.BytesIO(item.image.image_bytes)).convert("RGB")
output_tensors.append(torch.from_numpy(np.array(res_img).astype(np.float32) / 255.0))
return (torch.stack(output_tensors),)
except Exception as e:
print(f"Google Imagen Edit Error: {e}")
raise RuntimeError(f"Google Imagen Edit Error: {e}")
@classmethod
+44 -41
View File
@@ -1,63 +1,68 @@
import os
import json
import tempfile
import io
import json
import torch
import numpy as np
from PIL import Image
from google import genai
from google.genai import types
from google.oauth2 import service_account
class GoogleImagenGenerateVertex:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"multiline": True, "default": "A majestic lion in the savanna"}),
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": (["global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1", "us-west1", "us-west2", "us-west3", "us-west4", "northamerica-northeast1", "northamerica-northeast2", "southamerica-east1", "southamerica-west1", "africa-south1", "europe-west1", "europe-north1", "europe-west2", "europe-west3", "europe-west4", "europe-west6", "europe-west8", "europe-west9", "europe-west12", "europe-southwest1", "europe-central2", "asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2", "asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1", "asia-southeast2", "australia-southeast1", "australia-southeast2", "me-central1", "me-central2", "me-west1"], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
"model": (["imagen-4.0-ultra-generate-001", "imagen-4.0-generate-001", "imagen-4.0-fast-generate-001", "imagen-3.0-generate-002"], {"default": "imagen-4.0-generate-001"}),
"number_of_images": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
"aspect_ratio": (["1:1", "9:16", "16:9", "4:3", "3:4"], {"default": "1:1"}),
"image_size": (["1K", "2K"], {"default": "1K"}),
"seed": ("INT", {"default": 69, "min": 1, "max": 2147483646, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 20.0, "step": 0.1}),
"prompt": ("STRING", {"multiline": True, "default": "A majestic lion in the savanna"}),
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": (["global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1", "us-west1", "us-west2", "us-west3", "us-west4", "northamerica-northeast1", "northamerica-northeast2", "southamerica-east1", "southamerica-west1", "africa-south1", "europe-west1", "europe-north1", "europe-west2", "europe-west3", "europe-west4", "europe-west6", "europe-west8", "europe-west9", "europe-west12", "europe-southwest1", "europe-central2", "asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2", "asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1", "asia-southeast2", "australia-southeast1", "australia-southeast2", "me-central1", "me-central2", "me-west1"], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
"model": (["imagen-4.0-ultra-generate-001", "imagen-4.0-generate-001", "imagen-4.0-fast-generate-001", "imagen-3.0-generate-002"], {"default": "imagen-4.0-generate-001"}),
"number_of_images": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
"aspect_ratio": (["1:1", "9:16", "16:9", "4:3", "3:4"], {"default": "1:1"}),
"image_size": (["1K", "2K"], {"default": "1K"}),
"seed": ("INT", {"default": 69, "min": 1, "max": 2147483646, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 20.0, "step": 0.1}),
},
"optional": {
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "generate_images"
CATEGORY = "image/generation"
def setup_client(self, service_account_json, project_id, location):
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
if not project_id.strip():
raise ValueError("Project ID is required.")
try:
json.loads(service_account_json)
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
temp_file.write(service_account_json.strip())
temp_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
return genai.Client(
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
)
def generate_images(self, prompt, project_id, location, service_account, model, number_of_images, aspect_ratio, image_size, seed, guidance_scale, negative_prompt=""):
client = self.setup_client(service_account, project_id, location)
config = types.GenerateImagesConfig(
number_of_images=number_of_images,
aspect_ratio=aspect_ratio,
@@ -65,34 +70,32 @@ class GoogleImagenGenerateVertex:
seed=seed,
negative_prompt=negative_prompt.strip() if negative_prompt.strip() else None
)
if "imagen-4.0" in model and "fast" not in model:
config.image_size = image_size
try:
result = client.models.generate_images(model=model, prompt=prompt, config=config)
if not result.generated_images: raise ValueError("No images generated")
if not result.generated_images:
raise ValueError("No images generated")
tensors = []
for item in result.generated_images:
img_data = item.image
if hasattr(img_data, "image_bytes"):
pil_img = Image.open(io.BytesIO(img_data.image_bytes))
elif hasattr(img_data, "convert"):
pil_img = img_data
else:
pil_img = Image.open(io.BytesIO(img_data))
pil_img = pil_img.convert("RGB")
tensors.append(torch.from_numpy(np.array(pil_img).astype(np.float32) / 255.0))
tensors.append(torch.from_numpy(np.array(pil_img.convert("RGB")).astype(np.float32) / 255.0))
return (torch.stack(tensors),)
except Exception as e:
print(f"Google Imagen Error: {e}")
raise RuntimeError(f"Google Imagen Error: {e}")
@classmethod
def IS_CHANGED(cls, **kwargs):
return float("nan")
+57 -45
View File
@@ -1,39 +1,39 @@
import os
import io
import json
import tempfile
import torch
import numpy as np
from PIL import Image
from google import genai
from google.genai import types
from google.oauth2 import service_account
class NanoBananaVertexNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": ([
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"me-central1", "me-central2", "me-west1"
], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
"model": (["gemini-3-pro-image-preview", "gemini-2.5-flash-image", "gemini-3.1-flash-image-preview"],),
"aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "9:16", "16:9", "21:9"],),
"resolution": (["1K", "2K", "4K"], {"default": "1K"}),
"temperature": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"top_p": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1.0, "step": 0.01}),
"temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"top_p": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.01}),
"google_search": ("BOOLEAN", {"default": False}),
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
},
"optional": {
@@ -46,71 +46,85 @@ class NanoBananaVertexNode:
"image_5": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "generate"
CATEGORY = "image/generation"
def setup_client(self, service_account_json, project_id, location):
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
if not project_id.strip():
raise ValueError("Project ID is required.")
try:
json.loads(service_account_json)
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
temp_file.write(service_account_json.strip())
temp_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
return genai.Client(
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
)
def _convert_tensor_to_bytes(self, tensor):
if tensor.dim() == 4:
tensor = tensor[0]
arr = (tensor.cpu().numpy() * 255).astype(np.uint8)
buf = io.BytesIO()
Image.fromarray(arr).save(buf, format='PNG')
return buf.getvalue()
def generate(self, project_id, location, service_account, model, aspect_ratio, resolution, temperature, top_p, seed,
def generate(self, project_id, location, service_account, model, aspect_ratio, resolution,
temperature, top_p, google_search, seed,
prompt="", system_instruction="", **kwargs):
client = self.setup_client(service_account, project_id, location)
parts = []
input_images = [kwargs.get(f"image_{i}") for i in range(1, 6)]
for img in input_images:
for i in range(1, 6):
img = kwargs.get(f"image_{i}")
if img is not None:
img_bytes = self._convert_tensor_to_bytes(img)
parts.append(types.Part.from_bytes(mime_type="image/png", data=img_bytes))
parts.append(types.Part.from_bytes(mime_type="image/png", data=self._convert_tensor_to_bytes(img)))
if prompt.strip():
parts.append(types.Part.from_text(text=prompt))
if not parts:
raise ValueError("At least one image or prompt must be provided.")
tools = None
if google_search:
if "gemini-2.5" in model:
print(f"Ignoring google_search: {model} does not support it.")
else:
tools = [types.Tool(googleSearch=types.GoogleSearch())]
img_config_params = {"aspect_ratio": aspect_ratio}
if "gemini-3-pro" in model:
img_config_params["image_size"] = resolution
config = types.GenerateContentConfig(
temperature=temperature,
seed=seed,
top_p=top_p,
response_modalities=["IMAGE"],
image_config=types.ImageConfig(**img_config_params),
system_instruction=system_instruction.strip() if system_instruction.strip() else None
system_instruction=system_instruction.strip() if system_instruction.strip() else None,
tools=tools
)
try:
response = client.models.generate_content(
model=model,
@@ -119,14 +133,11 @@ class NanoBananaVertexNode:
)
except Exception as e:
raise RuntimeError(f"Gemini API Error: {str(e)}")
try:
img_data = response.candidates[0].content.parts[0].inline_data.data
result_pil = Image.open(io.BytesIO(img_data)).convert("RGB")
result_tensor = torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0)
return (result_tensor,)
return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),)
except (AttributeError, IndexError, TypeError):
raise ValueError("API returned a response, but no valid image data was found.")
@@ -134,5 +145,6 @@ class NanoBananaVertexNode:
def IS_CHANGED(cls, seed, **kwargs):
return seed
NODE_CLASS_MAPPINGS = {"NanoBananaVertexNode": NanoBananaVertexNode}
NODE_DISPLAY_NAME_MAPPINGS = {"NanoBananaVertexNode": "Nano Banana (Vertex AI)"}
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "vertexapi"
description = "A collection of powerful custom nodes for ComfyUI that connect your local workflows to closed-source AI models via Vertex AI."
version = "1.0.3"
version = "1.1.0"
license = {file = "LICENSE"}
# classifiers = [
# # For OS-independent nodes (works on all operating systems)
+90 -94
View File
@@ -1,16 +1,15 @@
import time
import os
import io
import json
import tempfile
import torch
import numpy as np
import av
from PIL import Image
from google import genai
from google.genai import types
from google.oauth2 import service_account
class GoogleVeoVertexVideoGenerator:
class GoogleVeoVertexVideoGenerator:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -18,16 +17,16 @@ class GoogleVeoVertexVideoGenerator:
"prompt": ("STRING", {"multiline": True, "default": "a cat reading a book"}),
"project_id": ("STRING", {"multiline": False, "default": ""}),
"location": ([
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
"us-west1", "us-west2", "us-west3", "us-west4",
"northamerica-northeast1", "northamerica-northeast2",
"southamerica-east1", "southamerica-west1", "africa-south1",
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
"europe-west12", "europe-southwest1", "europe-central2",
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
"asia-southeast2", "australia-southeast1", "australia-southeast2",
"me-central1", "me-central2", "me-west1"
], {"default": "us-central1"}),
"service_account": ("STRING", {"multiline": True, "default": ""}),
@@ -38,7 +37,8 @@ class GoogleVeoVertexVideoGenerator:
"veo-3.0-generate-001",
"veo-3.0-fast-generate-001",
"veo-3.1-generate-001",
"veo-3.1-fast-generate-001"
"veo-3.1-fast-generate-001",
"veo-3.1-lite-generate-preview"
], {"default": "veo-3.0-generate-001"}),
"resolution": (["720p", "1080p"], {"default": "720p"}),
"aspect_ratio": (["16:9", "9:16"], {"default": "16:9"}),
@@ -53,109 +53,105 @@ class GoogleVeoVertexVideoGenerator:
"last_frame": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE", "AUDIO")
RETURN_NAMES = ("frames", "audio")
FUNCTION = "generate_video"
CATEGORY = "video/generation"
OUTPUT_IS_LIST = (True, False)
def generate_video(self, prompt, project_id, location, service_account, model, resolution, aspect_ratio,
duration_seconds, seed, generate_audio, fps, negative_prompt=None,
first_frame=None, last_frame=None):
# Validate Service Account JSON
if not service_account.strip():
def setup_client(self, service_account_json, project_id, location):
if not service_account_json.strip():
raise ValueError("Service account JSON content is required.")
if not project_id.strip():
raise ValueError("Project ID is required.")
try:
json.loads(service_account)
sa_info = json.loads(service_account_json)
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON content: {str(e)}")
creds_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
creds_file.write(service_account.strip())
creds_file.close()
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = creds_file.name
credentials = service_account.Credentials.from_service_account_info(
sa_info,
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
try:
client = genai.Client(vertexai=True, project=project_id, location=location)
config = types.GenerateVideosConfig(
resolution=resolution,
aspect_ratio=aspect_ratio,
duration_seconds=duration_seconds,
generate_audio=generate_audio,
fps=int(fps),
seed=seed if seed != -1 else None,
negative_prompt=negative_prompt.strip() if negative_prompt and negative_prompt.strip() else None
return genai.Client(
vertexai=True,
project=project_id.strip(),
location=location.strip(),
credentials=credentials,
http_options=types.HttpOptions(
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
)
)
# Helper: Tensor -> Bytes
def tensor_to_bytes(t):
# Handle batch dimension if present
if t.dim() == 4:
t = t[0]
arr = (t.cpu().numpy() * 255).astype(np.uint8)
b = io.BytesIO()
Image.fromarray(arr).save(b, format="PNG")
return b.getvalue()
def generate_video(self, prompt, project_id, location, service_account, model, resolution, aspect_ratio,
duration_seconds, seed, generate_audio, fps, negative_prompt=None,
first_frame=None, last_frame=None):
gen_kwargs = {"model": model, "prompt": prompt, "config": config}
if first_frame is not None:
gen_kwargs["image"] = types.Image(image_bytes=tensor_to_bytes(first_frame), mime_type="image/png")
if last_frame is not None:
setattr(config, 'last_frame', types.Image(image_bytes=tensor_to_bytes(last_frame), mime_type="image/png"))
client = self.setup_client(service_account, project_id, location)
op = client.models.generate_videos(**gen_kwargs)
print(f"Veo Operation: {op.name}")
while not op.done:
time.sleep(5)
op = client.operations.get(op)
if op.error: raise Exception(f"Veo Error: {op.error}")
if not op.result.generated_videos: raise Exception("No videos generated")
def tensor_to_bytes(t):
if t.dim() == 4: t = t[0]
arr = (t.cpu().numpy() * 255).astype(np.uint8)
b = io.BytesIO()
Image.fromarray(arr).save(b, format="PNG")
return b.getvalue()
video_bytes = io.BytesIO(op.result.generated_videos[0].video.video_bytes)
# Decode Video
config = types.GenerateVideosConfig(
resolution=resolution,
aspect_ratio=aspect_ratio,
duration_seconds=duration_seconds,
generate_audio=generate_audio,
fps=int(fps),
seed=seed if seed != -1 else None,
negative_prompt=negative_prompt.strip() if negative_prompt and negative_prompt.strip() else None
)
gen_kwargs = {"model": model, "prompt": prompt, "config": config}
if first_frame is not None:
gen_kwargs["image"] = types.Image(image_bytes=tensor_to_bytes(first_frame), mime_type="image/png")
if last_frame is not None:
setattr(config, 'last_frame', types.Image(image_bytes=tensor_to_bytes(last_frame), mime_type="image/png"))
op = client.models.generate_videos(**gen_kwargs)
print(f"Veo Operation: {op.name}")
while not op.done:
time.sleep(5)
op = client.operations.get(op)
if op.error: raise Exception(f"Veo Error: {op.error}")
if not op.result.generated_videos: raise Exception("No videos generated")
video_bytes = io.BytesIO(op.result.generated_videos[0].video.video_bytes)
container = av.open(video_bytes)
frames = []
for frame in container.decode(video=0):
img = frame.to_rgb().to_ndarray().astype(np.float32) / 255.0
frames.append(torch.from_numpy(img).unsqueeze(0))
container.close()
audio = None
if generate_audio:
video_bytes.seek(0)
container = av.open(video_bytes)
frames = []
for frame in container.decode(video=0):
img = frame.to_rgb().to_ndarray().astype(np.float32) / 255.0
frames.append(torch.from_numpy(img).unsqueeze(0))
if container.streams.audio:
audio_data = [f.to_ndarray() for f in container.decode(audio=0)]
if audio_data:
waveform = torch.from_numpy(np.concatenate(audio_data, axis=1)).float()
if audio_data[0].dtype == np.int16: waveform /= 32768.0
elif audio_data[0].dtype == np.int32: waveform /= 2147483648.0
audio = {"waveform": waveform.unsqueeze(0), "sample_rate": container.streams.audio[0].rate}
container.close()
# Decode Audio
audio = None
if generate_audio:
video_bytes.seek(0)
container = av.open(video_bytes)
if container.streams.audio:
audio_data = [f.to_ndarray() for f in container.decode(audio=0)]
if audio_data:
waveform = torch.from_numpy(np.concatenate(audio_data, axis=1)).float()
# Normalize 16/32-bit audio
if audio_data[0].dtype == np.int16: waveform /= 32768.0
elif audio_data[0].dtype == np.int32: waveform /= 2147483648.0
audio = {
"waveform": waveform.unsqueeze(0),
"sample_rate": container.streams.audio[0].rate
}
container.close()
if not frames: raise Exception("Failed to decode video frames")
if not frames: raise Exception("Failed to decode video frames")
return ([torch.cat(frames, dim=0)], audio)
return ([torch.cat(frames, dim=0)], audio)
finally:
if os.path.exists(creds_file.name):
os.remove(creds_file.name)
NODE_CLASS_MAPPINGS = {"GoogleVeoVertexVideoGenerator": GoogleVeoVertexVideoGenerator}
NODE_DISPLAY_NAME_MAPPINGS = {"GoogleVeoVertexVideoGenerator": "Google Veo (Vertex AI)"}