import signal import subprocess import time import shutil from pathlib import Path import asyncio import uuid import os import json import random from typing import Dict, Any, Optional import mimetypes from fastapi import FastAPI, File, UploadFile, HTTPException, Form, Request, Response import requests from jsonpath_ng import jsonpath, parse from moviepy.editor import ImageClip from fastapi.middleware.cors import CORSMiddleware import urllib import websockets # CORS header for which websites can import this content in a frame. # This *technically* shouldn't be required, but the XHR requests are # being made from a controlled frame. # # https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/Content-Security-Policy/frame-ancestors # # NB(bt,2024-03-30): The *.netlify.app is dangerous. CORS_WHITELIST_ORIGINS = [ "https://storyteller.ai", "https://studio.storyteller.ai", "https://studio-testing.studio.storyteller.ai", "https://studio-staging.studio.storyteller.ai", "https://storytellerstudio.netlify.app", "https://branch-name--storytellerstudio.netlify.app", "https://deploy-preview-123--storytellerstudio.netlify.app", "https://pipeline-gottagofast.netlify.app", "https://branch-whatever--pipeline-gottagofast.netlify.app", "https://fakeyou.com", "https://*.fakeyou.com", "https://*.storyteller.ai", "https://*.netlify.app", "http://*.fakeyou.com:7000", "http://*.fakeyou.com:7001", "http://*.fakeyou.com:7002", "http://*.storyteller.ai:7000", "http://*.storyteller.ai:7001", "http://*.storyteller.ai:7002", ] # TODO(kasisnu): Remove ugly globals app = FastAPI() app.add_middleware( CORSMiddleware, allow_origins=CORS_WHITELIST_ORIGINS, allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) SERVER_IP = os.environ.get("SERVER_IP", "127.0.0.1:8188") TIMEOUT_SECONDS = os.environ.get("TIMEOUT_SECONDS", 1000) PROMPT_ENDPOINT = f"http://{SERVER_IP}/prompt" V2V_WORKFLOWS_DIRECTORY = os.environ.get("V2V_WORKFLOWS_DIRECTORY", "/home/kasisnu/go/src/github.com/storytold/V2V-StylesList/V2VWorkflows/") COMFY_ROOT_DIRECTORY = os.environ.get("COMFY_ROOT_DIRECTORY", "/home/kasisnu/go/src/github.com/comfyanonymous/ComfyUI") SERVER_START_TIMEOUT = int(os.environ.get("SERVER_START_TIMEOUT", 60)) STYLE_LOCK = os.environ.get("STYLE_LOCK", None) PREVIEW_WORKFLOW_NAME = os.environ.get("PREVIEW_WORKFLOW_NAME", 'yae_vid2vid_InstantPreview_14-05_API.json') MAPPING_NAME = os.environ.get("MAPPING_NAME", "mapping_two.json") AUTH_URL = os.environ.get("AUTH_URL", "https://api.fakeyou.com/v1/session") # websocket.enableTrace(True) client_id = str(uuid.uuid4()) style_to_filename = { "anime_2_5d": "1_2.5d_anime_model.json", "anime_2d_flat": "2_2d_flat_anime_model.json", "cartoon_3d": "3_3d_cartoon_style.json", "comic_book": "4_comic_book_model.json", "anime_ghibli": "5_ghibli_anime_model.json", "ink_punk": "6_ink_punk.json", "ink_splash": "7_ink_splash.json", "ink_bw_style": "8_ink_w_and_b_style.json", "jojo_style": "9_jojo_style.json", "paper_origami": "10_paper_origami.json", "pixel_art": "11_pixel_art.json", "pop_art": "12_pop_art.json", "realistic_1": "13_realistic_1.json", "realistic_2": "14_realistic_2.json", "anime_retro_neon": "15_retro_neon_anime_90.json", "anime_standard": "16_standard_anime_model.json", "hr_giger": "17_hr_giger.json", "simpsons": "18_simpsons.json", "carnage": "19_carnage.json", "pastel_cute_anime": "20_pastel_cute_anime.json", "bloom_lighting": "21_bloom_lighting.json", "25d_horror": "22_25D_Horror.json", "creepy": "23_creepy.json", "creepy_vhs": "24_creepy_vhs.json", "trail_cam_footage": "25_trail_cam_footage.json", "old_black_white_movie": "26_old_black_white_movie.json", "horror_noir_black_white": "27_horror_noir_black_white.json", "techno_noir_black_white": "28_techno_noir_black_white.json", "black_white_20s": "29_black_white_20s.json", "cyberpunk_anime": "30_cyberpunk_anime.json", "dragonball": "31_dragonball.json", "realistic_matrix": "32_realistic_matrix.json", "realistic_cyberpunk": "33_realistic_cyberpunk.json", "dreamer": "34_dreamer.json", } allowed_keys = ["style", "positive_prompt", "negative_prompt"] def can_use_vst(cookies): #todo(kasisnu): this didn't seem to work right away, might need digging into return True # FOR NOW """Check if the user is authenticated and if the 'video_style_transfer' feature flag is enabled.""" try: response = requests.get(AUTH_URL, cookies=cookies) except requests.exceptions.RequestException as e: print(f"Error during request: {e}") return False if response.status_code == 200: try: response_json = response.json() except ValueError: print("Invalid JSON response") return False if response_json.get("success") and response_json.get("logged_in"): username = response_json["user"]["core_info"]["username"] print(f"User is logged in: {username}") feature_flags = response_json["user"].get("maybe_feature_flags", {}) if "video_style_transfer" in feature_flags: print("VST feature flag is found") return True else: print("VST feature flag not found") else: print("User is not logged in or session is invalid") else: print(f"Failed to authenticate ({response.status_code}): {response.text}") return False def validate_request(request): try: request = json.loads(request) except json.JSONDecodeError: raise HTTPException(status_code=400, detail="Invalid JSON for 'request'") # check for style key if "style" not in request: raise HTTPException(status_code=400, detail="Missing 'style' key in 'request'") style = request["style"] if style not in style_to_filename: raise HTTPException(status_code=400, detail=f"Invalid 'style' key in 'request'. Valid styles: {', '.join(style_to_filename.keys())}") # check for other keys for key in request: if key not in allowed_keys: raise HTTPException(status_code=400, detail=f"Invalid key '{key}' in 'request'. Valid keys: {', '.join(allowed_keys)}") return request def terminate_process(proc): proc.terminate() try: proc.wait(timeout=10) except subprocess.TimeoutExpired: proc.kill() def signal_handler(signum, frame): print("Signal received, shutting down.") terminate_process(comfyui_process) print("ComfyUI server terminated") exit(0) @app.get('/health') def health(): return "All good" @app.post("/preview/") async def preview_request(request_obj: Request, request: str = Form(...), input_file: UploadFile = File(...)): cookies = request_obj.cookies if not can_use_vst(cookies): raise HTTPException(status_code=401, detail="No Video Style Transfer feature flag found") request = validate_request(request) if STYLE_LOCK: if request["style"] != STYLE_LOCK: raise HTTPException(status_code=400, detail=f"This endpoint is for '{STYLE_LOCK}' style only") try: request_style = request["style"] style_filename = style_to_filename[request["style"]] positive_prompt = request.get("positive_prompt", None) negative_prompt = request.get("negative_prompt", None) mime_type = input_file.content_type if not mime_type: mime_type, _ = mimetypes.guess_type(input_file.filename) input_filename_ext = '' if 'video' in mime_type: input_filename_ext = ".mp4" elif 'image' in mime_type: input_filename_ext = mimetypes.guess_extension(mime_type) else: raise HTTPException(status_code=400, detail="Unsupported file type") temp_filename = f"temp_files/{uuid.uuid4()}{input_filename_ext}" Path(temp_filename).parent.mkdir(parents=True, exist_ok=True) with open(temp_filename, "wb") as buffer: shutil.copyfileobj(input_file.file, buffer) if 'image' in mime_type: clip = ImageClip(temp_filename) fps = 24 clip_duration = 1 / fps clip = clip.set_duration(clip_duration) input_filename = f"temp_files/{uuid.uuid4()}.mp4" clip.write_videofile(input_filename, fps=fps) else: input_filename = temp_filename except Exception as e: raise HTTPException(status_code=500, detail=f"Failed to save file: {str(e)}") # Process video (or image converted to video) and generate preview preview_image_path = await process_video(request_style, input_filename, style_filename, positive_prompt, negative_prompt) return Response(preview_image_path, media_type='image/jpeg') async def process_video(request_style: str, video_path: str, style_filename, positive_prompt=None, negative_prompt=None) -> Path: # generate random seed (otherwise video will cache completely and not generate a new output) styles_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "styles" workflow_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "workflows" mappings_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "mappings" style_json = json.load(open(styles_directory / style_filename)) workflow_json = json.load(open(workflow_directory / PREVIEW_WORKFLOW_NAME)) mapping_json = json.load(open(mappings_directory / MAPPING_NAME)) json_mods = style_json["modifications"] json_mods['cn_soft_edge_strength'] = 1.0 jsonpath_mods = get_jsonpath_mods(json_mods, mapping_json, pos_in=positive_prompt, neg_in=negative_prompt) print(*jsonpath_mods.items(), sep="\n") json_prompt = apply_jsonpath_mods(jsonpath_mods, workflow_json) # generate random seed (otherwise video will cache completely and not generate a new output) try: json_prompt["173"]["inputs"]["seed"] = random.randint(0, 1000000) except KeyError: pass input_video_path = Path(COMFY_ROOT_DIRECTORY) / "input/input.mp4" shutil.move(video_path, input_video_path) p = {"prompt": json_prompt, "client_id": client_id, "extra_data": {"style": request_style}} data = json.dumps(p).encode('utf-8') execution_start_time = time.perf_counter() response = requests.post(PROMPT_ENDPOINT, data=data) if response.status_code == 200: print("Job enqueued successfully") else: raise HTTPException(status_code=500, detail=f"Failed to queue job") prompt_id = response.json()["prompt_id"] async with websockets.connect("ws://{}/ws?clientId={}".format(SERVER_IP, client_id)) as ws: images = await get_images(ws, prompt_id, '1517') current_time = time.perf_counter() execution_time = current_time - execution_start_time print('request executed in {:.2f} seconds'.format(execution_time)) return images['1517'][0] def get_jsonpath_mods(style_mods: Dict[str, Any], mapping_json: Dict[str, Any], pos_in: Optional[str] = None, neg_in: Optional[str] = None) -> Dict[str, Any]: modifications = {} new_mod_json = style_mods.copy() # Process "loras" differently loras = new_mod_json.get("loras", []) if len(loras) > 8: raise ValueError("Too many loras, max is 8") for index, lora in enumerate(loras): if "name" in lora and "strength" in lora: new_mod_json[f"lora_{index + 1}_strength"] = lora["strength"] new_mod_json[f"lora_{index + 1}_name"] = lora["name"] # Exclude "loras" from further processing new_mod_json.pop("loras", None) # Handling positive and negative prompts if pos_in: new_mod_json["positive_prompt"] = f"{pos_in}, {new_mod_json.get('positive_prompt', '')}," if neg_in: new_mod_json["negative_prompt"] = f"{neg_in}, {new_mod_json.get('negative_prompt', '')}," for key, value in new_mod_json.items(): mapping_key = f"$.{key}" jsonpath_expr = parse(mapping_key) # Finding matches using jsonpath_ng matches = [match.value for match in jsonpath_expr.find(mapping_json)] if matches: mapping_value = matches[0] modifications[mapping_value] = value else: print(f"No mapping found for key '{key}'") return modifications def apply_jsonpath_mods(jsonpath_mods: Dict[str, Any], workflow_json: Dict[str, Any]) -> Dict[str, Any]: for key, value in jsonpath_mods.items(): # jsonpath-ng expects numbers to be wrapped in quotes parts = key.split(".") new_parts = [f"'{part}'" if part.isdigit() else part for part in parts] key = ".".join(new_parts) jsonpath_expr = parse(key) matches = jsonpath_expr.find(workflow_json) if matches: parent = matches[0].context.value if isinstance(parent, dict): last_part = new_parts[-1].strip("'") parent[last_part] = value elif isinstance(parent, list): index = int(new_parts[-1].strip("'")) if 0 <= index < len(parent): parent[index] = value else: print(f"Index {index} is out of bounds for key '{key}'") else: print(f"No mapping found for key '{key}'") return workflow_json def get_image(filename, subfolder, folder_type): data = {"filename": filename, "subfolder": subfolder, "type": folder_type} url_values = urllib.parse.urlencode(data) with urllib.request.urlopen("http://{}/view?{}".format(SERVER_IP, url_values)) as response: return response.read() def get_history(prompt_id): with urllib.request.urlopen("http://{}/history/{}".format(SERVER_IP, prompt_id)) as response: return json.loads(response.read()) async def get_images(ws, prompt_id, filter_node_id=None): output_images = {} while True: out = await ws.recv() if isinstance(out, str): message = json.loads(out) if message['type'] == 'executing': data = message['data'] if data['node'] is None and data['prompt_id'] == prompt_id: break #Execution is done else: continue #previews are binary data history = get_history(prompt_id)[prompt_id] for o in history['outputs']: for node_id in history['outputs']: if filter_node_id is not None and node_id != filter_node_id: continue node_output = history['outputs'][node_id] images_output = [] if 'images' in node_output: for image in node_output['images']: image_data = get_image(image['filename'], image['subfolder'], image['type']) images_output.append(image_data) output_images[node_id] = images_output return output_images if __name__ == "__main__": import uvicorn # install ComfyUI # os.system("python3 install.py") signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGTERM, signal_handler) comfyui_process = subprocess.Popen(["../venv/bin/python", "main.py", "--highvram", "--disable-metadata"], cwd=COMFY_ROOT_DIRECTORY) # block until comfyUI server is up start_time = time.time() started = False while time.time() - start_time < SERVER_START_TIMEOUT: try: response = requests.get(PROMPT_ENDPOINT, timeout=5) if response.status_code == 200: print("Server is healthy") started = True break except requests.exceptions.ConnectionError: print(f"Server is not up yet. Time elapsed: {time.time() - start_time:.2f}") time.sleep(1) if not started: print(f"Server failed to start in {SERVER_START_TIMEOUT} seconds") exit(1) print("ComfyUI server running, starting FastAPI server") uvicorn.run(app, host="0.0.0.0", port=8000) # cleanup terminate_process(comfyui_process)