diff --git a/spark_executor/tools/connections.py b/spark_executor/tools/connections.py new file mode 100644 index 0000000..ccff910 --- /dev/null +++ b/spark_executor/tools/connections.py @@ -0,0 +1,47 @@ +# coding=utf-8 +""" +@Time :2026/6/24 +@Author :tao.chen +""" +from common.logging import logger +from spark_executor.core.connection_store import ConnectionStore, store +from spark_executor.models import Connection + + +def save_connection( + *, + name: str, + master: str, + deploy_mode: str = "cluster", + yarn_rm_url: str | None = None, + spark_conf: dict[str, str] | None = None, +) -> dict[str, str]: + conn = Connection( + name=name, + master=master, + deploy_mode=deploy_mode, + yarn_rm_url=yarn_rm_url, + spark_conf=spark_conf or {}, + ) + store.save(conn) + logger.info(f"save_connection name={name} master={master}") + return {"name": name, "status": "SAVED"} + + +def list_connections() -> list[dict[str, object]]: + return [c.model_dump() for c in store.list_all()] + + +def get_connection(name: str) -> dict[str, object]: + conn = store.get(name) + if conn is None: + raise KeyError(f"Unknown connection: {name}") + return conn.model_dump() + + +def delete_connection(name: str) -> dict[str, str]: + removed = store.delete(name) + if not removed: + raise KeyError(f"Unknown connection: {name}") + logger.info(f"delete_connection name={name}") + return {"name": name, "status": "DELETED"} diff --git a/tests/unit/test_connection_tools.py b/tests/unit/test_connection_tools.py new file mode 100644 index 0000000..4307622 --- /dev/null +++ b/tests/unit/test_connection_tools.py @@ -0,0 +1,85 @@ +# coding=utf-8 +from pathlib import Path + +import pytest + +from spark_executor.core import connection_store +from spark_executor.tools import connections + + +@pytest.fixture(autouse=True) +def _fresh_store(tmp_path: Path, monkeypatch): + monkeypatch.setattr(connection_store, "DEFAULT_DATA_DIR", str(tmp_path)) + monkeypatch.setattr(connection_store, "store", connection_store.ConnectionStore()) + monkeypatch.setattr(connections, "store", connection_store.store) + return connection_store.store + + +# --- save_connection --- + +def test_save_connection(_fresh_store): + out = connections.save_connection(name="prod", master="yarn") + assert out == {"name": "prod", "status": "SAVED"} + assert _fresh_store.get("prod") is not None + + +def test_save_connection_upserts(_fresh_store): + connections.save_connection(name="prod", master="yarn") + connections.save_connection(name="prod", master="spark://new:7077", deploy_mode="client") + assert _fresh_store.get("prod").master == "spark://new:7077" + assert _fresh_store.get("prod").deploy_mode == "client" + + +def test_save_connection_with_spark_conf(_fresh_store): + connections.save_connection( + name="dev", + master="yarn", + spark_conf={"spark.sql.shuffle.partitions": "200"}, + ) + assert _fresh_store.get("dev").spark_conf["spark.sql.shuffle.partitions"] == "200" + + +# --- list_connections --- + +def test_list_connections_empty(_fresh_store): + assert connections.list_connections() == [] + + +def test_list_connections_returns_all_saved(_fresh_store): + connections.save_connection(name="a", master="yarn") + connections.save_connection(name="b", master="spark://m:7077", deploy_mode="client") + out = connections.list_connections() + names = {c["name"] for c in out} + assert names == {"a", "b"} + masters = {c["name"]: c["master"] for c in out} + assert masters == {"a": "yarn", "b": "spark://m:7077"} + + +# --- get_connection --- + +def test_get_connection_returns_dict(_fresh_store): + connections.save_connection(name="prod", master="yarn", deploy_mode="cluster") + out = connections.get_connection("prod") + assert out["name"] == "prod" + assert out["master"] == "yarn" + assert out["deploy_mode"] == "cluster" + assert out["yarn_rm_url"] is None + assert out["spark_conf"] == {} + + +def test_get_connection_unknown_raises(_fresh_store): + with pytest.raises(KeyError): + connections.get_connection("missing") + + +# --- delete_connection --- + +def test_delete_connection_returns_status(_fresh_store): + connections.save_connection(name="prod", master="yarn") + assert connections.delete_connection("prod") == {"name": "prod", "status": "DELETED"} + assert _fresh_store.get("prod") is None + + +def test_delete_connection_unknown_raises(_fresh_store): + with pytest.raises(KeyError): + connections.delete_connection("missing")