diff --git a/workflows/preview_server/miku1f.mp4 b/workflows/preview_server/miku1f.mp4 new file mode 100644 index 0000000..5fca52a Binary files /dev/null and b/workflows/preview_server/miku1f.mp4 differ diff --git a/workflows/preview_server/preview_server.py b/workflows/preview_server/preview_server.py index 89ef883..27eddea 100644 --- a/workflows/preview_server/preview_server.py +++ b/workflows/preview_server/preview_server.py @@ -1,18 +1,20 @@ -import os import signal import subprocess import time -from fastapi import FastAPI, File, UploadFile, HTTPException -from fastapi.responses import FileResponse import shutil from pathlib import Path import asyncio import uuid import os import json -import requests import random +from typing import Dict, Any, Optional + +from fastapi import FastAPI, File, UploadFile, HTTPException, Form +from fastapi.responses import FileResponse +import requests +from jsonpath_ng import jsonpath, parse app = FastAPI() lock = asyncio.Lock() # Create an asyncio lock @@ -20,11 +22,53 @@ lock = asyncio.Lock() # Create an asyncio lock 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" -PROMPT_FILE = os.environ.get("PROMPT_FILE", "prompt.json") +V2V_WORKFLOWS_DIRECTORY = os.environ.get("V2V_WORKFLOWS_DIRECTORY", + "/home/salt/clone/storyteller/V2V-StylesList/V2VWorkflows/") COMFY_ROOT_DIRECTORY = os.environ.get("COMFY_ROOT_DIRECTORY", "/home/salt/clone/storyteller/storyteller-ml/workflows/comfy/ComfyUI") SERVER_START_TIMEOUT = int(os.environ.get("SERVER_START_TIMEOUT", 60)) +STYLE_LOCK = os.environ.get("STYLE_LOCK", None) +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", +} + +allowed_keys = ["style", "positive_prompt", "negative_prompt"] + + +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): @@ -42,9 +86,16 @@ def signal_handler(signum, frame): @app.post("/preview/") -async def upload_video(video: UploadFile = File(...)): +async def preview_request(request: str = Form(...), video: UploadFile = File(...)): + 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") async with lock: # Acquire the lock try: + style_filename = style_to_filename[request["style"]] + positive_prompt = request.get("positive_prompt", None) + negative_prompt = request.get("negative_prompt", None) # Save uploaded video to a temporary directory video_filename = f"temp_videos/{uuid.uuid4()}.mp4" # Use a unique filename Path("temp_videos").mkdir(parents=True, exist_ok=True) # Ensure the directory exists @@ -54,16 +105,32 @@ async def upload_video(video: UploadFile = File(...)): raise HTTPException(status_code=500, detail=f"Failed to save video: {str(e)}") # Process video and generate preview - preview_image_path = await process_video(video_filename) + preview_image_path = await process_video(video_filename, style_filename, positive_prompt, negative_prompt) return FileResponse(preview_image_path, media_type='image/jpeg') -async def process_video(video_path: str) -> Path: - json_prompt = json.load(open(PROMPT_FILE)) +async def process_video(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 / style_json["workflow_api_name"])) + mapping_json = json.load(open(mappings_directory / style_json["mapping_name"])) + + json_mods = style_json["modifications"] + 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) - json_prompt["173"]["inputs"]["seed"] = random.randint(0, 1000000) + try: + json_prompt["173"]["inputs"]["seed"] = random.randint(0, 1000000) + except KeyError: + pass input_video_path = Path(COMFY_ROOT_DIRECTORY) / "input/input.mp4" preview_image_path = Path(COMFY_ROOT_DIRECTORY) / "output/vid2vid/SparseImagePreview_00001_.png" @@ -116,8 +183,79 @@ async def process_video(video_path: str) -> Path: return preview_image_path +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"{new_mod_json.get('positive_prompt', '')}, {pos_in}" + if neg_in: + new_mod_json["negative_prompt"] = f"{new_mod_json.get('negative_prompt', '')}, {neg_in}" + + 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 + + if __name__ == "__main__": import uvicorn + import asyncio + + # install ComfyUI + os.system("python3 install.py") signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGTERM, signal_handler) diff --git a/workflows/preview_server/test_request.py b/workflows/preview_server/test_request.py new file mode 100644 index 0000000..4bc222d --- /dev/null +++ b/workflows/preview_server/test_request.py @@ -0,0 +1,17 @@ +import json + +import requests + +url = 'http://localhost:8000/preview/' +payload = {"style": "comic_book", "positive_prompt": "test pos in", "negative_prompt": "test neg in"} +files = {'video': open('miku1f.mp4', 'rb')} + +response = requests.post(url, data={'request': json.dumps(payload)}, files=files) # Use 'request': payload directly +if response.status_code == 200: + # Success + with open('preview_image.jpeg', 'wb') as f: + f.write(response.content) +else: + # Handle request error + print(f"Request failed: {response.status_code}") + print(response.text)