mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
379 lines
14 KiB
Python
379 lines
14 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
|
|
from fastapi.responses import FileResponse
|
|
import requests
|
|
from jsonpath_ng import jsonpath, parse
|
|
from moviepy.editor import ImageClip
|
|
|
|
app = FastAPI()
|
|
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"
|
|
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/preview_server/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", None)
|
|
MAPPING_NAME = os.environ.get("MAPPING_NAME", "mapping_two.json")
|
|
AUTH_URL = os.environ.get("AUTH_URL", "https://api.fakeyou.com/v1/session")
|
|
|
|
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"]
|
|
|
|
# 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_FRAME_ANCESTORS = [
|
|
"frame-ancestors 'self'",
|
|
"https://fakeyou.com",
|
|
"https://*.fakeyou.com",
|
|
"https://storyteller.ai",
|
|
"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",
|
|
";",
|
|
]
|
|
|
|
|
|
def can_use_vst(cookies):
|
|
"""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.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")
|
|
|
|
async with lock:
|
|
try:
|
|
style_filename = style_to_filename[request["style"]]
|
|
positive_prompt = request.get("positive_prompt", None)
|
|
negative_prompt = request.get("negative_prompt", None)
|
|
|
|
temp_filename = f"temp_files/{uuid.uuid4()}"
|
|
Path(temp_filename).parent.mkdir(parents=True, exist_ok=True)
|
|
with open(temp_filename, "wb") as buffer:
|
|
shutil.copyfileobj(input_file.file, buffer)
|
|
|
|
mime_type, _ = mimetypes.guess_type(input_file.filename)
|
|
|
|
if 'video' in mime_type:
|
|
input_filename = f"{temp_filename}.mp4"
|
|
elif 'image' in mime_type:
|
|
input_filename = f"{temp_filename}.{mime_type.split('/')[-1]}"
|
|
else:
|
|
raise HTTPException(status_code=400, detail="Unsupported file type")
|
|
|
|
Path(temp_filename).rename(input_filename)
|
|
|
|
if 'image' in mime_type:
|
|
clip = ImageClip(input_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)
|
|
|
|
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(input_filename, style_filename, positive_prompt, negative_prompt)
|
|
|
|
csp_frame_ancestors_value = ' '.join(CORS_FRAME_ANCESTORS[:-1])
|
|
headers = {
|
|
"Access-Control-Allow-Origin": "*", # Be cautious with this setting in production environments.
|
|
"Content-Security-Policy": f"frame-ancestors {csp_frame_ancestors_value}",
|
|
}
|
|
|
|
return FileResponse(preview_image_path, media_type='image/jpeg', headers=headers)
|
|
|
|
|
|
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 / PREVIEW_WORKFLOW_NAME))
|
|
# mapping_json = json.load(open(mappings_directory / style_json["mapping_name"]))
|
|
mapping_json = json.load(open(mappings_directory / 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)
|
|
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"
|
|
shutil.move(video_path, input_video_path)
|
|
if os.path.exists(preview_image_path):
|
|
os.remove(preview_image_path)
|
|
|
|
response = requests.get(PROMPT_ENDPOINT, timeout=5)
|
|
if response.status_code == 200:
|
|
print("Server is healthy")
|
|
else:
|
|
print("Server is unhealthy")
|
|
print(f"Status code: {response.status_code}")
|
|
print(f"Response: {response.text}")
|
|
raise HTTPException(status_code=500, detail="Server is unhealthy")
|
|
|
|
# make sure there are no jobs already
|
|
queue_json = response.json()
|
|
if queue_json["exec_info"]["queue_remaining"]:
|
|
print("There are jobs in the queue already. Exiting...")
|
|
raise HTTPException(status_code=500, detail="There are jobs in the queue already (this shouldn't happen)")
|
|
else:
|
|
print("Queue is empty")
|
|
|
|
p = {"prompt": json_prompt}
|
|
data = json.dumps(p).encode('utf-8')
|
|
response = requests.post(PROMPT_ENDPOINT, data=data)
|
|
if response.status_code == 200:
|
|
print("Job enqueued successfully")
|
|
else:
|
|
print("Job failed to queue")
|
|
print(f"Status code: {response.status_code}")
|
|
print(f"Response: {response.text}")
|
|
raise HTTPException(status_code=500, detail=f"Failed to queue job")
|
|
|
|
start_time = time.time()
|
|
while True:
|
|
response = requests.get(PROMPT_ENDPOINT, timeout=5)
|
|
queue_json = response.json()
|
|
if queue_json["exec_info"]["queue_remaining"]:
|
|
pass
|
|
else:
|
|
print("Job is done")
|
|
break
|
|
if time.time() - start_time > TIMEOUT_SECONDS:
|
|
print("Timeout reached and job didn't finish!")
|
|
raise HTTPException(status_code=500, detail="Timeout reached and job didn't finish!")
|
|
time.sleep(1)
|
|
|
|
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)
|
|
|
|
comfyui_process = subprocess.Popen(["venv/bin/python", "main.py"], 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)
|