diff --git a/media_stream.py b/media_stream.py index 916337d..30976f1 100644 --- a/media_stream.py +++ b/media_stream.py @@ -4,18 +4,21 @@ from PIL import Image import requests import io import logging +import imageio.v2 as imageio +import mimetypes -# Set up logging -logging.basicConfig(level=logging.INFO) - +# --- Setup Logging --- +logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') +# --- Node Categories --- category = "Nilor Nodes 👺" subcategories = { "streaming": "/Streaming", } +# --- MediaStreamInput: Universal Media Downloader --- class MediaStreamInput: """ - A custom node to download an image from a pre-signed URL and provide it as a tensor. + A custom node to download an image/video from a pre-signed URL and provide it as a tensor. """ def __init__(self): pass @@ -41,55 +44,84 @@ class MediaStreamInput: try: response = requests.get(presigned_download_url, timeout=180) response.raise_for_status() + media_bytes = response.content - # Open image from response content - img_bytes = response.content - image_pil = Image.open(io.BytesIO(img_bytes)) + # Use the Content-Type header to determine the file type + content_type = response.headers.get("Content-Type", "") + logging.info(f"Detected Content-Type: {content_type}") - # Convert PIL image to tensor - output_images = [] - output_masks = [] - - image_tensor = torch.from_numpy(np.array(image_pil).astype(np.float32) / 255.0).unsqueeze(0) - - if 'A' in image_pil.getbands(): - mask = torch.from_numpy(np.array(image_pil.getchannel('A')).astype(np.float32) / 255.0).unsqueeze(0) - output_masks.append(mask) - image_tensor = image_tensor[:, :, :, :3] # Drop alpha channel from image - - output_images.append(image_tensor) - - if not output_masks: - # Create a blank mask if one doesn't exist - mask = torch.zeros((1, image_pil.height, image_pil.width), dtype=torch.float32, device="cpu") - output_masks.append(mask) - - images_tensor = torch.cat(output_images, dim=0) - masks_tensor = torch.cat(output_masks, dim=0) - - logging.info("MediaStreamInput: Download and processing successful.") - return (images_tensor, masks_tensor) + if 'video' in content_type: + return self._process_video(media_bytes) + else: # Default to image processing + return self._process_image(media_bytes) except requests.RequestException as e: logging.error(f"MediaStreamInput: Failed to download file: {e}") return (None, None) except Exception as e: - logging.error(f"MediaStreamInput: Failed to process image: {e}") + logging.error(f"MediaStreamInput: Failed to process media: {e}") return (None, None) -class MediaStreamOutput: - """ - A custom node to upload an image tensor to a pre-signed URL. - """ - def __init__(self): - self.output_dir = "output" # Not used directly, but good practice - self.type = "output" + def _process_image(self, image_bytes): + logging.info("Processing as image...") + image_pil = Image.open(io.BytesIO(image_bytes)) + + output_images = [] + output_masks = [] + + # Ensure image is in RGB + rgb_image_pil = image_pil.convert("RGB") + image_tensor = torch.from_numpy(np.array(rgb_image_pil).astype(np.float32) / 255.0).unsqueeze(0) + + if 'A' in image_pil.getbands(): + mask = torch.from_numpy(np.array(image_pil.getchannel('A')).astype(np.float32) / 255.0).unsqueeze(0) + output_masks.append(mask) + + output_images.append(image_tensor) + + if not output_masks: + mask = torch.zeros((1, image_pil.height, image_pil.width), dtype=torch.float32, device="cpu") + output_masks.append(mask) + + images_tensor = torch.cat(output_images, dim=0) + masks_tensor = torch.cat(output_masks, dim=0) + + logging.info("Image processing successful.") + return (images_tensor, masks_tensor) + def _process_video(self, video_bytes): + logging.info("Processing as video...") + frames = [] + with imageio.get_reader(io.BytesIO(video_bytes), format='mp4') as reader: + for frame in reader: + # Convert frame to RGB PIL Image and then to tensor + pil_image = Image.fromarray(frame).convert("RGB") + numpy_image = np.array(pil_image).astype(np.float32) / 255.0 + tensor_frame = torch.from_numpy(numpy_image) + frames.append(tensor_frame) + + if not frames: + raise ValueError("No frames could be read from the video.") + + # Stack frames into a single tensor (batch of images) + video_tensor = torch.stack(frames) + + # THE FIX: The mask must have the same batch dimension as the image tensor. + batch_size, height, width, _ = video_tensor.shape + mask_tensor = torch.zeros((batch_size, height, width), dtype=torch.float32, device="cpu") + + logging.info(f"Video processing successful. Image Shape: {video_tensor.shape}, Mask Shape: {mask_tensor.shape}") + return (video_tensor, mask_tensor) + + +# --- MediaStreamOutput: Universal Media Uploader --- +class MediaStreamOutput: @classmethod def INPUT_TYPES(s): return { "required": { "images": ("IMAGE",), + "format": (["png", "mp4"],), "presigned_upload_url": ("STRING", { "multiline": True, "default": "http://example.com/upload_here" @@ -107,54 +139,66 @@ class MediaStreamOutput: OUTPUT_NODE = True CATEGORY = category + subcategories["streaming"] - def upload(self, images, presigned_upload_url: str, completion_webhook_url: str, prompt=None, extra_pnginfo=None): - try: - # For simplicity, we'll upload the first image of the batch. - # In a real scenario, this might loop and generate multiple upload URLs. - image_tensor = images[0] - - # Convert tensor to PIL Image + def upload(self, images, format, presigned_upload_url, completion_webhook_url, prompt=None, extra_pnginfo=None): + if format == "png": + # For "png", we upload the first image of the batch + self._upload_image(images[0], presigned_upload_url) + elif format == "mp4": + self._upload_video(images, presigned_upload_url) + + # The webhook call will be fully implemented in a later goal (Minor Goal 4.1) + # For now, just log it. + logging.info(f"Completion webhook (not sent): {completion_webhook_url}") + + return {"ui": {"images": []}} + + def _upload_image(self, image_tensor, url): + logging.info("Uploading as PNG image...") + i = 255. * image_tensor.cpu().numpy() + img_pil = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + + buffer = io.BytesIO() + img_pil.save(buffer, format='PNG', compress_level=4) + buffer.seek(0) + + self._perform_upload(buffer, url, 'image/png') + + def _upload_video(self, image_batch_tensor, url): + logging.info(f"Uploading as MP4 video. Frame count: {len(image_batch_tensor)}") + frames = [] + for image_tensor in image_batch_tensor: i = 255. * image_tensor.cpu().numpy() - img_pil = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + frame = np.clip(i, 0, 255).astype(np.uint8) + frames.append(frame) - # Save PIL image to a byte buffer as PNG - buffer = io.BytesIO() - img_pil.save(buffer, format='PNG', compress_level=4) - buffer.seek(0) - - logging.info(f"MediaStreamOutput: Uploading to {presigned_upload_url}") - - # Upload using a PUT request - headers = {'Content-Type': 'image/png'} - response = requests.put(presigned_upload_url, data=buffer, headers=headers, timeout=180) + buffer = io.BytesIO() + # Use mimwrite for format-agnostic writing, though we specify mp4 here + imageio.mimwrite(buffer, frames, format='mp4', fps=30, quality=8) + buffer.seek(0) + + self._perform_upload(buffer, url, 'video/mp4') + + def _perform_upload(self, buffer, url, content_type): + try: + logging.info(f"Uploading to {url} with Content-Type: {content_type}") + headers = {'Content-Type': content_type} + response = requests.put(url, data=buffer.read(), headers=headers, timeout=300) response.raise_for_status() - - logging.info("MediaStreamOutput: Upload successful.") - - # As per Minor Goal 3.2, the completion webhook is NOT called yet. - # This logic will be added in a later step. - # if completion_webhook_url: - # logging.info(f"MediaStreamOutput: Calling completion webhook: {completion_webhook_url}") - # # webhook_response = requests.post(completion_webhook_url, json={"status": "completed"}) - # # webhook_response.raise_for_status() - - return {"ui": {"images": []}} # Required return for output nodes - + logging.info("Upload successful.") except requests.RequestException as e: - logging.error(f"MediaStreamOutput: Failed to upload file: {e}") - return {"ui": {"error": [str(e)]}} + logging.error(f"MediaStreamOutput: Failed to upload media: {e}") + raise # Re-raise the exception to make ComfyUI aware of the failure except Exception as e: - logging.error(f"MediaStreamOutput: Failed to process and upload image: {e}") - return {"ui": {"error": [str(e)]}} + logging.error(f"MediaStreamOutput: Failed to process and upload media: {e}") + raise -# A dictionary that contains all nodes you want to export with their names -# NOTE: names should be globally unique + +# --- Node Mappings --- NODE_CLASS_MAPPINGS = { "MediaStreamInput": MediaStreamInput, "MediaStreamOutput": MediaStreamOutput, } -# A dictionary that contains the friendly/humanly readable titles for the nodes NODE_DISPLAY_NAME_MAPPINGS = { "MediaStreamInput": "👺 Media Stream Input (URL)", "MediaStreamOutput": "👺 Media Stream Output (URL)", diff --git a/requirements.txt b/requirements.txt index a2cc4b9..4c121d6 100644 --- a/requirements.txt +++ b/requirements.txt @@ -4,4 +4,6 @@ fastapi uvicorn python-multipart aiofiles -boto3 \ No newline at end of file +boto3 +imageio +imageio-ffmpeg \ No newline at end of file diff --git a/worker_consumer.py b/worker_consumer.py index 81d2f0f..6e71791 100644 --- a/worker_consumer.py +++ b/worker_consumer.py @@ -22,7 +22,7 @@ def submit_to_local_comfyui(payload: dict) -> bool: """ try: logging.info(f"Submitting job to local ComfyUI at {COMFYUI_URL}") - # The payload from the Brain API is now the complete, ready-to-run prompt. + # The payload from the Brain API is the complete, ready-to-run object. response = requests.post(COMFYUI_URL, json=payload, timeout=20) response.raise_for_status() logging.info(f"Successfully submitted job to local ComfyUI. Response: {response.json()}") @@ -75,8 +75,9 @@ def consume_jobs(): logging.info(f"Received job. Body: {message['Body']}") workflow_payload = json.loads(message['Body']) - # 4. Submit the job to the local ComfyUI server. The 'prompt' field within the body is the actual workflow. - submit_successful = submit_to_local_comfyui(workflow_payload['prompt']) + + # THE FIX: Pass the entire payload, not just the 'prompt' field, to the ComfyUI server. + submit_successful = submit_to_local_comfyui(workflow_payload) if submit_successful: sqs_client.delete_message(