Skip to content
39 changes: 31 additions & 8 deletions triton_viz/visualizer/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,31 @@
last_local_port = None


# Server state management
class ServerState:
"""Encapsulates server state to avoid global variables."""

def __init__(self):
self.last_public_url = None
self.last_local_port = None

def set_public_url(self, url):
self.last_public_url = url

def set_local_port(self, port):
self.last_local_port = port

def get_public_url(self):
return self.last_public_url

def get_local_port(self):
return self.last_local_port


# Create a single instance for the module
_server_state = ServerState()


def precompute_c_values(op_data):
input_data = op_data["input_data"]
other_data = op_data["other_data"]
Expand Down Expand Up @@ -254,10 +279,9 @@ def run_flask_with_cloudflared(port: int = 8000, tunnel_port: int | None = None)
cloudflared_port = port
if tunnel_port is None:
tunnel_port = cloudflared_port + 1
global last_public_url, last_local_port
tunnel_url = _run_cloudflared(cloudflared_port, tunnel_port)
last_public_url = tunnel_url
last_local_port = cloudflared_port
_server_state.set_public_url(tunnel_url)
_server_state.set_local_port(cloudflared_port)
print(f"Cloudflare tunnel URL: {tunnel_url}")
app.run(host="0.0.0.0", port=cloudflared_port, debug=False, use_reloader=False)

Expand All @@ -283,7 +307,7 @@ def launch(share: bool = True, port: int | None = None):

# Try to get the tunnel URL by making a request to the local server
local_url = f"http://localhost:{actual_port}"
public_url = last_public_url
public_url = _server_state.get_public_url()

try:
# touch local server to ensure it's up
Expand All @@ -304,8 +328,7 @@ def launch(share: bool = True, port: int | None = None):
local_url = f"http://localhost:{actual_port}"
print(f"Running on local URL: {local_url}")
print("--------")
global last_local_port
last_local_port = actual_port
_server_state.set_local_port(actual_port)

# Run Flask in a background thread so callers can continue (non-blocking)
def _run_local():
Expand All @@ -320,7 +343,7 @@ def _run_local():

def get_last_public_url():
"""Return the last Cloudflare public URL created by launch(share=True)."""
return last_public_url
return _server_state.get_public_url()


@app.route("/shutdown", methods=["POST", "GET"])
Expand All @@ -342,7 +365,7 @@ def stop_server(port: int | None = None):
Stop the running Flask server by calling the /shutdown endpoint.
If port is None, it will try the last used local port.
"""
target_port = port or last_local_port
target_port = port or _server_state.get_local_port()
if target_port is None:
return False
try:
Expand Down