"""Browser-based 2D viewer: drag-to-apply force on mouse release."""
import io
import json
import mimetypes
import os
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import parse_qs, urlparse
import numpy as np
import torch
from PIL import Image
from utils.interactive_scenes import scene_public_view
_FORCE_MIN = 1.0
_FORCE_MAX = 10.0
_INDEX_HTML = """
EndoGSim Interactive / 交互仿真
"""
def tensor_to_rgb(image_chw):
"""Convert (3, H, W) float tensor in [0, 1] to RGB uint8."""
img = image_chw.detach().clamp(0, 1).permute(1, 2, 0).cpu().numpy()
return (img * 255).astype(np.uint8)
def pixel_to_world_ray(camera, px, py):
"""Return (origin, direction) for a screen pixel. Both are CUDA float tensors."""
fx, fy = camera.Fx, camera.Fy
cx, cy = camera.Cx, camera.Cy
ray_cam = torch.tensor(
[(px - cx) / fx, (py - cy) / fy, 1.0],
dtype=torch.float32,
device="cuda",
)
ray_cam = ray_cam / torch.linalg.norm(ray_cam)
c2w = torch.linalg.inv(camera.world_view_transform)
rot = c2w[:3, :3]
origin = camera.camera_center
direction = rot @ ray_cam
direction = direction / torch.linalg.norm(direction)
return origin, direction
def pick_world_point_on_ray(world_positions, origin, direction, max_dist=0.15):
"""Pick world-space point on ray closest to simulated particles."""
oc = world_positions - origin.unsqueeze(0)
t = torch.sum(oc * direction.unsqueeze(0), dim=1)
t = torch.clamp(t, min=0.0)
closest = origin.unsqueeze(0) + t.unsqueeze(1) * direction.unsqueeze(0)
dist = torch.linalg.norm(world_positions - closest, dim=1)
valid = t > 0
if not torch.any(valid):
return None, None, None
dist = torch.where(valid, dist, torch.full_like(dist, float("inf")))
idx = int(torch.argmin(dist).item())
if float(dist[idx].item()) > max_dist:
return None, None, None
pick_t = float(t[idx].item())
return world_positions[idx].detach().cpu().numpy(), idx, pick_t
def compute_drag_world_force(cam_width, cam_height, camera, x0, y0, x1, y1, pick_t, force_magnitude):
"""Map screen drag to a world-space force along the drag direction."""
px0 = int(np.clip(x0, 0.0, 1.0) * (cam_width - 1))
py0 = int(np.clip(y0, 0.0, 1.0) * (cam_height - 1))
px1 = int(np.clip(x1, 0.0, 1.0) * (cam_width - 1))
py1 = int(np.clip(y1, 0.0, 1.0) * (cam_height - 1))
drag_screen = np.hypot(x1 - x0, y1 - y0)
if drag_screen < 0.01:
return None, "Drag too short / 拖动太短"
origin, dir0 = pixel_to_world_ray(camera, px0, py0)
_, dir1 = pixel_to_world_ray(camera, px1, py1)
p0 = origin + pick_t * dir0
p1 = origin + pick_t * dir1
drag_world = (p1 - p0).detach().cpu().numpy()
drag_norm = float(np.linalg.norm(drag_world))
if drag_norm < 1e-8:
return None, "Drag has no direction / 拖动方向无效"
direction = drag_world / drag_norm
magnitude = force_magnitude * drag_screen * 3.0
return (direction * magnitude).tolist(), None
class _ViewerHTTPHandler(BaseHTTPRequestHandler):
viewer = None
def log_message(self, format, *args):
return
def _send_json(self, payload: dict) -> None:
data = json.dumps(payload).encode("utf-8")
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(data)))
self.end_headers()
self.wfile.write(data)
def _send_bytes(self, data: bytes, content_type: str) -> None:
self.send_response(200)
self.send_header("Content-Type", content_type)
self.send_header("Content-Length", str(len(data)))
self.send_header("Cache-Control", "no-store")
self.end_headers()
self.wfile.write(data)
def _read_json_body(self):
length = int(self.headers.get("Content-Length", 0))
if length <= 0:
return {}
raw = self.rfile.read(length)
return json.loads(raw.decode("utf-8"))
def do_GET(self):
path = urlparse(self.path).path
viewer = self.viewer
if path in ("/", "/index.html"):
self._send_bytes(_INDEX_HTML.encode("utf-8"), "text/html; charset=utf-8")
elif path == "/frame.jpg":
with viewer._lock:
data = viewer._jpeg_bytes
if data is None:
self.send_error(503)
return
self._send_bytes(data, "image/jpeg")
elif path == "/api/status":
self._send_json(viewer.status_dict())
elif path == "/api/scenes":
self._send_json(viewer.scenes_dict())
elif path == "/api/scene_preview":
query = parse_qs(urlparse(self.path).query)
scene_id = (query.get("id") or [None])[0]
preview = viewer.get_scene_preview_bytes(scene_id)
if preview is None:
self.send_error(404)
return
data, content_type = preview
self._send_bytes(data, content_type)
else:
self.send_error(404)
def do_POST(self):
viewer = self.viewer
path = urlparse(self.path).path
try:
body = self._read_json_body()
except json.JSONDecodeError:
self.send_error(400)
return
if path == "/api/drag_release":
message = viewer.queue_drag_release(
float(body.get("x0", 0.5)),
float(body.get("y0", 0.5)),
float(body.get("x1", 0.5)),
float(body.get("y1", 0.5)),
)
self._send_json({"ok": True, "message": message})
elif path == "/api/pause":
viewer.paused = bool(body.get("paused", False))
self._send_json({"ok": True})
elif path == "/api/reset":
viewer.queue_reset()
self._send_json({"ok": True, "message": "Reset queued / 重置已排队"})
elif path == "/api/force":
viewer.set_force_magnitude(float(body.get("magnitude", viewer.force_magnitude)))
self._send_json({"ok": True})
elif path == "/api/select_scene":
scene_id = body.get("scene_id")
if not scene_id:
self._send_json({"ok": False, "error": "Missing scene_id"})
return
ok, message = viewer.queue_select_scene(str(scene_id))
self._send_json({"ok": ok, "message": message})
else:
self.send_error(404)
class InteractiveViewer:
"""Browser-based 2D image viewer (no 3D viewport)."""
def __init__(
self,
session,
session_factory,
scenes,
current_scene_id,
force_magnitude=1.0,
pick_radius=0.15,
impulse_radius=0.05,
host="0.0.0.0",
port=8080,
target_fps=15,
downsample=0.5,
white_bg=False,
):
self.session = session
self.session_factory = session_factory
self.scenes = scenes
self.scene_by_id = {s["id"]: s for s in scenes}
self.current_scene_id = current_scene_id
self.downsample = downsample
self.white_bg = white_bg
self.force_magnitude = self._clamp_force(force_magnitude)
self.pick_radius = pick_radius
self.impulse_radius = impulse_radius
self.host = host
self.port = port
self.target_fps = target_fps
self.paused = False
self.running = True
self.status = "Ready / 就绪"
self._loading_scene = False
self._lock = threading.Lock()
self._action_lock = threading.Lock()
self._pending_drag = None
self._pending_scene_id = None
self._burst_substeps = 0
self._reset_pending = False
self._jpeg_bytes = None
self._httpd = None
self._preview_cache = {}
self._frame_version = 0
def _clamp_force(self, value: float) -> float:
return float(np.clip(value, _FORCE_MIN, _FORCE_MAX))
def set_force_magnitude(self, value: float) -> None:
self.force_magnitude = self._clamp_force(value)
def queue_reset(self) -> None:
self._reset_pending = True
def queue_select_scene(self, scene_id: str) -> tuple[bool, str]:
if scene_id not in self.scene_by_id:
return False, f"Unknown scene / 未知场景: {scene_id}"
if scene_id == self.current_scene_id:
return False, "Already active / 已是当前场景"
with self._action_lock:
self._pending_scene_id = scene_id
return True, "Scene switch queued / 场景切换已排队"
def get_scene_preview_bytes(self, scene_id: str):
if not scene_id or scene_id not in self.scene_by_id:
return None
if scene_id in self._preview_cache:
return self._preview_cache[scene_id]
scene = self.scene_by_id[scene_id]
preview_path = scene.get("preview")
if not preview_path or not os.path.isfile(preview_path):
return None
mime, _ = mimetypes.guess_type(preview_path)
if not mime:
mime = "image/png"
with open(preview_path, "rb") as f:
data = f.read()
self._preview_cache[scene_id] = (data, mime)
return self._preview_cache[scene_id]
def scenes_dict(self):
return {
"current_scene_id": self.current_scene_id,
"scenes": [scene_public_view(s) for s in self.scenes],
}
def _reload_scene(self, scene_id: str) -> None:
scene = self.scene_by_id[scene_id]
self._loading_scene = True
self.status = f"Loading {scene['name']} / 加载中…"
print(f"Loading scene: {scene_id} ({scene['name']})")
try:
self.session = None
torch.cuda.empty_cache()
self.session = self.session_factory(scene)
self.current_scene_id = scene_id
self._burst_substeps = 0
self.status = f"Loaded: {scene['name']}"
print(f"Scene ready: {scene_id}")
except Exception as exc:
self.status = f"Scene load failed / 场景加载失败: {exc}"
print(f"Scene load failed: {scene_id}: {exc}")
finally:
self._loading_scene = False
def queue_drag_release(self, x0: float, y0: float, x1: float, y1: float) -> str:
with self._action_lock:
self._pending_drag = (x0, y0, x1, y1)
return "Drag queued / 拖动已排队"
def _process_drag_release(self, x0: float, y0: float, x1: float, y1: float) -> None:
camera = self.session.current_camera
w = self.session.cam_width
h = self.session.cam_height
px0 = int(np.clip(x0, 0.0, 1.0) * (w - 1))
py0 = int(np.clip(y0, 0.0, 1.0) * (h - 1))
origin, direction = pixel_to_world_ray(camera, px0, py0)
world_pos = self.session.get_world_positions()
picked, _, pick_t = pick_world_point_on_ray(
world_pos, origin, direction, max_dist=self.pick_radius
)
if picked is None:
self.status = f"Missed model / 未命中模型 ({px0}, {py0})"
return
force, err = compute_drag_world_force(
w, h, camera, x0, y0, x1, y1, pick_t, self.force_magnitude
)
if force is None:
self.status = err or "Invalid drag / 无效拖动"
return
self.session.apply_impulse_at_world(picked, force, radius=self.impulse_radius)
self._burst_substeps = self.session.substeps_per_frame
fx, fy, fz = force
self.status = (
f"Force ({fx:.2f}, {fy:.2f}, {fz:.2f}) / 施力 "
f"at ({picked[0]:.2f}, {picked[1]:.2f}, {picked[2]:.2f})"
)
def _process_pending(self) -> None:
scene_id = None
with self._action_lock:
if self._pending_scene_id is not None:
scene_id = self._pending_scene_id
self._pending_scene_id = None
self._pending_drag = None
self._reset_pending = False
if scene_id is not None:
self._reload_scene(scene_id)
return
if self._reset_pending:
self._reset_pending = False
self.session.reset_simulation()
self.status = "Simulation reset / 仿真已重置"
return
drag = None
with self._action_lock:
if self._pending_drag is not None:
drag = self._pending_drag
self._pending_drag = None
if drag is not None:
self._process_drag_release(*drag)
def status_dict(self):
scene = self.scene_by_id.get(self.current_scene_id, {})
return {
"status": self.status,
"paused": self.paused,
"loading": self._loading_scene,
"current_scene_id": self.current_scene_id,
"current_scene_name": scene.get("name", ""),
"force_magnitude": self.force_magnitude,
"force_min": _FORCE_MIN,
"force_max": _FORCE_MAX,
"frame_version": self._frame_version,
"time": float(self.session.mpm_solver.time) if self.session else 0.0,
}
def _encode_jpeg(self, rgb: np.ndarray) -> bytes:
buf = io.BytesIO()
Image.fromarray(rgb).save(buf, format="JPEG", quality=85)
return buf.getvalue()
def _start_http_server(self) -> None:
_ViewerHTTPHandler.viewer = self
self._httpd = ThreadingHTTPServer((self.host, self.port), _ViewerHTTPHandler)
thread = threading.Thread(target=self._httpd.serve_forever, daemon=True)
thread.start()
def run(self):
self._start_http_server()
display_host = "localhost" if self.host in ("0.0.0.0", "::") else self.host
print(f"2D viewer running at http://{display_host}:{self.port}")
print(" If remote: ssh -L 8080:localhost:8080 user@server")
print(f" Then open http://localhost:{self.port} in your browser")
print(" Drag & release to push / 拖动后松开施力")
frame_interval = 1.0 / max(self.target_fps, 1.0)
while self.running:
t0 = time.time()
self._process_pending()
if self._loading_scene or self.session is None:
time.sleep(0.05)
continue
n_substeps = 0
if self._burst_substeps > 0:
n_substeps = self._burst_substeps
self._burst_substeps = 0
elif not self.paused:
n_substeps = self.session.substeps_per_frame
if n_substeps > 0:
self.session.step(n_substeps=n_substeps)
rgb = tensor_to_rgb(self.session.render())
jpeg = self._encode_jpeg(rgb)
with self._lock:
self._jpeg_bytes = jpeg
self._frame_version += 1
if not self.status.startswith(
("Force", "Missed", "Drag", "Simulation", "Invalid", "Loading", "Loaded",
"施力", "未命中", "仿真", "无效", "太短", "加载")
):
self.status = f"t={self.session.mpm_solver.time:.4f}s"
elapsed = time.time() - t0
time.sleep(max(0.0, frame_interval - elapsed))