fix workflow
This commit is contained in:
@@ -1,11 +1,11 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Клиент для DaSiWa API Server.
|
||||
Запускается на ТВОЁМ ПК. Отправляет подписанные запросы на сервер.
|
||||
Клиент для DaSiWa API Server (асинхронный, как RunPod).
|
||||
Запускается на ТВОЁМ ПК. Отправляет задачу, поллит статус, забирает результат.
|
||||
|
||||
Использование:
|
||||
python client.py --server http://<ip>:5000 --image photo.png --prompt "woman dancing"
|
||||
python client.py --server http://<ip>:5000 --image start.png --last-image end.png --prompt "smooth transition"
|
||||
python client.py --server http://<ip>:8080 --image photo.png --prompt "woman dancing"
|
||||
python client.py --server http://<ip>:8080 --image start.png --last-image end.png --prompt "smooth transition"
|
||||
"""
|
||||
|
||||
import argparse
|
||||
@@ -40,29 +40,58 @@ def image_to_base64(path: str) -> str:
|
||||
return base64.b64encode(f.read()).decode()
|
||||
|
||||
|
||||
def send_request(server_url: str, payload: dict, client_id: str, secret_key: str) -> dict:
|
||||
"""Отправляет подписанный запрос на сервер."""
|
||||
def signed_post(server_url: str, path: str, payload: dict, client_id: str, secret_key: str):
|
||||
"""Отправляет подписанный POST запрос."""
|
||||
body = json.dumps(payload).encode("utf-8")
|
||||
auth_headers = sign_request(body, secret_key, client_id)
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
**auth_headers
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
f"{server_url}/generate",
|
||||
data=body,
|
||||
headers=headers,
|
||||
timeout=600
|
||||
)
|
||||
|
||||
headers = {"Content-Type": "application/json", **auth_headers}
|
||||
response = requests.post(f"{server_url}{path}", data=body, headers=headers, timeout=30)
|
||||
return response.status_code, response.json()
|
||||
|
||||
|
||||
def signed_get(server_url: str, path: str, client_id: str, secret_key: str):
|
||||
"""Отправляет подписанный GET запрос."""
|
||||
body = b""
|
||||
auth_headers = sign_request(body, secret_key, client_id)
|
||||
response = requests.get(f"{server_url}{path}", headers=auth_headers, timeout=30)
|
||||
return response.status_code, response.json()
|
||||
|
||||
|
||||
def submit_job(server_url: str, payload: dict, client_id: str, secret_key: str):
|
||||
"""Отправляет задачу на генерацию. Возвращает job_id."""
|
||||
code, data = signed_post(server_url, "/run", payload, client_id, secret_key)
|
||||
if code != 200:
|
||||
raise RuntimeError(f"Submit failed ({code}): {data.get('error', data)}")
|
||||
return data["id"]
|
||||
|
||||
|
||||
def wait_for_completion(server_url: str, job_id: str, client_id: str, secret_key: str,
|
||||
poll_interval: int = 5, max_wait: int = 1800):
|
||||
"""Поллит статус задачи до завершения."""
|
||||
start = time.time()
|
||||
while time.time() - start < max_wait:
|
||||
code, data = signed_get(server_url, f"/status/{job_id}", client_id, secret_key)
|
||||
if code != 200:
|
||||
raise RuntimeError(f"Status check failed ({code}): {data}")
|
||||
|
||||
status = data.get("status")
|
||||
elapsed = int(time.time() - start)
|
||||
|
||||
if status == "COMPLETED":
|
||||
print(f"\r✅ COMPLETED ({elapsed}s)")
|
||||
return data
|
||||
elif status == "FAILED":
|
||||
raise RuntimeError(f"Job failed: {data.get('error', 'Unknown error')}")
|
||||
else:
|
||||
print(f"\r⏳ {status}... ({elapsed}s)", end="", flush=True)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
raise RuntimeError(f"Timeout waiting for job ({max_wait}s)")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="DaSiWa API Client")
|
||||
parser.add_argument("--server", required=True, help="Server URL, e.g. http://1.2.3.4:5000")
|
||||
parser = argparse.ArgumentParser(description="DaSiWa API Client (async)")
|
||||
parser.add_argument("--server", required=True, help="Server URL, e.g. http://1.2.3.4:8080")
|
||||
parser.add_argument("--image", required=True, help="Path to first frame image")
|
||||
parser.add_argument("--last-image", default=None, help="Path to last frame image (FLF2V mode)")
|
||||
parser.add_argument("--prompt", required=True, help="Text prompt")
|
||||
@@ -74,10 +103,12 @@ def main():
|
||||
parser.add_argument("--cfg", type=float, default=1.0)
|
||||
parser.add_argument("--seed", type=int, default=-1)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--poll-interval", type=int, default=5, help="Status poll interval (seconds)")
|
||||
parser.add_argument("--output", "-o", default="output.mp4", help="Output video path")
|
||||
args = parser.parse_args()
|
||||
|
||||
keys = load_keys()
|
||||
cid, secret = keys["client_id"], keys["secret_key"]
|
||||
|
||||
# Формируем payload
|
||||
payload = {
|
||||
@@ -102,31 +133,45 @@ def main():
|
||||
print(f"🎬 Режим: I2V (image to video)")
|
||||
|
||||
print(f"📐 {args.width}x{args.height}, {args.length} frames, {args.steps} steps")
|
||||
print(f"📤 Отправляю запрос на {args.server}...")
|
||||
|
||||
start = time.time()
|
||||
status_code, result = send_request(
|
||||
args.server, payload, keys["client_id"], keys["secret_key"]
|
||||
)
|
||||
elapsed = time.time() - start
|
||||
# 1. Submit job
|
||||
print(f"📤 Отправляю задачу на {args.server}...")
|
||||
try:
|
||||
job_id = submit_job(args.server, payload, cid, secret)
|
||||
except RuntimeError as e:
|
||||
print(f"❌ {e}")
|
||||
sys.exit(1)
|
||||
print(f"📝 Job ID: {job_id}")
|
||||
|
||||
if status_code != 200:
|
||||
print(f"❌ Ошибка {status_code}: {result.get('error', 'Unknown')}")
|
||||
if "detail" in result:
|
||||
print(f" Детали: {result['detail']}")
|
||||
# 2. Poll for completion
|
||||
print(f"⏳ Жду результат (поллинг каждые {args.poll_interval}s)...")
|
||||
try:
|
||||
result = wait_for_completion(args.server, job_id, cid, secret,
|
||||
poll_interval=args.poll_interval)
|
||||
except RuntimeError as e:
|
||||
print(f"\n❌ {e}")
|
||||
sys.exit(1)
|
||||
|
||||
if "video" in result:
|
||||
video_bytes = base64.b64decode(result["video"])
|
||||
with open(args.output, "wb") as f:
|
||||
f.write(video_bytes)
|
||||
print(f"✅ Видео сохранено: {args.output} ({len(video_bytes) / 1024 / 1024:.1f} MB)")
|
||||
print(f"⏱ Время: {elapsed:.1f}s (сервер: {result.get('elapsed', '?')}s)")
|
||||
print(f"🌱 Seed: {result.get('seed', '?')}")
|
||||
else:
|
||||
print(f"❌ Ошибка: {result.get('error', 'No video in response')}")
|
||||
# 3. Save video
|
||||
output = result.get("output", {})
|
||||
video_b64 = output.get("video")
|
||||
if not video_b64:
|
||||
print(f"❌ Нет видео в ответе")
|
||||
sys.exit(1)
|
||||
|
||||
video_bytes = base64.b64decode(video_b64)
|
||||
with open(args.output, "wb") as f:
|
||||
f.write(video_bytes)
|
||||
|
||||
print(f"✅ Видео сохранено: {args.output} ({len(video_bytes) / 1024 / 1024:.1f} MB)")
|
||||
print(f"⏱ Сервер: {output.get('elapsed', '?')}s | Seed: {output.get('seed', '?')} | Mode: {output.get('mode', '?')}")
|
||||
|
||||
# 4. Purge job from server memory
|
||||
try:
|
||||
signed_post(args.server, f"/purge/{job_id}", {}, cid, secret)
|
||||
except Exception:
|
||||
pass # не критично
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
Reference in New Issue
Block a user