""" Point d'entree du microservice. Architecture de reference (dev, GPU dedie - pod) : une app FastAPI unique qui 1. expose un endpoint WebSocket /ws/edit (protocole reutilisable par n'importe quel client - pas seulement l'UI de demo), 2. monte l'UI Gradio sur "/" dans le MEME process Python (`gradio.mount_gradio_app`), pour que @spaces.GPU (Hugging Face Spaces + ZeroGPU) reste valide - le decorateur attend d'etre appele depuis le process reconnu comme le Space. Repli sous ZeroGPU reel (voir projet_2_avancement.md, Etape 6) : ce montage FastAPI+uvicorn.run() custom demarre correctement (modeles charges, port lie) puis le process est arrete silencieusement quelques centaines de ms apres, sans traceback - tous les exemples ZeroGPU officiels utilisent `demo.launch()` directement, jamais un uvicorn.run() manuel, et le superviseur ZeroGPU semble s'attendre a ce cycle de vie precis. Sous ZeroGPU (detecte via pipeline.IS_ZERO_GPU), on bascule donc sur un `demo.launch()` Gradio natif, SANS le endpoint WebSocket - le protocole FastAPI+WebSocket reste demontre et valide en dev, mais n'est pas ce qui tourne sur le Space public. """ import os # DOIT s'executer avant TOUT import touchant huggingface_hub (gradio inclus - # gradio importe huggingface_hub en interne). huggingface_hub fige HF_HOME # comme constante de module des son propre import ; le corriger plus tard # (ex. dans pipeline.py, importe apres gradio) ne change plus rien, la # constante est deja calculee avec l'ancienne valeur. Cf. PermissionError # rencontree au premier deploiement Spaces : /home/user/.cache/huggingface # est ecrit pendant le BUILD (preload_from_hub) mais pas inscriptible par # l'utilisateur runtime - Florence-2 (trust_remote_code) a besoin d'ecrire # dans /modules/. _hf_home = os.environ.get("HF_HOME", os.path.expanduser("~/.cache/huggingface")) try: os.makedirs(os.path.join(_hf_home, "modules"), exist_ok=True) except PermissionError: os.environ["HF_HOME"] = "/tmp/hf_home" print(f"[app] HF_HOME ({_hf_home}) non inscriptible au runtime -> repli sur /tmp/hf_home") import base64 import io import json import gradio as gr from fastapi import FastAPI, WebSocket, WebSocketDisconnect from PIL import Image from pydantic import ValidationError import pipeline import runner from gradio_ui import build_ui from ws_protocol import ConfirmRequest, LocateRequest def _b64_to_image(b64: str) -> Image.Image: return Image.open(io.BytesIO(base64.b64decode(b64))).convert("RGB") def _image_to_b64(img: Image.Image) -> str: buf = io.BytesIO() img.save(buf, format="PNG") return base64.b64encode(buf.getvalue()).decode() def _build_websocket_app() -> FastAPI: """Architecture de reference : FastAPI + endpoint WebSocket /ws/edit + UI Gradio montee dans le meme process (cf. docstring en tete de fichier). Utilisee en dev (GPU dedie) ; PAS utilisee sous ZeroGPU (voir _main).""" app = FastAPI(title="genai-image-editor") @app.websocket("/ws/edit") async def ws_edit(websocket: WebSocket): await websocket.accept() # Etat de session : le contexte (image + masque valide) vit le temps # de cette connexion, entre la phase "locate" et la phase "confirm". session = {"image": None, "dilated_mask": None} try: while True: raw = await websocket.receive_text() try: payload = json.loads(raw) except json.JSONDecodeError: await websocket.send_json({"type": "error", "message": "JSON invalide."}) continue msg_type = payload.get("type") if msg_type == "locate": try: req = LocateRequest.model_validate(payload) except ValidationError as e: await websocket.send_json({"type": "error", "message": str(e)}) continue image = _b64_to_image(req.image_b64) session["image"] = image try: result = await runner.run_locate(image, req.text_query) except Exception as e: await websocket.send_json({"type": "error", "message": str(e)}) continue session["dilated_mask"] = Image.open(io.BytesIO(result.dilated_mask_png)).convert("L") await websocket.send_json({ "type": "mask_ready", "mask_b64": base64.b64encode(result.dilated_mask_png).decode(), "box": result.box, "clip_confidence": result.clip_confidence, }) elif msg_type == "confirm": try: req = ConfirmRequest.model_validate(payload) except ValidationError as e: await websocket.send_json({"type": "error", "message": str(e)}) continue if session["image"] is None or session["dilated_mask"] is None: await websocket.send_json({ "type": "error", "message": "Aucun masque en attente - envoyer 'locate' d'abord.", }) continue try: async for kind, *rest in runner.run_inpaint_stream( session["image"], session["dilated_mask"], req.prompt, req.negative_prompt, req.steps, req.guidance_scale, req.seed, ): if kind == "progress": step, total = rest await websocket.send_json({ "type": "progress", "stage": "diffusion", "step": step, "total": total, }) elif kind == "result": (image_result,) = rest await websocket.send_json({ "type": "result", "image_b64": _image_to_b64(image_result), }) except Exception as e: await websocket.send_json({"type": "error", "message": str(e)}) else: await websocket.send_json({"type": "error", "message": f"Type de message inconnu : {msg_type}"}) except WebSocketDisconnect: pass # ssr_mode=False : sans ca, Gradio 5.x demarre un serveur Node.js # compagnon (rendu cote serveur du frontend) qui tente de se binder sur # un port separe (7861, avec repli 7862...) - source d'un "address # already in use" observe lors du premier essai de deploiement Spaces. gr.mount_gradio_app(app, build_ui(), path="/", ssr_mode=False) return app if pipeline.IS_ZERO_GPU: # Repli ZeroGPU (cf. docstring) : demo.launch() natif, pattern garanti # compatible par tous les exemples officiels, sans le endpoint WebSocket. if __name__ == "__main__": demo = build_ui() demo.launch(server_name="0.0.0.0", server_port=int(os.environ.get("PORT", 7860))) else: # Architecture de reference (dev). `app` doit exister au niveau module # pour `uvicorn app:app --port 8000` (voir README pod). app = _build_websocket_app() if __name__ == "__main__": # Point d'entree si jamais execute directement (python app.py) hors # ZeroGPU - en dev on utilise plutot `uvicorn app:app --port 8000`. import uvicorn uvicorn.run(app, host="0.0.0.0", port=int(os.environ.get("PORT", 7860)))