diff --git a/mamba_gator/envmanager.py b/mamba_gator/envmanager.py index 5d9a7d7a..50a8bad8 100644 --- a/mamba_gator/envmanager.py +++ b/mamba_gator/envmanager.py @@ -340,18 +340,25 @@ async def conda_config(self) -> Dict[str, Any]: _, output = ans return self._clean_conda_json(output) - async def clone_env(self, env: str, name: str) -> Dict[str, str]: + async def clone_env(self, env: str, name: str, channels: Optional[List[str]] = None) -> Dict[str, str]: """Clone an environment. Args: env (str): To-be-cloned environment name name (str): New environment name + channels (List[str], optional): Channel priority list (e.g., ["conda-forge", "defaults"]) Returns: Dict[str, str]: Clone command output. """ + cmd_args = ["-y", "-q", "--json", "-n", name] + + if channels: + for channel in channels: + cmd_args.extend(["-c", channel]) + ans = await self._execute( - self.manager, "create", "-y", "-q", "--json", "-n", name, "--clone", env + self.manager, "create", *cmd_args, "--clone", env ) rcode, output = ans @@ -360,18 +367,25 @@ async def clone_env(self, env: str, name: str) -> Dict[str, str]: return self._clean_conda_json(output) - async def create_env(self, env: str, *args) -> Dict[str, str]: + async def create_env(self, env: str, *args, channels: Optional[List[str]] = None) -> Dict[str, str]: """Create a environment from a list of packages. Args: env (str): Name of the environment *args (List[str]): optional, packages to install + channels (List[str], optional): Channel priority list (e.g., ["conda-forge", "defaults"]) Returns: Dict[str, str]: Create command output """ + cmd_args = ["-y", "-q", "--json", "-n", env] + + if channels: + for channel in channels: + cmd_args.extend(["-c", channel]) + ans = await self._execute( - self.manager, "create", "-y", "-q", "--json", "-n", env, *args + self.manager, "create", *cmd_args, *args ) rcode, output = ans @@ -565,7 +579,7 @@ def get_info(env): return {"environments": envs_list} async def update_env( - self, env: str, file_content: str, file_name: str = "environment.yml" + self, env: str, file_content: str, file_name: str = "environment.yml", channels: Optional[List[str]] = None ) -> Dict[str, str]: """Update a environment from a file. @@ -573,6 +587,7 @@ async def update_env( env (str): Name of the environment file_content (str): File content file_name (str): optional, Original filename + channels (List[str], optional): Channel priority list (e.g., ["conda-forge", "defaults"]) Returns: Dict[str, str]: Update command output @@ -590,9 +605,14 @@ async def update_env( if file_name.endswith('.txt'): # For .txt files (explicit package lists), use conda install self.log.debug(f"Updating environment {env} with txt file using conda install") - ans = await self._execute( - self.manager, "install", "-y", "-q", "--json", "-n", env, "--file", name - ) + cmd_args = ["-y", "-q", "--json", "-n", env] + + if channels: + for channel in channels: + cmd_args.extend(["-c", channel]) + + cmd_args.extend(["--file", name]) + ans = await self._execute(self.manager, "install", *cmd_args) else: # For .yml files (environment definitions), use conda env update self.log.debug(f"Updating environment {env} with yml file using conda env update") @@ -1032,19 +1052,26 @@ async def check_update( # no action plan returned means everything is already up to date return {"updates": []} - async def install_packages(self, env: str, packages: List[str]) -> Dict[str, str]: + async def install_packages(self, env: str, packages: List[str], channels: Optional[List[str]] = None) -> Dict[str, str]: """Install packages in an environment. Args: env (str): Environment name packages (List[str]): List of packages to install + channels (List[str], optional): Channel priority list (e.g., ["conda-forge", "defaults"]) Returns: Dict[str, str]: Install command output. """ - ans = await self._execute( - self.manager, "install", "-y", "-q", "--json", "-n", env, *packages - ) + cmd_args = ["-y", "-q", "--json", "-n", env] + + if channels: + for channel in channels: + cmd_args.extend(["-c", channel]) + + cmd_args.extend(packages) + + ans = await self._execute(self.manager, "install", *cmd_args) _, output = ans return self._clean_conda_json(output) @@ -1134,19 +1161,26 @@ async def develop_packages( return {"packages": result} - async def update_packages(self, env: str, packages: List[str]) -> Dict[str, str]: + async def update_packages(self, env: str, packages: List[str], channels: Optional[List[str]] = None) -> Dict[str, str]: """Update packages in an environment. Args: env (str): Environment name packages (List[str]): List of packages to update + channels (List[str], optional): Channel priority list (e.g., ["conda-forge", "defaults"]) Returns: Dict[str, str]: Update command output. """ - ans = await self._execute( - self.manager, "update", "-y", "-q", "--json", "-n", env, *packages - ) + cmd_args = ["-y", "-q", "--json", "-n", env] + + if channels: + for channel in channels: + cmd_args.extend(["-c", channel]) + + cmd_args.extend(packages) + + ans = await self._execute(self.manager, "update", *cmd_args) _, output = ans return self._clean_conda_json(output) diff --git a/mamba_gator/handlers.py b/mamba_gator/handlers.py index b722e759..aebd7cd8 100644 --- a/mamba_gator/handlers.py +++ b/mamba_gator/handlers.py @@ -77,12 +77,13 @@ def get(self, idx: int) -> Any: else: return None - def put(self, task: Callable, *args) -> int: + def put(self, task: Callable, *args, **kwargs) -> int: """Add a asynchronous task into the queue. Args: task (Callable): Asynchronous task *args : arguments of the task + **kwargs : keyword arguments of the task Returns: int: Task id @@ -90,10 +91,10 @@ def put(self, task: Callable, *args) -> int: ActionsStack.__last_index += 1 idx = ActionsStack.__last_index - async def execute_task(idx, f, *args) -> Any: + async def execute_task(idx, f, *args, **kwargs) -> Any: try: get_logger().debug("Will execute task {}.".format(idx)) - result = await f(*args) + result = await f(*args, **kwargs) except asyncio.CancelledError: raise except Exception as e: @@ -110,7 +111,7 @@ async def execute_task(idx, f, *args) -> Any: return result - self.__tasks[idx] = asyncio.ensure_future(execute_task(idx, task, *args)) + self.__tasks[idx] = asyncio.ensure_future(execute_task(idx, task, *args, **kwargs)) return idx def __del__(self): @@ -191,6 +192,7 @@ def post(self): twin (str): optional, environment name to clone file (str): optional, environment file (TXT or YAML format) filename (str): optional, environment filename of the `file` content + channels (List[str]): optional, channel priority list (e.g., ["conda-forge", "defaults"]) } """ data = self.get_json_body() @@ -199,17 +201,18 @@ def post(self): twin = data.get("twin", None) file_content = data.get("file", None) file_name = data.get("filename", "environment.txt") + channels = data.get("channels", None) if packages is not None: - idx = self._stack.put(self.env_manager.create_env, name, *packages) + idx = self._stack.put(self.env_manager.create_env, name, *packages, channels=channels) elif twin is not None: - idx = self._stack.put(self.env_manager.clone_env, twin, name) + idx = self._stack.put(self.env_manager.clone_env, twin, name, channels=channels) elif file_content is not None: idx = self._stack.put( self.env_manager.import_env, name, file_content, file_name ) else: - idx = self._stack.put(self.env_manager.create_env, name) + idx = self._stack.put(self.env_manager.create_env, name, channels=channels) self.redirect_to_task(idx) @@ -270,13 +273,15 @@ def patch(self, env: str): { file (str): optional, environment file (TXT or YAML format) filename (str): optional, environment filename of the `file` content + channels (List[str]): optional, channel priority list (e.g., ["conda-forge", "defaults"]) } """ data = self.get_json_body() file_content = data["file"] file_name = data.get("filename", "environment.yml") + channels = data.get("channels", None) - idx = self._stack.put(self.env_manager.update_env, env, file_content, file_name) + idx = self._stack.put(self.env_manager.update_env, env, file_content, file_name, channels=channels) self.redirect_to_task(idx) @@ -307,11 +312,13 @@ def patch(self, env: str): Request json body: { packages (List[str]): optional, list of packages to update + channels (List[str]): optional, channel priority list (e.g., ["conda-forge", "defaults"]) } """ body = self.get_json_body() or {} packages = body.get("packages", ["--all"]) - idx = self._stack.put(self.env_manager.update_packages, env, packages) + channels = body.get("channels", None) + idx = self._stack.put(self.env_manager.update_packages, env, packages, channels=channels) self.redirect_to_task(idx) @tornado.web.authenticated @@ -327,16 +334,18 @@ def post(self, env: str): Request json body: { packages (List[str]): list of packages to install + channels (List[str]): optional, channel priority list (e.g., ["conda-forge", "defaults"]) } """ body = self.get_json_body() packages = body["packages"] + channels = body.get("channels", None) develop = int(self.get_query_argument("develop", 0)) if develop: idx = self._stack.put(self.env_manager.develop_packages, env, packages) else: - idx = self._stack.put(self.env_manager.install_packages, env, packages) + idx = self._stack.put(self.env_manager.install_packages, env, packages, channels=channels) self.redirect_to_task(idx) diff --git a/mamba_gator/tests/test_api.py b/mamba_gator/tests/test_api.py index e3e09034..dcfa8d87 100644 --- a/mamba_gator/tests/test_api.py +++ b/mamba_gator/tests/test_api.py @@ -44,7 +44,7 @@ def wait_for_task(self, call, *args, **kwargs): self.assertRegex(location, r"^/conda/tasks/\d+$") return self.wait_task(location) - def mk_env(self, name=None, packages=None, remove_if_exists=True): + def mk_env(self, name=None, packages=None, remove_if_exists=True, channels=None): envs = self.conda_api.envs() env_names = map(lambda env: env["name"], envs["environments"]) new_name = name or generate_name() @@ -54,10 +54,10 @@ def mk_env(self, name=None, packages=None, remove_if_exists=True): # TODO: Remove this once we have a way to test the environment creation with packages from different channels # or once packages are available in the default channel - return self.conda_api.post( - ["environments"], - body={"name": new_name, "packages": packages or ["python!=3.14.0"]}, - ) + body = {"name": new_name, "packages": packages or ["python"]} + if channels: + body["channels"] = channels + return self.conda_api.post(["environments"], body=body) def rm_env(self, name): answer = self.conda_api.delete(["environments", name]) @@ -263,7 +263,7 @@ def test_environment_yaml_import(self): - conda-forge - defaults dependencies: -- python=3.9 +- python - astroid prefix: /home/user/.conda/envs/lab_conda """ @@ -294,7 +294,7 @@ def test_environment_text_import(self): self.env_names.append(n) content = """# This file may be used to create an environment using: # $ conda create --name --file -python=3.9 +python astroid """ @@ -317,7 +317,7 @@ def test_update_env_yaml(self): self.skipTest("FIXME not working with mamba") n = generate_name() - response = self.wait_for_task(self.mk_env, n, ["python=3.9"]) + response = self.wait_for_task(self.mk_env, n, ["python"]) self.assertEqual(response.status_code, 200) content = """name: test_conda @@ -353,7 +353,7 @@ def test_update_env_no_filename(self): self.skipTest("FIXME not working with mamba") n = generate_name() - response = self.wait_for_task(self.mk_env, n, ["python=3.9"]) + response = self.wait_for_task(self.mk_env, n, ["python"]) self.assertEqual(response.status_code, 200) content = """name: test_conda @@ -384,7 +384,7 @@ def g(): def test_update_env_txt(self): n = generate_name() - response = self.wait_for_task(self.mk_env, n, ["python=3.9"]) + response = self.wait_for_task(self.mk_env, n, ["python"], channels=["conda-forge", "defaults"]) self.assertEqual(response.status_code, 200) content = """# This file may be used to create an environment using: @@ -394,7 +394,7 @@ def test_update_env_txt(self): def g(): return self.conda_api.patch( - ["environments", n], body={"file": content, "filename": "testenv.txt"}, + ["environments", n], body={"file": content, "filename": "testenv.txt", "channels": ["conda-forge", "defaults"]}, ) response = self.wait_for_task(g) @@ -485,7 +485,7 @@ def test_get_has_update(self): self.wait_for_task( self.conda_api.post, ["environments"], - body={"name": n, "packages": ["python=3.9", "astroid"]}, + body={"name": n, "packages": ["python", "astroid"]}, ) r = self.wait_for_task( @@ -516,7 +516,7 @@ def test_env_export(self): def test_env_export_history(self): n = generate_name() - self.wait_for_task(self.mk_env, n, packages=["python=3.9"]) + self.wait_for_task(self.mk_env, n, packages=["python"]) r = self.conda_api.get( ["environments", n], params={"download": 1, "history": 1} ) @@ -524,7 +524,7 @@ def test_env_export_history(self): content = " ".join(r.text.splitlines()) self.assertRegex( - content, r"^name:\s" + n + r"\s+channels:(\s+-\s+[^\s]+)+\s+dependencies:\s+-\s+python=3\.9\s+prefix:" + content, r"^name:\s" + n + r"\s+channels:(\s+-\s+[^\s]+)+\s+dependencies:\s+-\s+python\s+prefix:" ) def test_env_export_not_supporting_history(self): @@ -560,12 +560,13 @@ def test_version(self): class TestPackagesEnvironmentHandler(JupyterCondaAPITest): def test_pkg_install_and_remove(self): n = generate_name() - self.wait_for_task(self.mk_env, n) + self.wait_for_task(self.mk_env, n, packages=["python"], channels=["conda-forge", "defaults"]) + body = {"packages": [self.pkg_name], "channels": ["conda-forge", "defaults"]} r = self.wait_for_task( self.conda_api.post, ["environments", n, "packages"], - body={"packages": [self.pkg_name]}, + body=body, ) self.assertEqual(r.status_code, 200) r = self.conda_api.get(["environments", n]) @@ -595,12 +596,14 @@ def test_pkg_install_and_remove(self): def test_pkg_install_with_version_constraints(self): test_pkg = "astroid" n = generate_name() - self.wait_for_task(self.mk_env, n, packages=["python=3.9"]) + self.wait_for_task(self.mk_env, n, packages=["python"], channels=["conda-forge", "defaults"]) + + body = {"packages": [test_pkg + "==4.0.1"], "channels": ["conda-forge", "defaults"]} r = self.wait_for_task( self.conda_api.post, ["environments", n, "packages"], - body={"packages": [test_pkg + "==2.14.2"]}, + body=body ) self.assertEqual(r.status_code, 200) r = self.conda_api.get(["environments", n]) @@ -610,14 +613,14 @@ def test_pkg_install_with_version_constraints(self): if p["name"] == test_pkg: v = p break - self.assertEqual(v["version"], "2.14.2") + self.assertEqual(v["version"], "4.0.1") n = generate_name() - self.wait_for_task(self.mk_env, n, packages=["python=3.9"]) + self.wait_for_task(self.mk_env, n, packages=["python"], channels=["conda-forge", "defaults"]) r = self.wait_for_task( self.conda_api.post, ["environments", n, "packages"], - body={"packages": [test_pkg + ">=2.14.0"]}, + body={"packages": [test_pkg + ">=4.0.1"], "channels": ["conda-forge", "defaults"]}, ) self.assertEqual(r.status_code, 200) r = self.conda_api.get(["environments", n]) @@ -627,14 +630,14 @@ def test_pkg_install_with_version_constraints(self): if p["name"] == test_pkg: v = tuple(map(int, p["version"].split("."))) break - self.assertGreaterEqual(v, (2, 14, 0)) + self.assertGreaterEqual(v, (4, 0, 1)) n = generate_name() - self.wait_for_task(self.mk_env, n, packages=["python=3.9"]) + self.wait_for_task(self.mk_env, n, packages=["python"], channels=["conda-forge", "defaults"]) r = self.wait_for_task( self.conda_api.post, ["environments", n, "packages"], - body={"packages": [test_pkg + ">=2.14.0,<3.0.0"]}, + body={"packages": [test_pkg + ">=4.0.1,<5.0.0"], "channels": ["conda-forge", "defaults"]}, ) self.assertEqual(r.status_code, 200) r = self.conda_api.get(["environments", n]) @@ -644,8 +647,8 @@ def test_pkg_install_with_version_constraints(self): if p["name"] == test_pkg: v = tuple(map(int, p["version"].split("."))) break - self.assertGreaterEqual(v, (2, 14, 0)) - self.assertLess(v, (3, 0, 0)) + self.assertGreaterEqual(v, (4, 0, 1)) + self.assertLess(v, (5, 0, 0)) def test_package_install_development_mode(self): n = generate_name()