256 lines
10 KiB
Python
256 lines
10 KiB
Python
#!/usr/bin/env python3
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import mimetypes
|
|
import os
|
|
import random
|
|
import re
|
|
import shutil
|
|
import subprocess
|
|
import threading
|
|
import time
|
|
from http import HTTPStatus
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
from urllib.parse import unquote, urlparse
|
|
|
|
|
|
HOST = os.getenv("YUE2_UI_HOST", "0.0.0.0")
|
|
PORT = int(os.getenv("YUE2_UI_PORT", "8014"))
|
|
ROOT = Path("/workspace/runs")
|
|
REQUESTS = ROOT / ".playground_requests"
|
|
LOGS = ROOT / ".playground_logs"
|
|
INDEX = Path(__file__).with_name("index.html")
|
|
SAFE_ID = re.compile(r"^[a-z0-9][a-z0-9_-]{0,63}$")
|
|
LOCK = threading.Lock()
|
|
JOBS: dict[str, dict] = {}
|
|
ACTIVE: str | None = None
|
|
|
|
|
|
def now_ms() -> int:
|
|
return int(time.time() * 1000)
|
|
|
|
|
|
def existing_jobs() -> list[dict]:
|
|
found: list[dict] = []
|
|
for result_file in ROOT.glob("*/result.json"):
|
|
try:
|
|
result = json.loads(result_file.read_text(encoding="utf-8"))
|
|
request_file = result_file.parent / "request.json"
|
|
request = json.loads(request_file.read_text(encoding="utf-8"))
|
|
found.append({
|
|
"id": result_file.parent.name,
|
|
"state": "complete",
|
|
"style": request.get("style", ""),
|
|
"seed": request.get("seed"),
|
|
"instrumental": not bool(request.get("lyrics")),
|
|
"audio_seconds": result.get("audio_seconds"),
|
|
"elapsed_seconds": (result.get("timing") or {}).get("e2e_seconds"),
|
|
"audio_url": f"/audio/{result_file.parent.name}",
|
|
"created": int(result_file.stat().st_mtime * 1000),
|
|
})
|
|
except (OSError, ValueError, TypeError):
|
|
continue
|
|
return sorted(found, key=lambda item: item["created"], reverse=True)
|
|
|
|
|
|
def snapshot() -> dict:
|
|
with LOCK:
|
|
live = [dict(item) for item in JOBS.values()]
|
|
active = ACTIVE
|
|
known = {item["id"] for item in live}
|
|
live.extend(item for item in existing_jobs() if item["id"] not in known)
|
|
return {"active": active, "jobs": sorted(live, key=lambda item: item["created"], reverse=True)}
|
|
|
|
|
|
def run_job(job_id: str, request: dict) -> None:
|
|
global ACTIVE
|
|
output = ROOT / job_id
|
|
request_file = REQUESTS / f"{job_id}.json"
|
|
log_file = LOGS / f"{job_id}.log"
|
|
command = [
|
|
"yue2", "generate", "--offline", "--device", "cuda:0", "--budget", "16",
|
|
"--request", str(request_file), "--output", str(output),
|
|
]
|
|
started = time.monotonic()
|
|
try:
|
|
with log_file.open("w", encoding="utf-8") as log:
|
|
process = subprocess.Popen(command, stdout=log, stderr=subprocess.STDOUT, text=True)
|
|
with LOCK:
|
|
JOBS[job_id]["pid"] = process.pid
|
|
code = process.wait()
|
|
if code != 0:
|
|
raise RuntimeError(f"YuE2 wurde mit Exit-Code {code} beendet")
|
|
result = json.loads((output / "result.json").read_text(encoding="utf-8"))
|
|
update = {
|
|
"state": "complete",
|
|
"audio_seconds": result.get("audio_seconds"),
|
|
"elapsed_seconds": round(time.monotonic() - started, 1),
|
|
"audio_url": f"/audio/{job_id}",
|
|
}
|
|
except Exception as exc:
|
|
tail = ""
|
|
try:
|
|
tail = "\n".join(log_file.read_text(encoding="utf-8", errors="replace").splitlines()[-30:])
|
|
except OSError:
|
|
pass
|
|
update = {"state": "failed", "error": str(exc), "log": tail,
|
|
"elapsed_seconds": round(time.monotonic() - started, 1)}
|
|
with LOCK:
|
|
JOBS[job_id].update(update)
|
|
ACTIVE = None
|
|
|
|
|
|
class Handler(BaseHTTPRequestHandler):
|
|
server_version = "YuE2Playground/1.0"
|
|
|
|
def log_message(self, fmt: str, *args: object) -> None:
|
|
print(f"{self.address_string()} - {fmt % args}", flush=True)
|
|
|
|
def json_response(self, status: int, payload: object) -> None:
|
|
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
|
self.send_response(status)
|
|
self.send_header("Content-Type", "application/json; charset=utf-8")
|
|
self.send_header("Content-Length", str(len(body)))
|
|
self.send_header("Cache-Control", "no-store")
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
|
|
def do_GET(self) -> None: # noqa: N802
|
|
path = urlparse(self.path).path
|
|
if path == "/health":
|
|
self.json_response(200, {"status": "ok", "active": ACTIVE})
|
|
return
|
|
if path == "/api/jobs":
|
|
self.json_response(200, snapshot())
|
|
return
|
|
if path.startswith("/audio/"):
|
|
self.send_audio(unquote(path.removeprefix("/audio/")))
|
|
return
|
|
if path in {"/", "/index.html"}:
|
|
body = INDEX.read_bytes()
|
|
self.send_response(200)
|
|
self.send_header("Content-Type", "text/html; charset=utf-8")
|
|
self.send_header("Content-Length", str(len(body)))
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
return
|
|
self.send_error(404)
|
|
|
|
def do_POST(self) -> None: # noqa: N802
|
|
global ACTIVE
|
|
if urlparse(self.path).path != "/api/jobs":
|
|
self.send_error(404)
|
|
return
|
|
try:
|
|
length = int(self.headers.get("Content-Length", "0"))
|
|
if length <= 0 or length > 65536:
|
|
raise ValueError("Ungültige Anfragegröße")
|
|
data = json.loads(self.rfile.read(length))
|
|
style = str(data.get("style", "")).strip()
|
|
lyrics = str(data.get("lyrics", "")).strip()
|
|
instrumental = bool(data.get("instrumental", False))
|
|
cot = str(data.get("cot", "full"))
|
|
if not style or len(style) > 3000:
|
|
raise ValueError("Bitte eine Stilbeschreibung mit höchstens 3000 Zeichen eingeben")
|
|
if len(lyrics) > 20000:
|
|
raise ValueError("Der Liedtext ist zu lang")
|
|
if cot not in {"full", "off"}:
|
|
raise ValueError("Unbekannter Planungsmodus")
|
|
if not instrumental and not lyrics:
|
|
raise ValueError("Für einen Song mit Gesang fehlt der Liedtext")
|
|
if instrumental:
|
|
lyrics = ""
|
|
if "instrumental" not in style.casefold():
|
|
style = "Instrumental, no vocals, " + style
|
|
raw_seed = data.get("seed")
|
|
seed = int(raw_seed) if str(raw_seed).strip() else random.SystemRandom().randrange(1, 2**31)
|
|
if not 0 <= seed < 2**32:
|
|
raise ValueError("Seed muss zwischen 0 und 4294967295 liegen")
|
|
except (ValueError, TypeError, json.JSONDecodeError) as exc:
|
|
self.json_response(400, {"error": str(exc)})
|
|
return
|
|
|
|
with LOCK:
|
|
if ACTIVE is not None:
|
|
self.json_response(409, {"error": f"Auftrag {ACTIVE} läuft bereits"})
|
|
return
|
|
job_id = time.strftime("song-%Y%m%d-%H%M%S") + f"-{seed % 10000:04d}"
|
|
request = {"id": job_id, "style": style, "lyrics": lyrics, "cot": cot, "seed": seed}
|
|
REQUESTS.mkdir(parents=True, exist_ok=True)
|
|
LOGS.mkdir(parents=True, exist_ok=True)
|
|
request_file = REQUESTS / f"{job_id}.json"
|
|
request_file.write_text(json.dumps(request, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
|
JOBS[job_id] = {"id": job_id, "state": "running", "style": style,
|
|
"seed": seed, "instrumental": instrumental,
|
|
"created": now_ms(), "elapsed_seconds": 0}
|
|
ACTIVE = job_id
|
|
threading.Thread(target=run_job, args=(job_id, request), daemon=True).start()
|
|
self.json_response(HTTPStatus.ACCEPTED, JOBS[job_id])
|
|
|
|
def do_DELETE(self) -> None: # noqa: N802
|
|
path = urlparse(self.path).path
|
|
job_id = unquote(path.removeprefix("/api/jobs/"))
|
|
if not path.startswith("/api/jobs/") or not SAFE_ID.fullmatch(job_id):
|
|
self.send_error(404)
|
|
return
|
|
with LOCK:
|
|
if ACTIVE == job_id:
|
|
self.json_response(409, {"error": "Ein laufender Auftrag kann nicht gelöscht werden"})
|
|
return
|
|
JOBS.pop(job_id, None)
|
|
shutil.rmtree(ROOT / job_id, ignore_errors=True)
|
|
for directory, suffix in ((REQUESTS, ".json"), (LOGS, ".log")):
|
|
try:
|
|
(directory / f"{job_id}{suffix}").unlink()
|
|
except FileNotFoundError:
|
|
pass
|
|
self.json_response(200, {"status": "deleted", "id": job_id})
|
|
|
|
def send_audio(self, job_id: str) -> None:
|
|
if not SAFE_ID.fullmatch(job_id):
|
|
self.send_error(404)
|
|
return
|
|
path = ROOT / job_id / "audio.flac"
|
|
if not path.is_file():
|
|
self.send_error(404)
|
|
return
|
|
size = path.stat().st_size
|
|
start, end = 0, size - 1
|
|
status = 200
|
|
range_header = self.headers.get("Range", "")
|
|
if range_header.startswith("bytes="):
|
|
try:
|
|
left, right = range_header[6:].split("-", 1)
|
|
start = int(left) if left else 0
|
|
end = min(int(right), size - 1) if right else size - 1
|
|
if start < 0 or start > end:
|
|
raise ValueError
|
|
status = 206
|
|
except ValueError:
|
|
self.send_error(416)
|
|
return
|
|
self.send_response(status)
|
|
self.send_header("Content-Type", mimetypes.guess_type(path.name)[0] or "audio/flac")
|
|
self.send_header("Accept-Ranges", "bytes")
|
|
self.send_header("Content-Length", str(end - start + 1))
|
|
if status == 206:
|
|
self.send_header("Content-Range", f"bytes {start}-{end}/{size}")
|
|
self.end_headers()
|
|
with path.open("rb") as source:
|
|
source.seek(start)
|
|
remaining = end - start + 1
|
|
while remaining:
|
|
chunk = source.read(min(1024 * 1024, remaining))
|
|
if not chunk:
|
|
break
|
|
self.wfile.write(chunk)
|
|
remaining -= len(chunk)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
ROOT.mkdir(parents=True, exist_ok=True)
|
|
print(f"YuE2 Playground listening on {HOST}:{PORT}", flush=True)
|
|
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
|