import os import io import time import uuid import base64 import logging import asyncio import traceback from typing import Optional, Dict, Any, Union import requests from PIL import Image import gradio as gr import aiohttp from aiohttp import ClientTimeout # Configure logging logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') logger = logging.getLogger(__name__) if not os.environ.get("SPACE_ID"): try: from dotenv import load_dotenv load_dotenv() except Exception as e: logger.error(f"Failed to load environment variables: {str(e)}") # API configuration API_CONFIG = { # Get API endpoints from environment variables, use default values if not exist "SUBMIT_API_URL": os.environ.get("FACESWAP_SUBMIT_API", ""), "RESULT_API_URL": os.environ.get("FACESWAP_RESULT_API", ""), "SUBMIT_API_KEY_ID": os.environ.get("FACESWAP_API_KEY_ID", ""), "SUBMIT_API_KEY_SECRET": os.environ.get("FACESWAP_API_KEY_SECRET", ""), "RESULT_API_KEY_ID": os.environ.get("RESULT_API_KEY_ID", ""), "RESULT_API_KEY_SECRET": os.environ.get("RESULT_API_KEY_SECRET", ""), "TIMEOUT": int(os.environ.get("FACESWAP_TIMEOUT", "26")), # Polling timeout (seconds) "POLL_INTERVAL": int(os.environ.get("FACESWAP_POLL_INTERVAL", "2")), # Polling interval (seconds) } def image_to_base64(img: Union[str, Image.Image]) -> str: """ Convert an image to a base64 encoded string with MIME type prefix Parameters: img: Image file path or PIL.Image object Returns: Base64 encoded string with MIME type prefix """ if isinstance(img, str): # If input is a file path img_extension = os.path.splitext(img)[-1].lower() mime_type_map = { '.jpg': 'image/jpeg', '.jpeg': 'image/jpeg', '.png': 'image/png', '.gif': 'image/gif', '.bmp': 'image/bmp', '.webp': 'image/webp' } mime_type = mime_type_map.get(img_extension, 'image/jpeg') # Default to jpeg # Open image file and convert to base64 with open(img, 'rb') as img_file: img_data = base64.b64encode(img_file.read()).decode('utf-8') else: # If input is a PIL.Image object buffered = io.BytesIO() img_format = img.format if img.format else 'PNG' mime_type = f'image/{img_format.lower()}' if img_format else 'image/png' img.save(buffered, format=img_format if img_format else 'PNG') img_data = base64.b64encode(buffered.getvalue()).decode('utf-8') # Add MIME type prefix return f"data:{mime_type};base64,{img_data}" def submit_face_swap_task(source_base64: str, target_base64: str, input_task_id: str) -> Optional[str]: """Submit face swap task to async API, return task ID""" try: api_url = API_CONFIG["SUBMIT_API_URL"] payload = { "sourceUrl": source_base64, "targetUrl": target_base64, "businessTaskId": input_task_id, "imageDetect": True, "imageBase64Format": True } # Add API keys to request headers if available headers = { 'X-Request-Req-Accesskeyid': API_CONFIG["SUBMIT_API_KEY_ID"], 'X-Request-Req-Accesskeysecret': API_CONFIG["SUBMIT_API_KEY_SECRET"], 'Content-Type': 'application/json' } response = requests.post(api_url, json=payload, headers=headers, timeout=10) response.raise_for_status() result = response.json() # Extract task ID based on actual API response format task_id = result.get("data", {}).get("task_id") logger.info(f"Successfully submitted face swap task, task ID: {task_id}") return task_id except Exception as e: logger.exception(f"Failed to submit face swap task: {str(e)}") return None def poll_face_swap_result(task_id: str, timeout: int = 30) -> Optional[Dict[str, Any]]: """Poll for face swap task result until success or timeout""" start_time = time.time() poll_interval = 2 # Polling interval in seconds api_url = API_CONFIG["RESULT_API_URL"] payload = { "saasTaskId": task_id } # Add API keys to request headers if available headers = { 'X-Request-Req-Accesskeyid': API_CONFIG["RESULT_API_KEY_ID"], 'X-Request-Req-Accesskeysecret': API_CONFIG["RESULT_API_KEY_SECRET"], 'Content-Type': 'application/json' } while time.time() - start_time < timeout: try: response = requests.post(api_url, json=payload, headers=headers, timeout=2) response.raise_for_status() result = response.json() # Check if task is completed if result: if result.get("code", None) == "200": data = result.get("data", None) if data is None: time.sleep(poll_interval) else: return data else: logger.error(f"Task failed") return None else: # Task not completed, wait and continue polling time.sleep(poll_interval) except Exception as e: logger.warning(f"Error during polling: {str(e)}") time.sleep(poll_interval) logger.error(f"Polling timeout, task ID: {task_id}") return None def download_image(url: str) -> Optional[Image.Image]: """Download image from URL""" try: response = requests.get(url, timeout=10) response.raise_for_status() image_data = io.BytesIO(response.content) image = Image.open(image_data) logger.info(f"Successfully downloaded result image, URL: {url}") return image except Exception as e: logger.exception(f"Failed to download image: {str(e)}") return None async def download_image_async(encoding: str, logger: logging.Logger, max_retries: int=2, total: int=7, connect: float=1.5, sock_read: float=5) -> Image.Image: timeout = ClientTimeout(total=total, connect=connect, sock_read=sock_read) async with aiohttp.ClientSession(timeout=timeout) as session: retries = 0 while retries < max_retries: try: if encoding.startswith("http://") or encoding.startswith("https://"): t = time.time() async with session.get(encoding) as response: response.raise_for_status() # Check HTTP status code content = await response.read() # Asynchronously read response content image = Image.open((io.BytesIO(content))) logger.info(f"Image with link {encoding} downloaded, costs time {time.time() - t}s") return image else: logger.error(f"Download failed: {encoding}, it is not a downloadable url") return None except TimeoutError as e: retries += 1 logger.warning(f"Attempt {retries} failed with error: {e}") if retries >= max_retries: logger.error(f"Download {encoding} failed after {max_retries} attempts.") return None except Exception as e: logger.error(f'download {encoding} failed with {e}, {traceback.format_exc()}') return None # Face swap processing function - calling async API def face_swap(source_img: str, target_img: str) -> Optional[str]: """ Synthesize facial features from source image onto target image using async API Parameters: source_img: Source image (containing the face to extract) target_img: Target image (image whose face will be replaced) Returns: Face-swapped image, or None if failed """ # Input validation if source_img is None or target_img is None: logger.warning("Source image or target image is empty") return None, "Source image or target image is empty" try: # Generate task ID input_task_id = str(uuid.uuid4()) # Convert images to base64 encoding source_base64 = image_to_base64(source_img) target_base64 = image_to_base64(target_img) # Step 1: Call async API to initiate face swap task task_id = submit_face_swap_task(source_base64, target_base64, input_task_id) if task_id is None: logger.error("Unable to get face swap task ID") return None, "Unable to get face swap task ID" time.sleep(4) # Step 2: Poll for results with 30-second timeout timeout = API_CONFIG["TIMEOUT"] output_url = poll_face_swap_result(task_id, timeout=timeout) if output_url is None: logger.error("Could not get face swap result within specified time") return None, "Could not get face swap result within specified time" # Step 3: Return result image if output_url == "IMAGE_DETECT_ILLEGAL": logger.error("Illegal image detected") image = Image.open("example/NSFW.jpg") return image, "Illegal image detected" else: image = asyncio.run(download_image_async(output_url, logger)) if image is None: logger.error("Failed to download result image") return None, "Failed to download result image" else: return image, "Face swap successful" except Exception as e: logger.exception(f"Error during face swap process: {str(e)}") return None, "Error during face swap process" # Create Gradio interface with gr.Blocks(theme=gr.themes.Default()) as demo: # Title and introduction gr.Markdown("# AI Face Swap") gr.Markdown("The AI Face Swap Developed by [ZhiQITeAI](https://huggingface.co/izhiqiteai) enables seamless and accurate face swapping in images. yangzhi@zhiqiteai.cn for business inquiries.") with gr.Row(): with gr.Column(): # Source image upload area source_image = gr.Image(label="Source Image", type="filepath", elem_id="source_image") # Target image upload area target_image = gr.Image(label="Target Image", type="filepath", elem_id="target_image") with gr.Column(): # Output image display area output_image = gr.Image(label="Output", type="pil") # Status message status_message = gr.Markdown("Ready. Please upload source and target images, then click Submit.", visible=True) # Button area with gr.Row(): clear_btn = gr.Button("Clear", variant="secondary") submit_btn = gr.Button("Submit", variant="primary") with gr.Row(): with gr.Column(): gr.Markdown("## Source Image Examples") source_exm = gr.Examples( examples=[ os.path.join("example", "source", file) for file in os.listdir(os.path.join("example", "source")) ], examples_per_page=4, inputs=source_image ) with gr.Row(): with gr.Column(): gr.Markdown("## Target Image Examples") target_exm = gr.Examples( examples=[ os.path.join("example", "target", file) for file in os.listdir(os.path.join("example", "target")) ], examples_per_page=4, inputs=target_image ) # Set up button events def clear_inputs(): return None, None, None, "Ready. Please upload source and target images, then click Submit." clear_btn.click( fn=clear_inputs, inputs=None, outputs=[source_image, target_image, output_image, status_message] ) submit_btn.click( fn=face_swap, inputs=[source_image, target_image], outputs=[output_image, status_message] ) # Launch Gradio application if __name__ == "__main__": demo.queue(max_size=30) demo.launch()