vmtk-centerline-api / server.py
mosta02's picture
Upload 6 files
a1b81e3 verified
Raw
History Blame Contribute Delete
4.72 kB
import os
import sys
import gc
import base64
import tempfile
import traceback
from flask import Flask, request, jsonify
from flask_cors import CORS
# Import the logic from the backend modules
from centerline_extraction import compute_centerlines_network, load_nrrd_as_vtk_image
from ostium_detection import extract_centerline_data
from endograft_generation import build_endograft_stl_payload
# Startup diagnostics (prints to Docker logs)
def _print_env_info():
import vtk
print(f"Python : {sys.version}", flush=True)
print(f"VTK : {vtk.vtkVersion.GetVTKVersion()}", flush=True)
try:
from vmtk import vtkvmtkMiscPython as m
f = m.vtkvmtkPolyDataNetworkExtraction()
assert hasattr(f, 'SetInputData')
print("vmtk: SetInputData OK", flush=True)
except Exception as e:
print(f"vmtk: FAILED - {e}", flush=True)
_print_env_info()
# Flask API server that wraps the VMTK centerline extraction, ostium detection, and endograft generation pipelines
app = Flask(__name__)
CORS(app) # Allow cross-origin requests from the Next.js frontend
app.config["MAX_CONTENT_LENGTH"] = 200 * 1024 * 1024 # Maximum file upload size (200 MB)
# Simple health check so the deployment platform knows the server is running
@app.route("/health", methods=["GET"])
def health():
return jsonify({"status": "ok"})
# Runs the full VMTK centerline extraction pipeline on NRRD segmentation and returns the result as JSON
@app.route("/extract", methods=["POST"])
def extract():
# Check that the segmentation file sent to the centerline extraction endpoint exists and is an NRRD file
if "file" not in request.files:
return jsonify({"status": "error", "error": "No file provided"}), 400
uploaded = request.files["file"]
file_name = uploaded.filename.lower()
if not file_name.endswith(".nrrd"):
return jsonify({"status": "error", "error": "Invalid File Type. Allowed: .nrrd"}), 400
# Save to a temp file so VMTK can read it from disk
tmp = tempfile.NamedTemporaryFile(suffix=".nrrd", delete=False)
try:
uploaded.save(tmp.name)
tmp.close()
# Read the NRRD segmentation via SimpleITK (preserves full direction matrix)
image, origin, direction = load_nrrd_as_vtk_image(tmp.name)
# Slicer-style surface processing -> Voronoi centerline extraction
centerlines, branch_junctions = compute_centerlines_network(image, origin, direction)
# Extract labelled segments and artery bifurcation data
data = extract_centerline_data(centerlines, branch_junctions=branch_junctions)
# Free the large VTK pipeline objects immediately
del centerlines, branch_junctions
gc.collect()
# Include the NRRD image's physical extent so the frontend can normalize the centerline
dims = image.GetDimensions()
spacing = image.GetSpacing()
phys_size = [dims[i] * spacing[i] for i in range(3)]
phys_center = [origin[i] + phys_size[i] / 2.0 for i in range(3)]
data['volume_info'] = {'center': {'x': phys_center[0], 'y': phys_center[1], 'z': phys_center[2]}, 'size': {'x': phys_size[0], 'y': phys_size[1], 'z': phys_size[2]}}
del image
gc.collect()
return jsonify({"status": "success", "data": data})
except TimeoutError as e:
return jsonify({"status": "error", "error": str(e)}), 504
except Exception as e:
traceback.print_exc()
return jsonify({"status": "error", "error": str(e)}), 500
finally:
# Always clean up the temp file
try:
os.unlink(tmp.name)
except OSError:
pass
# Runs the endograft generation pipeline on VMTK centerline and returns the result as JSON
@app.route("/endograft", methods=["POST"])
def endograft():
body = request.get_json(silent=True) or {}
centerline = body.get("centerline")
params = body.get("params") or {}
if not isinstance(centerline, list) or len(centerline) < 2:
return jsonify({"status": "error", "error": "centerline must be a list of at least 2 points"}), 400
try:
payload = build_endograft_stl_payload(centerline, **params)
result = {
"endograft_stl_b64": base64.b64encode(payload["endograft"]).decode("utf-8"),
"struts_stl_b64": base64.b64encode(payload["struts"]).decode("utf-8"),
"combined_stl_b64": base64.b64encode(payload["combined"]).decode("utf-8"),
}
gc.collect()
return jsonify({"status": "success", "data": result})
except Exception as e:
traceback.print_exc()
return jsonify({"status": "error", "error": str(e)}), 500
if __name__ == "__main__":
port = int(os.environ.get("PORT", 5050))
app.run(host="0.0.0.0", port=port)