From defd41aad7067f515badce6e77118ba37d2273f0 Mon Sep 17 00:00:00 2001 From: fyrestone Date: Sat, 10 Oct 2020 13:27:05 +0800 Subject: [PATCH] [Dashboard] http route handler cache (#10921) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Add aiohttp_cache to dashboard * Add comments; Refine code * Keep NODE_STATS_UPDATE_INTERVAL_SECONDS 1 second; Change AIOHTTP_CACHE_TTL_SECONDS to 2 seconds * Update merge Co-authored-by: 刘宝 --- dashboard/consts.py | 4 + dashboard/modules/log/test_log.py | 2 +- .../modules/logical_view/logical_view_head.py | 6 +- dashboard/modules/reporter/reporter_head.py | 10 +- .../stats_collector/stats_collector_head.py | 22 ++-- .../stats_collector/test_stats_collector.py | 5 +- dashboard/modules/test/test_head.py | 21 +++- dashboard/modules/tune/tune_head.py | 12 +- dashboard/tests/conftest.py | 7 ++ dashboard/tests/test_dashboard.py | 80 +++++++++++++- dashboard/utils.py | 103 +++++++++++++++++- 11 files changed, 235 insertions(+), 37 deletions(-) diff --git a/dashboard/consts.py b/dashboard/consts.py index 5b321800a..80c658882 100644 --- a/dashboard/consts.py +++ b/dashboard/consts.py @@ -13,6 +13,10 @@ REDIS_KEY_DASHBOARD_RPC = "dashboard_rpc" REDIS_KEY_GCS_SERVER_ADDRESS = "GcsServerAddress" REPORT_METRICS_TIMEOUT_SECONDS = 2 REPORT_METRICS_INTERVAL_SECONDS = 10 +# aiohttp_cache +AIOHTTP_CACHE_TTL_SECONDS = 2 +AIOHTTP_CACHE_MAX_SIZE = 128 +AIOHTTP_CACHE_DISABLE_ENVIRONMENT_KEY = "RAY_DASHBOARD_NO_CACHE" # Named signals SIGNAL_NODE_INFO_FETCHED = "node_info_fetched" # Default param for RotatingFileHandler diff --git a/dashboard/modules/log/test_log.py b/dashboard/modules/log/test_log.py index f73e9bc0c..c0cece624 100644 --- a/dashboard/modules/log/test_log.py +++ b/dashboard/modules/log/test_log.py @@ -36,7 +36,7 @@ class LogUrlParser(html.parser.HTMLParser): return self._urls -def test_log(ray_start_with_dashboard): +def test_log(disable_aiohttp_cache, ray_start_with_dashboard): @ray.remote def write_log(s): print(s) diff --git a/dashboard/modules/logical_view/logical_view_head.py b/dashboard/modules/logical_view/logical_view_head.py index f8168f319..674ac64e9 100644 --- a/dashboard/modules/logical_view/logical_view_head.py +++ b/dashboard/modules/logical_view/logical_view_head.py @@ -24,7 +24,7 @@ class LogicalViewHead(dashboard_utils.DashboardHeadModule): # hence we merge them together before constructing actor groups. actors.update(actor_creation_tasks) actor_groups = actor_utils.construct_actor_groups(actors) - return await rest_response( + return rest_response( success=True, message="Fetched actor groups.", actor_groups=actor_groups) @@ -36,7 +36,7 @@ class LogicalViewHead(dashboard_utils.DashboardHeadModule): ip_address = req.query["ipAddress"] port = req.query["port"] except KeyError: - return await rest_response(success=False, message="Bad Request") + return rest_response(success=False, message="Bad Request") try: channel = aiogrpc.insecure_channel(f"{ip_address}:{port}") stub = core_worker_pb2_grpc.CoreWorkerServiceStub(channel) @@ -51,7 +51,7 @@ class LogicalViewHead(dashboard_utils.DashboardHeadModule): # before this handler, however it deletes the actor correctly. pass - return await rest_response( + return rest_response( success=True, message=f"Killed actor with id {actor_id}") async def run(self, server): diff --git a/dashboard/modules/reporter/reporter_head.py b/dashboard/modules/reporter/reporter_head.py index f0ff25dd8..b9ce287bb 100644 --- a/dashboard/modules/reporter/reporter_head.py +++ b/dashboard/modules/reporter/reporter_head.py @@ -48,7 +48,7 @@ class ReportHead(dashboard_utils.DashboardHeadModule): reporter_pb2.GetProfilingStatsRequest(pid=pid, duration=duration)) profiling_info = (json.loads(reply.profiling_stats) if reply.profiling_stats else reply.std_out) - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Profiling success.", profiling_info=profiling_info) @@ -61,14 +61,14 @@ class ReportHead(dashboard_utils.DashboardHeadModule): with open(config_path) as f: cfg = yaml.safe_load(f) except yaml.YAMLError: - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=False, message=f"No config found at {config_path}.", ) except FileNotFoundError: - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=False, - message=f"Invalid config, could not load YAML.") + message="Invalid config, could not load YAML.") payload = { "min_workers": cfg["min_workers"], @@ -90,7 +90,7 @@ class ReportHead(dashboard_utils.DashboardHeadModule): self._ray_config = payload - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Fetched ray config.", **self._ray_config, diff --git a/dashboard/modules/stats_collector/stats_collector_head.py b/dashboard/modules/stats_collector/stats_collector_head.py index 23f0ebde7..38e45933e 100644 --- a/dashboard/modules/stats_collector/stats_collector_head.py +++ b/dashboard/modules/stats_collector/stats_collector_head.py @@ -64,17 +64,18 @@ class StatsCollector(dashboard_utils.DashboardHeadModule): self._stubs[node_id] = stub @routes.get("/nodes") + @dashboard_utils.aiohttp_cache async def get_all_nodes(self, req) -> aiohttp.web.Response: view = req.query.get("view") if view == "summary": all_node_summary = await DataOrganizer.get_all_node_summary() - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Node summary fetched.", summary=all_node_summary) elif view == "details": all_node_details = await DataOrganizer.get_all_node_details() - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="All node details fetched", clients=all_node_details, @@ -84,19 +85,20 @@ class StatsCollector(dashboard_utils.DashboardHeadModule): for node in DataSource.nodes.values(): if node["state"] == "ALIVE": alive_hostnames.add(node["nodeManagerHostname"]) - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Node hostname list fetched.", host_name_list=list(alive_hostnames)) else: - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=False, message=f"Unknown view {view}") @routes.get("/nodes/{node_id}") + @dashboard_utils.aiohttp_cache async def get_node(self, req) -> aiohttp.web.Response: node_id = req.match_info.get("node_id") node_info = await DataOrganizer.get_node_info(node_id) - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Node details fetched.", detail=node_info) @routes.get("/memory/memory_table") @@ -110,7 +112,7 @@ class StatsCollector(dashboard_utils.DashboardHeadModule): kwargs["sort_by"] = SortingType(sort_by) memory_table = await DataOrganizer.get_memory_table(**kwargs) - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Fetched memory table", memory_table=memory_table.as_dict()) @@ -123,10 +125,10 @@ class StatsCollector(dashboard_utils.DashboardHeadModule): elif should_fetch == "false": self._collect_memory_info = False else: - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=False, message=f"Unknown argument to set_fetch {should_fetch}") - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message=f"Successfully set fetching to {should_fetch}") @@ -136,7 +138,7 @@ class StatsCollector(dashboard_utils.DashboardHeadModule): pid = req.query.get("pid") node_logs = DataSource.ip_and_pid_to_logs[ip] payload = node_logs.get(pid, []) if pid else node_logs - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Fetched logs.", logs=payload) @routes.get("/node_errors") @@ -145,7 +147,7 @@ class StatsCollector(dashboard_utils.DashboardHeadModule): pid = req.query.get("pid") node_errors = DataSource.ip_and_pid_to_errors[ip] filtered_errs = node_errors.get(pid, []) if pid else node_errors - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Fetched errors.", errors=filtered_errs) async def _update_actors(self): diff --git a/dashboard/modules/stats_collector/test_stats_collector.py b/dashboard/modules/stats_collector/test_stats_collector.py index dea19997c..6e371dfa8 100644 --- a/dashboard/modules/stats_collector/test_stats_collector.py +++ b/dashboard/modules/stats_collector/test_stats_collector.py @@ -19,7 +19,7 @@ os.environ["RAY_USE_NEW_DASHBOARD"] = "1" logger = logging.getLogger(__name__) -def test_node_info(ray_start_with_dashboard): +def test_node_info(disable_aiohttp_cache, ray_start_with_dashboard): @ray.remote class Actor: def getpid(self): @@ -168,7 +168,8 @@ def test_get_all_node_details(ray_start_with_dashboard): "ray_start_cluster_head", [{ "include_dashboard": True }], indirect=True) -def test_multi_nodes_info(enable_test_module, ray_start_cluster_head): +def test_multi_nodes_info(enable_test_module, disable_aiohttp_cache, + ray_start_cluster_head): cluster = ray_start_cluster_head assert (wait_until_server_available(cluster.webui_url) is True) webui_url = cluster.webui_url diff --git a/dashboard/modules/test/test_head.py b/dashboard/modules/test/test_head.py index 699fc1545..bcb52add3 100644 --- a/dashboard/modules/test/test_head.py +++ b/dashboard/modules/test/test_head.py @@ -1,3 +1,4 @@ +import time import logging import aiohttp.web @@ -53,20 +54,20 @@ class TestHead(dashboard_utils.DashboardHeadModule): for k, v in DataSource.__dict__.items() if not k.startswith("_") } - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Fetch all data from datacenter success.", **all_data) else: data = dict(DataSource.__dict__.get(key)) - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message=f"Fetch {key} from datacenter success.", **{key: data}) @routes.get("/test/notified_agents") async def get_notified_agents(self, req) -> aiohttp.web.Response: - return await dashboard_utils.rest_response( + return dashboard_utils.rest_response( success=True, message="Fetch notified agents success.", **self._notified_agents) @@ -78,5 +79,19 @@ class TestHead(dashboard_utils.DashboardHeadModule): url) return aiohttp.web.json_response(result) + @routes.get("/test/aiohttp_cache/{sub_path}") + @dashboard_utils.aiohttp_cache(ttl_seconds=1) + async def test_aiohttp_cache(self, req) -> aiohttp.web.Response: + value = req.query["value"] + return dashboard_utils.rest_response( + success=True, message="OK", value=value, timestamp=time.time()) + + @routes.get("/test/aiohttp_cache_lru/{sub_path}") + @dashboard_utils.aiohttp_cache(ttl_seconds=60, maxsize=5) + async def test_aiohttp_cache_lru(self, req) -> aiohttp.web.Response: + value = req.query.get("value") + return dashboard_utils.rest_response( + success=True, message="OK", value=value, timestamp=time.time()) + async def run(self, server): pass diff --git a/dashboard/modules/tune/tune_head.py b/dashboard/modules/tune/tune_head.py index a0c6b32c0..3f10e5df6 100644 --- a/dashboard/modules/tune/tune_head.py +++ b/dashboard/modules/tune/tune_head.py @@ -37,7 +37,7 @@ class TuneController(dashboard_utils.DashboardHeadModule): @routes.get("/tune/info") async def tune_info(self, req) -> aiohttp.web.Response: stats = self.get_stats() - return await rest_response( + return rest_response( success=True, message="Fetched tune info", result=stats) @routes.get("/tune/availability") @@ -46,7 +46,7 @@ class TuneController(dashboard_utils.DashboardHeadModule): "available": Analysis is not None, "trials_available": self._trials_available } - return await rest_response( + return rest_response( success=True, message="Fetched tune availability", result=availability) @@ -56,17 +56,17 @@ class TuneController(dashboard_utils.DashboardHeadModule): experiment = req.query["experiment"] err, experiment = self.set_experiment(experiment) if err: - return await rest_response(success=False, error=err) - return await rest_response( + return rest_response(success=False, error=err) + return rest_response( success=True, message="Successfully set experiment", **experiment) @routes.get("/tune/enable_tensorboard") async def enable_tensorboard(self, req) -> aiohttp.web.Response: self._enable_tensorboard() if not self._tensor_board_dir: - return await rest_response( + return rest_response( success=False, message="Error enabling tensorboard") - return await rest_response(success=True, message="Enabled tensorboard") + return rest_response(success=True, message="Enabled tensorboard") def get_stats(self): tensor_board_info = { diff --git a/dashboard/tests/conftest.py b/dashboard/tests/conftest.py index 030e25806..cb49e8bfc 100644 --- a/dashboard/tests/conftest.py +++ b/dashboard/tests/conftest.py @@ -8,3 +8,10 @@ def enable_test_module(): os.environ["RAY_DASHBOARD_MODULE_TEST"] = "true" yield os.environ.pop("RAY_DASHBOARD_MODULE_TEST", None) + + +@pytest.fixture +def disable_aiohttp_cache(): + os.environ["RAY_DASHBOARD_NO_CACHE"] = "true" + yield + os.environ.pop("RAY_DASHBOARD_NO_CACHE", None) diff --git a/dashboard/tests/test_dashboard.py b/dashboard/tests/test_dashboard.py index a88061edf..47d5926fa 100644 --- a/dashboard/tests/test_dashboard.py +++ b/dashboard/tests/test_dashboard.py @@ -4,6 +4,7 @@ import json import time import logging import asyncio +import collections import aiohttp.web import ray @@ -104,10 +105,14 @@ def test_basic(ray_start_with_dashboard): agent_proc.kill() agent_proc.wait() # The agent will be restarted for imports failure. - for x in range(40): + for x in range(50): agent_proc = _search_agent(raylet_proc.children()) if agent_proc: agent_pids.add(agent_proc.pid) + # The agent should be restarted, + # so we can break if the len(agent_pid) > 1 + if len(agent_pids) > 1: + break time.sleep(0.1) finally: cleanup_test_files() @@ -364,5 +369,78 @@ print("success") run_string_as_driver(test_code) +def test_aiohttp_cache(enable_test_module, ray_start_with_dashboard): + assert (wait_until_server_available(ray_start_with_dashboard["webui_url"]) + is True) + webui_url = ray_start_with_dashboard["webui_url"] + webui_url = format_web_url(webui_url) + + timeout_seconds = 5 + start_time = time.time() + value1_timestamps = [] + while True: + time.sleep(1) + try: + for x in range(10): + response = requests.get(webui_url + + "/test/aiohttp_cache/t1?value=1") + response.raise_for_status() + timestamp = response.json()["data"]["timestamp"] + value1_timestamps.append(timestamp) + assert len(collections.Counter(value1_timestamps)) > 1 + break + except (AssertionError, requests.exceptions.ConnectionError) as e: + logger.info("Retry because of %s", e) + finally: + if time.time() > start_time + timeout_seconds: + raise Exception("Timed out while testing.") + + sub_path_timestamps = [] + for x in range(10): + response = requests.get(webui_url + + f"/test/aiohttp_cache/tt{x}?value=1") + response.raise_for_status() + timestamp = response.json()["data"]["timestamp"] + sub_path_timestamps.append(timestamp) + assert len(collections.Counter(sub_path_timestamps)) == 10 + + volatile_value_timestamps = [] + for x in range(10): + response = requests.get(webui_url + + f"/test/aiohttp_cache/tt?value={x}") + response.raise_for_status() + timestamp = response.json()["data"]["timestamp"] + volatile_value_timestamps.append(timestamp) + assert len(collections.Counter(volatile_value_timestamps)) == 10 + + response = requests.get(webui_url + "/test/aiohttp_cache/raise_exception") + response.raise_for_status() + result = response.json() + assert result["result"] is False + assert "KeyError" in result["msg"] + + volatile_value_timestamps = [] + for x in range(10): + response = requests.get(webui_url + + f"/test/aiohttp_cache_lru/tt{x % 4}") + response.raise_for_status() + timestamp = response.json()["data"]["timestamp"] + volatile_value_timestamps.append(timestamp) + assert len(collections.Counter(volatile_value_timestamps)) == 4 + + volatile_value_timestamps = [] + data = collections.defaultdict(set) + for x in [0, 1, 2, 3, 4, 5, 2, 1, 0, 3]: + response = requests.get(webui_url + + f"/test/aiohttp_cache_lru/t1?value={x}") + response.raise_for_status() + timestamp = response.json()["data"]["timestamp"] + data[x].add(timestamp) + volatile_value_timestamps.append(timestamp) + assert len(collections.Counter(volatile_value_timestamps)) == 8 + assert len(data[3]) == 2 + assert len(data[0]) == 2 + + if __name__ == "__main__": sys.exit(pytest.main(["-v", __file__])) diff --git a/dashboard/utils.py b/dashboard/utils.py index c4f8010dc..e9cfcaf68 100644 --- a/dashboard/utils.py +++ b/dashboard/utils.py @@ -1,5 +1,6 @@ import abc import socket +import time import asyncio import collections import copy @@ -18,6 +19,7 @@ from typing import Any import os import aioredis import aiohttp.web +import ray.new_dashboard.consts as dashboard_consts from aiohttp import hdrs from aiohttp.frozenlist import FrozenList from aiohttp.typedefs import PathLike @@ -25,6 +27,12 @@ from aiohttp.web import RouteDef import aiohttp.signals from google.protobuf.json_format import MessageToDict from ray.utils import binary_to_hex +from ray.ray_constants import env_bool + +try: + create_task = asyncio.create_task +except AttributeError: + create_task = asyncio.ensure_future logger = logging.getLogger(__name__) @@ -112,13 +120,16 @@ class ClassMethodRouteTable: handler.__code__.co_firstlineno, None) @functools.wraps(handler) - async def _handler_route(*args, **kwargs): - if len(args) and args[0] == bind_info.instance: - args = args[1:] + async def _handler_route(*args) -> aiohttp.web.Response: try: - return await handler(bind_info.instance, *args, **kwargs) + # Make the route handler as a bound method. + # The args may be: + # * (Request, ) + # * (self, Request) + req = args[-1] + return await handler(bind_info.instance, req) except Exception: - return await rest_response( + return rest_response( success=False, message=traceback.format_exc()) cls._bind_map[method][path] = bind_info @@ -217,7 +228,7 @@ class CustomEncoder(json.JSONEncoder): return json.JSONEncoder.default(self, obj) -async def rest_response(success, message, **kwargs) -> aiohttp.web.Response: +def rest_response(success, message, **kwargs) -> aiohttp.web.Response: # In the dev context we allow a dev server running on a # different port to consume the API, meaning we need to allow # cross-origin access @@ -292,6 +303,86 @@ def message_to_dict(message, decode_keys=None, **kwargs): return MessageToDict(message, use_integers_for_enums=False, **kwargs) +# The cache value type used by aiohttp_cache. +_AiohttpCacheValue = namedtuple("AiohttpCacheValue", + ["data", "expiration", "task"]) +# The methods with no request body used by aiohttp_cache. +_AIOHTTP_CACHE_NOBODY_METHODS = {hdrs.METH_GET, hdrs.METH_DELETE} + + +def aiohttp_cache( + ttl_seconds=dashboard_consts.AIOHTTP_CACHE_TTL_SECONDS, + maxsize=dashboard_consts.AIOHTTP_CACHE_MAX_SIZE, + enable=not env_bool( + dashboard_consts.AIOHTTP_CACHE_DISABLE_ENVIRONMENT_KEY, False)): + assert maxsize > 0 + cache = collections.OrderedDict() + + def _wrapper(handler): + if enable: + + @functools.wraps(handler) + async def _cache_handler(*args) -> aiohttp.web.Response: + # Make the route handler as a bound method. + # The args may be: + # * (Request, ) + # * (self, Request) + req = args[-1] + # Make key. + if req.method in _AIOHTTP_CACHE_NOBODY_METHODS: + key = req.path_qs + else: + key = (req.path_qs, await req.read()) + # Query cache. + value = cache.get(key) + if value is not None: + cache.move_to_end(key) + if (not value.task.done() + or value.expiration >= time.time()): + # Update task not done or the data is not expired. + return aiohttp.web.Response(**value.data) + + def _update_cache(task): + try: + response = task.result() + except Exception: + response = rest_response( + success=False, message=traceback.format_exc()) + data = { + "status": response.status, + "headers": dict(response.headers), + "body": response.body, + } + cache[key] = _AiohttpCacheValue(data, + time.time() + ttl_seconds, + task) + cache.move_to_end(key) + if len(cache) > maxsize: + cache.popitem(last=False) + return response + + task = create_task(handler(*args)) + task.add_done_callback(_update_cache) + if value is None: + return await task + else: + return aiohttp.web.Response(**value.data) + + suffix = f"[cache ttl={ttl_seconds}, max_size={maxsize}]" + _cache_handler.__name__ += suffix + _cache_handler.__qualname__ += suffix + return _cache_handler + else: + return handler + + if inspect.iscoroutinefunction(ttl_seconds): + target_func = ttl_seconds + ttl_seconds = dashboard_consts.AIOHTTP_CACHE_TTL_SECONDS + return _wrapper(target_func) + else: + return _wrapper + + class SignalManager: _signals = FrozenList()