diff --git a/triton_viz/visualizer/interface.py b/triton_viz/visualizer/interface.py index d01fd8cf..9fa30b5d 100644 --- a/triton_viz/visualizer/interface.py +++ b/triton_viz/visualizer/interface.py @@ -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"] @@ -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) @@ -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 @@ -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(): @@ -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"]) @@ -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: