Gemini-Image-service / handler.py
Tim-canova
Add custom inference handler for Gemini image editing
be413cd
Raw History Blame Contribute Delete
6.42 kB
import json
import os
import time
import uuid
import tempfile
import io
from PIL import Image, ImageDraw, ImageFont
import base64
import mimetypes
from google import genai
from google.genai import types
from typing import Dict, Any
def save_binary_file(file_name, data):
with open(file_name, "wb") as f:
f.write(data)
def generate(text, file_name, api_key, model="gemini-2.0-flash-exp"):
# Initialize client using provided api_key (or fallback to env variable)
client = genai.Client(api_key=(api_key.strip() if api_key and api_key.strip() != ""
else os.environ.get("GEMINI_API_KEY")))
files = [ client.files.upload(file=file_name) ]
contents = [
types.Content(
role="user",
parts=[
types.Part.from_uri(
file_uri=files[0].uri,
mime_type=files[0].mime_type,
),
types.Part.from_text(text=text),
],
),
]
generate_content_config = types.GenerateContentConfig(
temperature=1,
top_p=0.95,
top_k=40,
max_output_tokens=8192,
response_modalities=["image", "text"],
response_mime_type="text/plain",
)
text_response = ""
image_path = None
# Create a temporary file to potentially store image data.
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
temp_path = tmp.name
for chunk in client.models.generate_content_stream(
model=model,
contents=contents,
config=generate_content_config,
):
if not chunk.candidates or not chunk.candidates[0].content or not chunk.candidates[0].content.parts:
continue
candidate = chunk.candidates[0].content.parts[0]
# Check for inline image data
if candidate.inline_data:
save_binary_file(temp_path, candidate.inline_data.data)
print(f"File of mime type {candidate.inline_data.mime_type} saved to: {temp_path} and prompt input: {text}")
image_path = temp_path
# If an image is found, we assume that is the desired output.
break
else:
# Accumulate text response if no inline_data is present.
text_response += chunk.text + "\n"
del files
return image_path, text_response
class EndpointHandler:
def __init__(self, path=""):
"""
Initialize the handler
Args:
path (str): Path to the model directory
"""
# Nothing to initialize for this custom handler
pass
def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
"""
Process the request
Args:
data (Dict): The request payload containing:
- inputs: Dict with 'image', 'prompt', and optionally 'gemini_api_key'
Returns:
Dict: Response with processed results
"""
try:
inputs = data.get("inputs", {})
# Extract parameters
image_data = inputs.get("image")
prompt = inputs.get("prompt", "")
gemini_api_key = inputs.get("gemini_api_key", "")
if not image_data or not prompt:
return {"error": "Missing required inputs: 'image' and 'prompt'"}
# Process image - convert from base64 or file path to PIL Image
if isinstance(image_data, str):
# If base64 string
if image_data.startswith('data:image'):
image_data = image_data.split(',')[1]
image_bytes = base64.b64decode(image_data)
composite_pil = Image.open(io.BytesIO(image_bytes))
elif isinstance(image_data, dict) and 'path' in image_data:
# If file path provided
composite_pil = Image.open(image_data['path'])
else:
return {"error": "Invalid image format"}
# Save the composite image to a temporary file
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
composite_path = tmp.name
composite_pil.save(composite_path)
file_name = composite_path
input_text = prompt
model = "gemini-2.0-flash-exp"
image_path, text_response = generate(text=input_text, file_name=file_name, api_key=gemini_api_key, model=model)
if image_path:
# Load and convert the image if needed
result_img = Image.open(image_path)
if result_img.mode == "RGBA":
result_img = result_img.convert("RGB")
# Convert to base64 for response
output_buffer = io.BytesIO()
result_img.save(output_buffer, format='PNG')
img_base64 = base64.b64encode(output_buffer.getvalue()).decode()
# Clean up temp files
os.unlink(composite_path)
os.unlink(image_path)
return {
"generated_outputs": [{
"image": {
"path": None,
"url": f"data:image/png;base64,{img_base64}",
"size": len(output_buffer.getvalue()),
"orig_name": "edited_image.png",
"mime_type": "image/png",
"is_stream": False,
"meta": {}
},
"caption": None
}],
"gemini_output": "",
"prompt_used": prompt
}
else:
# Clean up temp file
os.unlink(composite_path)
# Return text response if no image generated
return {
"generated_outputs": [],
"gemini_output": text_response,
"prompt_used": prompt
}
except Exception as e:
return {"error": f"Processing failed: {str(e)}"}