chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,250 @@
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
from threading import RLock
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import requests
|
||||
|
||||
from ray._common.network_utils import build_address
|
||||
from ray.autoscaler.node_launch_exception import NodeLaunchException
|
||||
from ray.autoscaler.node_provider import NodeProvider
|
||||
from ray.autoscaler.tags import (
|
||||
NODE_KIND_HEAD,
|
||||
NODE_KIND_WORKER,
|
||||
STATUS_SETTING_UP,
|
||||
STATUS_UP_TO_DATE,
|
||||
TAG_RAY_NODE_KIND,
|
||||
TAG_RAY_NODE_NAME,
|
||||
TAG_RAY_NODE_STATUS,
|
||||
TAG_RAY_USER_NODE_TYPE,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
HEAD_NODE_ID = 0
|
||||
HEAD_NODE_TYPE = "ray.head.default"
|
||||
|
||||
|
||||
class SparkNodeProvider(NodeProvider):
|
||||
"""A node provider that implements provider for nodes of Ray on spark."""
|
||||
|
||||
def __init__(self, provider_config, cluster_name):
|
||||
NodeProvider.__init__(self, provider_config, cluster_name)
|
||||
self.lock = RLock()
|
||||
|
||||
self._nodes = {
|
||||
str(HEAD_NODE_ID): {
|
||||
"tags": {
|
||||
TAG_RAY_NODE_KIND: NODE_KIND_HEAD,
|
||||
TAG_RAY_USER_NODE_TYPE: HEAD_NODE_TYPE,
|
||||
TAG_RAY_NODE_NAME: HEAD_NODE_ID,
|
||||
TAG_RAY_NODE_STATUS: STATUS_UP_TO_DATE,
|
||||
}
|
||||
},
|
||||
}
|
||||
self._next_node_id = 0
|
||||
|
||||
self.ray_head_ip = self.provider_config["ray_head_ip"]
|
||||
# The port of spark job server. We send http request to spark job server
|
||||
# to launch spark jobs, ray worker nodes are launched by spark task in
|
||||
# spark jobs.
|
||||
spark_job_server_port = self.provider_config["spark_job_server_port"]
|
||||
self.spark_job_server_url = (
|
||||
f"http://{build_address(self.ray_head_ip, spark_job_server_port)}"
|
||||
)
|
||||
self.ray_head_port = self.provider_config["ray_head_port"]
|
||||
# The unique id for the Ray on spark cluster.
|
||||
self.cluster_id = self.provider_config["cluster_unique_id"]
|
||||
|
||||
def get_next_node_id(self):
|
||||
with self.lock:
|
||||
self._next_node_id += 1
|
||||
return self._next_node_id
|
||||
|
||||
def non_terminated_nodes(self, tag_filters):
|
||||
with self.lock:
|
||||
nodes = []
|
||||
|
||||
died_nodes = []
|
||||
for node_id in self._nodes:
|
||||
if node_id == str(HEAD_NODE_ID):
|
||||
status = "running"
|
||||
else:
|
||||
status = self._query_node_status(node_id)
|
||||
|
||||
if status == "running":
|
||||
if (
|
||||
self._nodes[node_id]["tags"][TAG_RAY_NODE_STATUS]
|
||||
== STATUS_SETTING_UP
|
||||
):
|
||||
self._nodes[node_id]["tags"][
|
||||
TAG_RAY_NODE_STATUS
|
||||
] = STATUS_UP_TO_DATE
|
||||
logger.info(
|
||||
f"Spark node provider node {node_id} starts running."
|
||||
)
|
||||
|
||||
if status == "terminated":
|
||||
died_nodes.append(node_id)
|
||||
else:
|
||||
tags = self.node_tags(node_id)
|
||||
ok = True
|
||||
for k, v in tag_filters.items():
|
||||
if tags.get(k) != v:
|
||||
ok = False
|
||||
if ok:
|
||||
nodes.append(node_id)
|
||||
|
||||
for died_node_id in died_nodes:
|
||||
self._nodes.pop(died_node_id)
|
||||
|
||||
return nodes
|
||||
|
||||
def _query_node_status(self, node_id):
|
||||
spark_job_group_id = self._gen_spark_job_group_id(node_id)
|
||||
|
||||
response = requests.post(
|
||||
url=self.spark_job_server_url + "/query_task_status",
|
||||
json={"spark_job_group_id": spark_job_group_id},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
decoded_resp = response.content.decode("utf-8")
|
||||
json_res = json.loads(decoded_resp)
|
||||
return json_res["status"]
|
||||
|
||||
def is_running(self, node_id):
|
||||
with self.lock:
|
||||
return (
|
||||
node_id in self._nodes
|
||||
and self._nodes[node_id]["tags"][TAG_RAY_NODE_STATUS]
|
||||
== STATUS_UP_TO_DATE
|
||||
)
|
||||
|
||||
def is_terminated(self, node_id):
|
||||
with self.lock:
|
||||
return node_id not in self._nodes
|
||||
|
||||
def node_tags(self, node_id):
|
||||
with self.lock:
|
||||
return self._nodes[node_id]["tags"]
|
||||
|
||||
def _get_ip(self, node_id: str) -> Optional[str]:
|
||||
return node_id
|
||||
|
||||
def external_ip(self, node_id):
|
||||
return self._get_ip(node_id)
|
||||
|
||||
def internal_ip(self, node_id):
|
||||
return self._get_ip(node_id)
|
||||
|
||||
def set_node_tags(self, node_id, tags):
|
||||
assert node_id in self._nodes
|
||||
self._nodes[node_id]["tags"].update(tags)
|
||||
|
||||
def create_node(
|
||||
self, node_config: Dict[str, Any], tags: Dict[str, str], count: int
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
raise AssertionError("This method should not be called.")
|
||||
|
||||
def _gen_spark_job_group_id(self, node_id):
|
||||
return (
|
||||
f"ray-cluster-{self.ray_head_port}-{self.cluster_id}"
|
||||
f"-worker-node-{node_id}"
|
||||
)
|
||||
|
||||
def create_node_with_resources_and_labels(
|
||||
self, node_config, tags, count, resources, labels
|
||||
):
|
||||
for _ in range(count):
|
||||
self._create_node_with_resources_and_labels(
|
||||
node_config, tags, resources, labels
|
||||
)
|
||||
|
||||
def _create_node_with_resources_and_labels(
|
||||
self, node_config, tags, resources, labels
|
||||
):
|
||||
from ray.util.spark.cluster_init import _append_resources_config
|
||||
|
||||
with self.lock:
|
||||
resources = resources.copy()
|
||||
node_type = tags[TAG_RAY_USER_NODE_TYPE]
|
||||
# NOTE:
|
||||
# "NODE_ID_AS_RESOURCE" value must be an integer,
|
||||
# but `node_id` used by autoscaler must be a string.
|
||||
node_id = str(self.get_next_node_id())
|
||||
resources["NODE_ID_AS_RESOURCE"] = int(node_id)
|
||||
|
||||
conf = self.provider_config.copy()
|
||||
|
||||
num_cpus_per_node = resources.pop("CPU")
|
||||
num_gpus_per_node = resources.pop("GPU")
|
||||
heap_memory_per_node = resources.pop("memory")
|
||||
object_store_memory_per_node = resources.pop("object_store_memory")
|
||||
|
||||
conf["worker_node_options"] = _append_resources_config(
|
||||
conf["worker_node_options"], resources
|
||||
)
|
||||
response = requests.post(
|
||||
url=self.spark_job_server_url + "/create_node",
|
||||
json={
|
||||
"spark_job_group_id": self._gen_spark_job_group_id(node_id),
|
||||
"spark_job_group_desc": (
|
||||
"This job group is for spark job which runs the Ray "
|
||||
f"cluster worker node {node_id} connecting to ray "
|
||||
f"head node {build_address(self.ray_head_ip, self.ray_head_port)}"
|
||||
),
|
||||
"using_stage_scheduling": conf["using_stage_scheduling"],
|
||||
"ray_head_ip": self.ray_head_ip,
|
||||
"ray_head_port": self.ray_head_port,
|
||||
"ray_temp_dir": conf["ray_temp_dir"],
|
||||
"num_cpus_per_node": num_cpus_per_node,
|
||||
"num_gpus_per_node": num_gpus_per_node,
|
||||
"heap_memory_per_node": heap_memory_per_node,
|
||||
"object_store_memory_per_node": object_store_memory_per_node,
|
||||
"worker_node_options": conf["worker_node_options"],
|
||||
"collect_log_to_path": conf["collect_log_to_path"],
|
||||
"node_id": resources["NODE_ID_AS_RESOURCE"],
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
# Spark job server is locally launched, if spark job server request
|
||||
# failed, it is unlikely network error but probably unrecoverable
|
||||
# error, so we make it fast-fail.
|
||||
response.raise_for_status()
|
||||
except Exception:
|
||||
raise NodeLaunchException(
|
||||
"Node creation failure",
|
||||
f"Starting ray worker node {node_id} failed",
|
||||
sys.exc_info(),
|
||||
)
|
||||
|
||||
self._nodes[node_id] = {
|
||||
"tags": {
|
||||
TAG_RAY_NODE_KIND: NODE_KIND_WORKER,
|
||||
TAG_RAY_USER_NODE_TYPE: node_type,
|
||||
TAG_RAY_NODE_NAME: node_id,
|
||||
TAG_RAY_NODE_STATUS: STATUS_SETTING_UP,
|
||||
},
|
||||
}
|
||||
logger.info(f"Spark node provider creates node {node_id}.")
|
||||
|
||||
def terminate_node(self, node_id):
|
||||
if node_id in self._nodes:
|
||||
response = requests.post(
|
||||
url=self.spark_job_server_url + "/terminate_node",
|
||||
json={"spark_job_group_id": self._gen_spark_job_group_id(node_id)},
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
with self.lock:
|
||||
if node_id in self._nodes:
|
||||
self._nodes.pop(node_id)
|
||||
|
||||
logger.info(f"Spark node provider terminates node {node_id}")
|
||||
|
||||
@staticmethod
|
||||
def bootstrap_config(cluster_config):
|
||||
return cluster_config
|
||||
@@ -0,0 +1,244 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
|
||||
from pyspark.util import inheritable_thread_target
|
||||
|
||||
from ray.util.spark.cluster_init import _start_ray_worker_nodes
|
||||
|
||||
|
||||
class SparkJobServerRequestHandler(BaseHTTPRequestHandler):
|
||||
def setup(self) -> None:
|
||||
super().setup()
|
||||
self._handler_lock = threading.RLock()
|
||||
self._created_node_id_set = set()
|
||||
self._logger = logging.getLogger(__name__)
|
||||
if "RAY_ON_SPARK_JOB_SERVER_VERBOSE" in os.environ:
|
||||
self._logger.setLevel(logging.DEBUG)
|
||||
else:
|
||||
self._logger.setLevel(logging.WARN)
|
||||
|
||||
def _set_headers(self):
|
||||
self.send_response(200)
|
||||
self.send_header("Content-type", "application/json")
|
||||
self.end_headers()
|
||||
|
||||
def handle_POST(self, path, data):
|
||||
path_parts = Path(path).parts[1:]
|
||||
|
||||
spark_job_group_id = data["spark_job_group_id"]
|
||||
|
||||
if path_parts[0] == "create_node":
|
||||
assert len(path_parts) == 1, f"Illegal request path: {path}"
|
||||
spark_job_group_desc = data["spark_job_group_desc"]
|
||||
using_stage_scheduling = data["using_stage_scheduling"]
|
||||
ray_head_ip = data["ray_head_ip"]
|
||||
ray_head_port = data["ray_head_port"]
|
||||
ray_temp_dir = data["ray_temp_dir"]
|
||||
num_cpus_per_node = data["num_cpus_per_node"]
|
||||
num_gpus_per_node = data["num_gpus_per_node"]
|
||||
heap_memory_per_node = data["heap_memory_per_node"]
|
||||
object_store_memory_per_node = data["object_store_memory_per_node"]
|
||||
worker_node_options = data["worker_node_options"]
|
||||
collect_log_to_path = data["collect_log_to_path"]
|
||||
node_id = data["node_id"]
|
||||
self._created_node_id_set.add(node_id)
|
||||
|
||||
def start_ray_worker_thread_fn():
|
||||
try:
|
||||
err_msg = _start_ray_worker_nodes(
|
||||
spark_job_server=self.server,
|
||||
spark_job_group_id=spark_job_group_id,
|
||||
spark_job_group_desc=spark_job_group_desc,
|
||||
num_worker_nodes=1,
|
||||
using_stage_scheduling=using_stage_scheduling,
|
||||
ray_head_ip=ray_head_ip,
|
||||
ray_head_port=ray_head_port,
|
||||
ray_temp_dir=ray_temp_dir,
|
||||
num_cpus_per_node=num_cpus_per_node,
|
||||
num_gpus_per_node=num_gpus_per_node,
|
||||
heap_memory_per_node=heap_memory_per_node,
|
||||
object_store_memory_per_node=object_store_memory_per_node,
|
||||
worker_node_options=worker_node_options,
|
||||
collect_log_to_path=collect_log_to_path,
|
||||
node_id=node_id,
|
||||
)
|
||||
if err_msg:
|
||||
self._logger.warning(
|
||||
f"Spark job {spark_job_group_id} hosting Ray worker node "
|
||||
f"launching failed, error:\n{err_msg}"
|
||||
)
|
||||
except Exception:
|
||||
if spark_job_group_id in self.server.task_status_dict:
|
||||
self.server.task_status_dict.pop(spark_job_group_id)
|
||||
|
||||
msg = (
|
||||
f"Spark job {spark_job_group_id} hosting Ray worker node exit."
|
||||
)
|
||||
if self._logger.level > logging.DEBUG:
|
||||
self._logger.warning(
|
||||
f"{msg} To see details, you can set "
|
||||
"'RAY_ON_SPARK_JOB_SERVER_VERBOSE' environmental variable "
|
||||
"to '1' before calling 'ray.util.spark.setup_ray_cluster'."
|
||||
)
|
||||
else:
|
||||
# This branch is only for debugging Ray-on-Spark purpose.
|
||||
# User can configure 'RAY_ON_SPARK_JOB_SERVER_VERBOSE'
|
||||
# environment variable to make the spark job server logging
|
||||
# showing full exception stack here.
|
||||
self._logger.debug(msg, exc_info=True)
|
||||
|
||||
threading.Thread(
|
||||
target=inheritable_thread_target(start_ray_worker_thread_fn),
|
||||
args=(),
|
||||
daemon=True,
|
||||
).start()
|
||||
|
||||
self.server.task_status_dict[spark_job_group_id] = "pending"
|
||||
return {}
|
||||
|
||||
elif path_parts[0] == "check_node_id_availability":
|
||||
node_id = data["node_id"]
|
||||
with self._handler_lock:
|
||||
if node_id in self._created_node_id_set:
|
||||
# If the node with the node id has been created,
|
||||
# it shouldn't be created twice so fail fast here.
|
||||
# The case happens when a Ray node is down unexpected
|
||||
# caused by spark worker node down and spark tries to
|
||||
# reschedule the spark task, so it triggers node
|
||||
# creation with duplicated node id.
|
||||
return {"available": False}
|
||||
else:
|
||||
self._created_node_id_set.add(node_id)
|
||||
return {"available": True}
|
||||
|
||||
elif path_parts[0] == "terminate_node":
|
||||
assert len(path_parts) == 1, f"Illegal request path: {path}"
|
||||
self.server.spark.sparkContext.cancelJobGroup(spark_job_group_id)
|
||||
if spark_job_group_id in self.server.task_status_dict:
|
||||
self.server.task_status_dict.pop(spark_job_group_id)
|
||||
return {}
|
||||
|
||||
elif path_parts[0] == "notify_task_launched":
|
||||
if spark_job_group_id in self.server.task_status_dict:
|
||||
# Note that if `spark_job_group_id` not in task_status_dict,
|
||||
# the task has been terminated
|
||||
self.server.task_status_dict[spark_job_group_id] = "running"
|
||||
self._logger.info(f"Spark task in {spark_job_group_id} has started.")
|
||||
return {}
|
||||
|
||||
elif path_parts[0] == "query_task_status":
|
||||
if spark_job_group_id in self.server.task_status_dict:
|
||||
return {"status": self.server.task_status_dict[spark_job_group_id]}
|
||||
else:
|
||||
return {"status": "terminated"}
|
||||
|
||||
elif path_parts[0] == "query_last_worker_err":
|
||||
return {"last_worker_err": self.server.last_worker_error}
|
||||
|
||||
else:
|
||||
raise ValueError(f"Illegal request path: {path}")
|
||||
|
||||
def do_POST(self):
|
||||
"""Reads post request body"""
|
||||
self._set_headers()
|
||||
content_len = int(self.headers["content-length"])
|
||||
content_type = self.headers["content-type"]
|
||||
assert content_type == "application/json"
|
||||
path = self.path
|
||||
post_body = self.rfile.read(content_len).decode("utf-8")
|
||||
post_body_json = json.loads(post_body)
|
||||
with self._handler_lock:
|
||||
response_body_json = self.handle_POST(path, post_body_json)
|
||||
response_body = json.dumps(response_body_json)
|
||||
self.wfile.write(response_body.encode("utf-8"))
|
||||
|
||||
def log_request(self, code="-", size="-"):
|
||||
# Make logs less verbose.
|
||||
pass
|
||||
|
||||
|
||||
class SparkJobServer(ThreadingHTTPServer):
|
||||
"""
|
||||
High level design:
|
||||
|
||||
1. In Ray on spark autoscaling mode, How to start and terminate Ray worker node ?
|
||||
|
||||
It uses spark job to launch Ray worker node,
|
||||
and each spark job contains only one spark task, the corresponding spark task
|
||||
creates Ray worker node as subprocess.
|
||||
When autoscaler request terminating specific Ray worker node, it cancels
|
||||
corresponding spark job to trigger Ray worker node termination.
|
||||
Because we can only cancel spark job not spark task when we need to scale
|
||||
down a Ray worker node. So we have to have one spark job for each Ray worker node.
|
||||
|
||||
2. How to create / cancel spark job from spark node provider?
|
||||
|
||||
Spark node provider runs in autoscaler process that is different process
|
||||
than the one that executes "setup_ray_cluster" API. User calls "setup_ray_cluster"
|
||||
API in spark application driver node, and the semantic is "setup_ray_cluster"
|
||||
requests spark resources from this spark application.
|
||||
Internally, "setup_ray_cluster" should use "spark session" instance to request
|
||||
spark application resources. But spark node provider runs in another python
|
||||
process, in order to share spark session to the separate NodeProvider process,
|
||||
it sets up a spark job server that runs inside spark application driver process
|
||||
(the process that calls "setup_ray_cluster" API), and in NodeProvider process,
|
||||
it sends RPC request to the spark job server for creating spark jobs in the
|
||||
spark application.
|
||||
Note that we cannot create another spark session in NodeProvider process,
|
||||
because if doing so, it means we create another spark application, and then
|
||||
it causes NodeProvider requests resources belonging to the new spark application,
|
||||
but we need to ensure all requested spark resources belong to
|
||||
the original spark application that calls "setup_ray_cluster" API.
|
||||
|
||||
Note:
|
||||
The server must inherit ThreadingHTTPServer because request handler uses
|
||||
the active spark session in current process to create spark jobs, so all request
|
||||
handler must be running in current process.
|
||||
"""
|
||||
|
||||
def __init__(self, server_address, spark, ray_node_custom_env):
|
||||
super().__init__(server_address, SparkJobServerRequestHandler)
|
||||
self.spark = spark
|
||||
|
||||
# For ray on spark autoscaling mode,
|
||||
# for each ray worker node, we create an individual spark job
|
||||
# to launch it, the corresponding spark job has only one
|
||||
# spark task that starts ray worker node, and the spark job
|
||||
# is assigned with a unique spark job group ID that is used
|
||||
# to cancel this spark job (i.e., kill corresponding ray worker node).
|
||||
# Each spark task has status of pending, running, or terminated.
|
||||
# the task_status_dict key is spark job group id,
|
||||
# and value is the corresponding spark task status.
|
||||
# each spark task holds a ray worker node.
|
||||
self.task_status_dict = {}
|
||||
self.last_worker_error = None
|
||||
self.ray_node_custom_env = ray_node_custom_env
|
||||
|
||||
def shutdown(self) -> None:
|
||||
super().shutdown()
|
||||
for spark_job_group_id in list(self.task_status_dict.keys()):
|
||||
self.spark.sparkContext.cancelJobGroup(spark_job_group_id)
|
||||
# Sleep 1 second to wait for all spark job cancellation
|
||||
# The spark job cancellation will do things asyncly in a background thread,
|
||||
# On Databricks platform, when detaching a notebook, it triggers SIGTERM
|
||||
# and then sigterm handler triggers Ray cluster shutdown, without sleep,
|
||||
# after the SIGTERM handler execution the process is killed and then
|
||||
# these cancelling spark job background threads are killed.
|
||||
time.sleep(1)
|
||||
|
||||
|
||||
def _start_spark_job_server(host, port, spark, ray_node_custom_env):
|
||||
server = SparkJobServer((host, port), spark, ray_node_custom_env)
|
||||
|
||||
def run_server():
|
||||
server.serve_forever()
|
||||
|
||||
server_thread = threading.Thread(target=run_server, daemon=True)
|
||||
server_thread.start()
|
||||
|
||||
return server
|
||||
Reference in New Issue
Block a user