99 lines
3.0 KiB
Python
99 lines
3.0 KiB
Python
from __future__ import annotations
|
|
|
|
import ctypes
|
|
import os
|
|
import sys
|
|
import winreg
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
from typing import Protocol
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class EnvironmentUpdate:
|
|
path_entries: list[Path] = field(default_factory=list)
|
|
env_vars: dict[str, str] = field(default_factory=dict)
|
|
|
|
|
|
class RegistryLike(Protocol):
|
|
def get_value(self, name: str) -> str: ...
|
|
|
|
def set_value(self, name: str, value: str) -> None: ...
|
|
|
|
def broadcast_change(self) -> None: ...
|
|
|
|
|
|
class UserEnvironmentRegistry:
|
|
key_path = r"Environment"
|
|
|
|
def get_value(self, name: str) -> str:
|
|
try:
|
|
with winreg.OpenKey(winreg.HKEY_CURRENT_USER, self.key_path, 0, winreg.KEY_READ) as key:
|
|
value, _ = winreg.QueryValueEx(key, name)
|
|
return str(value)
|
|
except FileNotFoundError:
|
|
return ""
|
|
|
|
def set_value(self, name: str, value: str) -> None:
|
|
with winreg.CreateKeyEx(winreg.HKEY_CURRENT_USER, self.key_path, 0, winreg.KEY_SET_VALUE) as key:
|
|
winreg.SetValueEx(key, name, 0, winreg.REG_EXPAND_SZ, value)
|
|
|
|
def broadcast_change(self) -> None:
|
|
if sys.platform != "win32":
|
|
return
|
|
hwnd_broadcast = 0xFFFF
|
|
wm_settingchange = 0x001A
|
|
smto_abortifhung = 0x0002
|
|
result = ctypes.c_ulong()
|
|
ctypes.windll.user32.SendMessageTimeoutW(
|
|
hwnd_broadcast,
|
|
wm_settingchange,
|
|
0,
|
|
"Environment",
|
|
smto_abortifhung,
|
|
5000,
|
|
ctypes.byref(result),
|
|
)
|
|
|
|
|
|
def merge_path_entries(current_path: str, entries: list[Path]) -> str:
|
|
parts = [part for part in current_path.split(os.pathsep) if part.strip()]
|
|
seen = {_normalize_path(part) for part in parts}
|
|
|
|
for entry in entries:
|
|
value = str(entry)
|
|
normalized = _normalize_path(value)
|
|
if normalized not in seen:
|
|
parts.append(value)
|
|
seen.add(normalized)
|
|
|
|
return os.pathsep.join(parts)
|
|
|
|
|
|
def build_environment_update(tool_name: str, executable: Path) -> EnvironmentUpdate:
|
|
exe = Path(executable)
|
|
bin_dir = exe.parent
|
|
env_vars: dict[str, str] = {}
|
|
|
|
if tool_name == "JDK":
|
|
env_vars["JAVA_HOME"] = str(bin_dir.parent)
|
|
elif tool_name == "Maven":
|
|
env_vars["MAVEN_HOME"] = str(bin_dir.parent)
|
|
elif tool_name == "Gradle":
|
|
env_vars["GRADLE_HOME"] = str(bin_dir.parent)
|
|
|
|
return EnvironmentUpdate(path_entries=[bin_dir], env_vars=env_vars)
|
|
|
|
|
|
def apply_user_environment(update: EnvironmentUpdate, registry: RegistryLike | None = None) -> None:
|
|
env_registry = registry or UserEnvironmentRegistry()
|
|
current_path = env_registry.get_value("Path")
|
|
env_registry.set_value("Path", merge_path_entries(current_path, update.path_entries))
|
|
for name, value in update.env_vars.items():
|
|
env_registry.set_value(name, value)
|
|
env_registry.broadcast_change()
|
|
|
|
|
|
def _normalize_path(value: str) -> str:
|
|
return os.path.normcase(os.path.normpath(os.path.expandvars(value.strip().strip('"'))))
|