Files

262 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",
# YuE2 creates a child directory from request["id"] itself.
"--request", str(request_file), "--output", str(ROOT),
]
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
try:
self.wfile.write(chunk)
except (BrokenPipeError, ConnectionResetError):
# Browsers routinely close an old range request after a
# seek or metadata probe. This is not a server failure.
break
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()