# coding=utf-8 import base64 from unittest.mock import patch import httpx import pytest from common import config from spark_executor.models import Connection from spark_executor.core.yarn_client import ( YarnConfigError, YarnError, YarnClientConfig, get_application_logs, get_application_status, kill_application, ) RM = "http://rm:8088" HTTPS_RM = "https://rm:8088" @pytest.fixture(autouse=True) def _restore_settings(): """Each test gets a clean copy of `settings` so env-var-style overrides in one test don't leak into the next.""" snapshot = config.Settings( data_dir=config.settings.data_dir, jobs_dir=config.settings.jobs_dir, yarn_resource_manager_url=config.settings.yarn_resource_manager_url, log_level=config.settings.log_level, ssl_verify_default=config.settings.ssl_verify_default, ssl_ca_bundle_default=config.settings.ssl_ca_bundle_default, ) yield config.settings.data_dir = snapshot.data_dir config.settings.jobs_dir = snapshot.jobs_dir config.settings.yarn_resource_manager_url = snapshot.yarn_resource_manager_url config.settings.log_level = snapshot.log_level config.settings.ssl_verify_default = snapshot.ssl_verify_default config.settings.ssl_ca_bundle_default = snapshot.ssl_ca_bundle_default def _resp(status: int, *, json_data=None, text: str | None = None, headers: dict | None = None) -> httpx.Response: kwargs: dict = {} if headers: kwargs["headers"] = headers if json_data is not None: return httpx.Response(status, json=json_data, **kwargs) return httpx.Response(status, text=text or "", **kwargs) # --- get_application_status --- def test_status_parses_app_state(): fake = _resp(200, json_data={"app": {"id": "application_1", "state": "RUNNING"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: state, raw = get_application_status("application_1", YarnClientConfig(yarn_rm_url=RM)) assert state == "RUNNING" assert "RUNNING" in raw args = m.call_args.args assert args == ("GET", f"{RM}/ws/v1/cluster/apps/application_1") assert m.call_args.kwargs["verify"] is True def test_request_sends_accept_json_header(): fake = _resp(200, json_data={"app": {"id": "application_1", "state": "RUNNING"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", YarnClientConfig(yarn_rm_url=RM)) assert m.call_args.kwargs["headers"] == {"Accept": "application/json"} def test_status_raises_on_404(): with patch("spark_executor.core.yarn_client.httpx.request", return_value=_resp(404)): with pytest.raises(YarnError, match="not found"): get_application_status("application_x", YarnClientConfig(yarn_rm_url=RM)) def test_status_raises_on_5xx(): with patch( "spark_executor.core.yarn_client.httpx.request", return_value=_resp(503, text="upstream down"), ): with pytest.raises(YarnError, match="503"): get_application_status("application_1", YarnClientConfig(yarn_rm_url=RM)) def test_status_raises_when_state_field_missing(): fake = _resp(200, json_data={"app": {"id": "application_1"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake): with pytest.raises(YarnError, match="Could not parse YARN state"): get_application_status("application_1", YarnClientConfig(yarn_rm_url=RM)) def test_status_requires_rm_url(): config.settings.yarn_resource_manager_url = None with pytest.raises(YarnConfigError): get_application_status("application_1", YarnClientConfig(yarn_rm_url=None)) def test_status_falls_back_to_settings_yarn_rm_url(): config.settings.yarn_resource_manager_url = "http://env-rm:8088" fake = _resp(200, json_data={"app": {"state": "FINISHED"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: state, _ = get_application_status("application_1", YarnClientConfig(yarn_rm_url=None)) assert state == "FINISHED" assert "env-rm:8088" in m.call_args.args[1] def test_status_rejects_non_http_url(): with pytest.raises(YarnConfigError, match="must start with"): get_application_status("application_1", YarnClientConfig(yarn_rm_url="rm:8088")) # --- get_application_logs --- def test_logs_returns_text_on_h3_aggregated_endpoint(): fake = _resp(200, text="log line 1\nlog line 2\n") with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: out = get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) assert out == "log line 1\nlog line 2\n" assert m.call_args.args == ( "GET", f"{RM}/ws/v1/cluster/apps/application_1/aggregated-logs", ) assert m.call_args.kwargs["verify"] is True def test_logs_falls_back_to_am_container_on_404(): responses = [ _resp(404), # aggregated-logs endpoint missing _resp( 200, json_data={ "app": { "amContainerLogs": "http://nm1:8042/node/containerlogs/container_1/hdfs", } }, ), _resp(200, text="am container log content"), ] with patch( "spark_executor.core.yarn_client.httpx.request", side_effect=responses ) as m: out = get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) assert out == "am container log content" calls = m.call_args_list assert calls[0].args == ( "GET", f"{RM}/ws/v1/cluster/apps/application_1/aggregated-logs", ) assert calls[1].args == ("GET", f"{RM}/ws/v1/cluster/apps/application_1") assert calls[2].args == ( "GET", "http://nm1:8042/node/containerlogs/container_1/hdfs/stdout", ) def test_logs_raises_on_404_when_am_container_field_missing(): responses = [ _resp(404), # aggregated-logs endpoint missing _resp(200, json_data={"app": {}}), ] with patch( "spark_executor.core.yarn_client.httpx.request", side_effect=responses ): with pytest.raises(YarnError, match="log-aggregation-enable"): get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) def test_logs_returns_am_container_stdout(): responses = [ _resp(404), # aggregated-logs endpoint missing _resp( 200, json_data={ "app": { "amContainerLogs": "http://nm1:8042/node/containerlogs/container_1/hdfs", } }, ), _resp(200, text="driver stdout"), ] with patch( "spark_executor.core.yarn_client.httpx.request", side_effect=responses ) as m: out = get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) assert out == "driver stdout" assert m.call_args_list[2].args == ( "GET", "http://nm1:8042/node/containerlogs/container_1/hdfs/stdout", ) def test_logs_follows_307_redirect_to_other_host(): """RM 307-redirects the log fetch to the actual NodeManager host. The helper must follow Location and re-apply auth/verify on the new host.""" responses = [ _resp(404), # aggregated-logs missing _resp( 200, json_data={ "app": { "amContainerLogs": "http://rm:8088/node/containerlogs/container_1/hdfs", } }, ), _resp(307, headers={"Location": "http://nm2:8042/node/containerlogs/container_1/hdfs/stdout"}), _resp(200, text="driver stdout via redirect"), ] with patch( "spark_executor.core.yarn_client.httpx.request", side_effect=responses ) as m: out = get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) assert out == "driver stdout via redirect" calls = m.call_args_list assert calls[2].args == ("GET", "http://rm:8088/node/containerlogs/container_1/hdfs/stdout") assert calls[3].args == ("GET", "http://nm2:8042/node/containerlogs/container_1/hdfs/stdout") # Auth must be re-applied on the redirected request. assert calls[3].kwargs["auth"] == calls[2].kwargs["auth"] def test_logs_307_without_location_raises_unavailable(): """3xx without a Location header is a misconfigured server — surface as YarnError rather than silently returning empty text.""" responses = [ _resp(404), # aggregated-logs missing _resp( 200, json_data={ "app": { "amContainerLogs": "http://rm:8088/node/containerlogs/container_1/hdfs", } }, ), _resp(307), # no Location header ] with patch( "spark_executor.core.yarn_client.httpx.request", side_effect=responses ): with pytest.raises(YarnError, match="no Location header"): get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) def test_logs_follows_redirect_chain_then_200(): """Two consecutive 307s followed by a 200 — the helper should walk the chain and return the final body.""" responses = [ _resp(404), # aggregated-logs missing _resp( 200, json_data={ "app": { "amContainerLogs": "http://lb:8088/node/containerlogs/container_1/hdfs", } }, ), _resp(307, headers={"Location": "http://rm:8088/node/containerlogs/container_1/hdfs/stdout"}), _resp(307, headers={"Location": "http://nm2:8042/node/containerlogs/container_1/hdfs/stdout"}), _resp(200, text="chained redirect body"), ] with patch( "spark_executor.core.yarn_client.httpx.request", side_effect=responses ) as m: out = get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) assert out == "chained redirect body" urls = [c.args[1] for c in m.call_args_list[2:]] assert urls == [ "http://lb:8088/node/containerlogs/container_1/hdfs/stdout", "http://rm:8088/node/containerlogs/container_1/hdfs/stdout", "http://nm2:8042/node/containerlogs/container_1/hdfs/stdout", ] def test_logs_5xx_on_aggregated_endpoint_raises_immediately(): with patch( "spark_executor.core.yarn_client.httpx.request", return_value=_resp(500, text="boom"), ) as m: with pytest.raises(YarnError, match="500"): get_application_logs("application_1", YarnClientConfig(yarn_rm_url=RM)) assert m.call_count == 1 # --- kill_application --- def test_kill_sends_put_with_killed_state(): fake = _resp(200, json_data={"app": {"state": "KILLED"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: kill_application("application_1", YarnClientConfig(yarn_rm_url=RM)) args = m.call_args.args assert args == ("PUT", f"{RM}/ws/v1/cluster/apps/application_1/state") assert m.call_args.kwargs["json"] == {"state": "KILLED"} assert m.call_args.kwargs["verify"] is True def test_kill_raises_on_5xx(): with patch( "spark_executor.core.yarn_client.httpx.request", return_value=_resp(403, text="forbidden"), ): with pytest.raises(YarnError, match="403"): kill_application("application_1", YarnClientConfig(yarn_rm_url=RM)) # --- connection errors --- def test_status_wraps_httpx_errors_as_yarn_error(): with patch( "spark_executor.core.yarn_client.httpx.request", side_effect=httpx.ConnectError("connection refused"), ): with pytest.raises(YarnError, match="connection failed"): get_application_status("application_1", YarnClientConfig(yarn_rm_url=RM)) # --- SSL/TLS configuration --- def test_request_passes_ssl_verify_true(): config.settings.ssl_verify_default = True fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status( "application_1", YarnClientConfig(yarn_rm_url=RM, ssl_verify=True), ) assert m.call_args.kwargs["verify"] is True def test_request_passes_ssl_ca_bundle_path(): fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status( "application_1", YarnClientConfig( yarn_rm_url=HTTPS_RM, ssl_ca_bundle="/path/to/ca.crt", ), ) assert m.call_args.kwargs["verify"] == "/path/to/ca.crt" def test_request_uses_global_ssl_default_when_connection_unset(): config.settings.ssl_verify_default = True config.settings.ssl_ca_bundle_default = "/global/ca.crt" config.settings.yarn_resource_manager_url = "https://rm:8088" fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) conn = Connection(name="prod", master="yarn", ssl_verify=None, ssl_ca_bundle=None) cfg = YarnClientConfig.from_connection(conn) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", cfg) assert m.call_args.kwargs["verify"] == "/global/ca.crt" def test_request_per_connection_ssl_overrides_global(): config.settings.ssl_verify_default = True config.settings.ssl_ca_bundle_default = "/global/ca.crt" config.settings.yarn_resource_manager_url = "https://rm:8088" fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) conn = Connection( name="prod", master="yarn", ssl_verify=False, ssl_ca_bundle="/conn/ca.crt", ) cfg = YarnClientConfig.from_connection(conn) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", cfg) assert m.call_args.kwargs["verify"] == "/conn/ca.crt" def test_ssl_verify_false_without_ca_bundle_uses_global_default(): """If a connection only sets ssl_verify=False and no CA bundle is configured globally, httpx should skip certificate verification.""" config.settings.ssl_verify_default = True config.settings.ssl_ca_bundle_default = None config.settings.yarn_resource_manager_url = "https://rm:8088" fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) conn = Connection(name="prod", master="yarn", ssl_verify=False) cfg = YarnClientConfig.from_connection(conn) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", cfg) assert m.call_args.kwargs["verify"] is False # --- authentication configuration --- def test_status_no_auth_for_none(): fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", YarnClientConfig(yarn_rm_url=RM, auth_type="none")) assert m.call_args.kwargs["auth"] is None def test_status_basic_auth_header(): fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) cfg = YarnClientConfig( yarn_rm_url=RM, auth_type="basic", auth_user="alice", auth_password="pw", ) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", cfg) auth = m.call_args.kwargs["auth"] assert isinstance(auth, httpx.BasicAuth) encoded = auth._auth_header.split(" ", 1)[1] assert base64.b64decode(encoded).decode() == "alice:pw" def test_status_no_auth_for_simple(): fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", YarnClientConfig(yarn_rm_url=RM, auth_type="simple")) assert m.call_args.kwargs["auth"] is None def test_status_kerberos_auth(): fake = _resp(200, json_data={"app": {"state": "RUNNING"}}) cfg = YarnClientConfig(yarn_rm_url=RM, auth_type="kerberos") with patch("spark_executor.core.yarn_client.httpx.request", return_value=fake) as m: get_application_status("application_1", cfg) auth = m.call_args.kwargs["auth"] assert auth is not None assert "Kerberos" in type(auth).__name__ def test_basic_auth_missing_user_raises(): cfg = YarnClientConfig(yarn_rm_url=RM, auth_type="basic") with pytest.raises(YarnConfigError, match="auth_type='basic' requires auth_user"): get_application_status("application_1", cfg)