mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
424 lines
16 KiB
Python
424 lines
16 KiB
Python
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", 'workflow_api_instantpreview.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"
|
|
}
|
|
|
|
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)
|