diff --git a/AGENTS.md b/AGENTS.md index 52f67d0..b476180 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -137,10 +137,14 @@ Tests use `conftest.py` fixtures for consistent setup: ### Environment Setup ```bash -pipx install pypeline-runner # Install pypeline globally -pypeline run # Creates .venv and installs dependencies +pipx install pypeline-runner # Install pypeline globally +pypeline run --step CreateVEnv # Creates .venv and installs project in dev mode +pypeline run # Run full pipeline (tests, linting, etc.) ``` +> [!IMPORTANT] +> **When developing new built-in steps**: If you add a new step module under `src/pypeline/steps/` and reference it via `module:` in `pypeline.yaml`, you must first run `pypeline run --step CreateVEnv` to create/update the virtual environment. This installs the project in development mode so the new module becomes importable. Running `pypeline run` directly will fail because it tries to load all step modules before executing any step, and new modules won't be found in the globally installed package. + ### Running Tasks Use VS Code tasks or Python commands: diff --git a/pypeline.yaml b/pypeline.yaml index bb263d9..178f528 100644 --- a/pypeline.yaml +++ b/pypeline.yaml @@ -4,6 +4,10 @@ pipeline: config: package_manager: poetry==2.2.1 python_version: 3.11 + - step: ResolveLicenseServer + module: pypeline.steps.resolve_license + config: + license_config: license_servers.json - step: GenerateEnvSetupScript module: pypeline.steps.env_setup_script - step: PreCommit diff --git a/src/pypeline/steps/resolve_license.py b/src/pypeline/steps/resolve_license.py new file mode 100644 index 0000000..267c747 --- /dev/null +++ b/src/pypeline/steps/resolve_license.py @@ -0,0 +1,106 @@ +import json +import time +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Dict, List, Optional + +from py_app_dev.core.config import BaseConfigJSONMixin +from py_app_dev.core.exceptions import UserNotificationException +from py_app_dev.core.logging import logger + +from pypeline.domain.execution_context import ExecutionContext +from pypeline.domain.pipeline import PipelineStep + + +@dataclass +class ResolveLicenseExecutionInfo(BaseConfigJSONMixin): + env_vars: Dict[str, Any] = field(default_factory=dict) + + def to_json_file(self, file_path: Path) -> None: + file_path.parent.mkdir(parents=True, exist_ok=True) + super().to_json_file(file_path) + + +@dataclass +class ResolveLicenseConfig(BaseConfigJSONMixin): + license_config: str = "license_servers.json" + + +class ResolveLicenseServer(PipelineStep[ExecutionContext]): + def __init__(self, execution_context: ExecutionContext, group_name: Optional[str], config: Optional[Dict[str, Any]] = None) -> None: + super().__init__(execution_context, group_name, config) + self.user_config = ResolveLicenseConfig.from_dict(config) if config else ResolveLicenseConfig() + self.execution_info = ResolveLicenseExecutionInfo() + + @property + def _license_config_file(self) -> Path: + return self.execution_context.project_root_dir / self.user_config.license_config + + @property + def _execution_info_file(self) -> Path: + return self.output_dir / "resolve_license_exec_info.json" + + def _detect_timezone(self) -> str: + return time.tzname[0] + + def _resolve_site(self, timezones: Dict[str, List[str]], tz_name: str) -> Optional[str]: + for site, tz_names in timezones.items(): + if tz_name in tz_names: + return site + return None + + def run(self) -> int: + if not self._license_config_file.exists(): + raise UserNotificationException(f"License config file not found: {self._license_config_file}") + + with self._license_config_file.open("r") as f: + license_config = json.load(f) + + tz_name = self._detect_timezone() + logger.info(f"Detected timezone: {tz_name}") + + timezones: Dict[str, List[str]] = license_config.get("timezones", {}) + servers: Dict[str, Dict[str, Any]] = license_config.get("servers", {}) + + site = self._resolve_site(timezones, tz_name) + + if site: + logger.info(f"Matched site: {site}") + else: + if "default" in servers: + site = "default" + logger.info(f"No site matched for timezone '{tz_name}', using default") + else: + raise UserNotificationException( + f"No license server configured for timezone '{tz_name}' and no default configured in {self._license_config_file}" + ) + + env_vars = servers.get(site, {}) + if not env_vars: + raise UserNotificationException(f"No server configuration found for site '{site}' in {self._license_config_file}") + + self.execution_info.env_vars.update(env_vars) + self.execution_info.to_json_file(self._execution_info_file) + self.execution_context.add_env_vars(env_vars) + + for key, value in env_vars.items(): + logger.info(f"Set {key}={value}") + + return 0 + + def get_inputs(self) -> List[Path]: + if self._license_config_file.exists(): + return [self._license_config_file] + return [] + + def get_outputs(self) -> List[Path]: + return [self._execution_info_file] + + def get_name(self) -> str: + return self.__class__.__name__ + + def update_execution_context(self) -> None: + if self._execution_info_file.exists(): + execution_info = ResolveLicenseExecutionInfo.from_json_file(self._execution_info_file) + if execution_info.env_vars: + self.execution_context.add_env_vars(execution_info.env_vars) diff --git a/tests/test_resolve_license.py b/tests/test_resolve_license.py new file mode 100644 index 0000000..512b09b --- /dev/null +++ b/tests/test_resolve_license.py @@ -0,0 +1,134 @@ +import json +from pathlib import Path +from unittest.mock import Mock, patch + +import pytest +from py_app_dev.core.exceptions import UserNotificationException + +from pypeline.domain.execution_context import ExecutionContext +from pypeline.steps.env_setup_script import GenerateEnvSetupScript +from pypeline.steps.resolve_license import ResolveLicenseServer + + +@pytest.fixture +def license_config(execution_context: Mock) -> Path: + config = { + "timezones": { + "germany": ["W. Europe Standard Time", "Central European Standard Time", "CET"], + "pune": ["India Standard Time"], + "shanghai": ["China Standard Time"], + "romania": ["GTB Standard Time", "E. Europe Standard Time", "EET"], + }, + "servers": { + "germany": {"TASKING_LICENSE_SERVER": "1234@license-de.example.com"}, + "pune": {"TASKING_LICENSE_SERVER": "1234@license-pune.example.com"}, + "shanghai": {"TASKING_LICENSE_SERVER": "1234@license-sha.example.com"}, + "romania": {"TASKING_LICENSE_SERVER": "1234@license-ro.example.com"}, + "default": {"TASKING_LICENSE_SERVER": "1234@license-de.example.com"}, + }, + } + config_file = execution_context.project_root_dir / "license_servers.json" + config_file.write_text(json.dumps(config)) + return config_file + + +@patch.object(ResolveLicenseServer, "_detect_timezone") +def test_resolve_matching_site(mock_tz: Mock, execution_context: Mock, license_config: Path) -> None: + mock_tz.return_value = "India Standard Time" + execution_context.env_vars = {} + + step = ResolveLicenseServer(execution_context, None, {"license_config": "license_servers.json"}) + step.run() + + execution_context.add_env_vars.assert_called_once_with({"TASKING_LICENSE_SERVER": "1234@license-pune.example.com"}) + + +@patch.object(ResolveLicenseServer, "_detect_timezone") +def test_resolve_fallback_to_default(mock_tz: Mock, execution_context: Mock, license_config: Path) -> None: + mock_tz.return_value = "Eastern Standard Time" + execution_context.env_vars = {} + + step = ResolveLicenseServer(execution_context, None, {"license_config": "license_servers.json"}) + step.run() + + execution_context.add_env_vars.assert_called_once_with({"TASKING_LICENSE_SERVER": "1234@license-de.example.com"}) + + +@patch.object(ResolveLicenseServer, "_detect_timezone") +def test_resolve_no_match_no_default_raises_error(mock_tz: Mock, execution_context: Mock) -> None: + config = { + "timezones": {"germany": ["W. Europe Standard Time"]}, + "servers": {"germany": {"TASKING_LICENSE_SERVER": "1234@license-de.example.com"}}, + } + config_file = execution_context.project_root_dir / "license_servers.json" + config_file.write_text(json.dumps(config)) + + mock_tz.return_value = "Eastern Standard Time" + execution_context.env_vars = {} + + step = ResolveLicenseServer(execution_context, None, {"license_config": "license_servers.json"}) + + with pytest.raises(UserNotificationException, match="No license server configured for timezone"): + step.run() + + +def test_resolve_config_file_not_found_raises_error(execution_context: Mock) -> None: + execution_context.env_vars = {} + + step = ResolveLicenseServer(execution_context, None, {"license_config": "nonexistent.json"}) + + with pytest.raises(UserNotificationException, match="License config file not found"): + step.run() + + +@patch.object(ResolveLicenseServer, "_detect_timezone") +def test_resolve_germany_with_cet(mock_tz: Mock, execution_context: Mock, license_config: Path) -> None: + mock_tz.return_value = "CET" + execution_context.env_vars = {} + + step = ResolveLicenseServer(execution_context, None, {"license_config": "license_servers.json"}) + step.run() + + execution_context.add_env_vars.assert_called_once_with({"TASKING_LICENSE_SERVER": "1234@license-de.example.com"}) + + +@patch.object(ResolveLicenseServer, "_detect_timezone") +def test_update_execution_context_restores_cached(mock_tz: Mock, execution_context: Mock, license_config: Path) -> None: + mock_tz.return_value = "India Standard Time" + execution_context.env_vars = {} + + step = ResolveLicenseServer(execution_context, None, {"license_config": "license_servers.json"}) + step.run() + + execution_context.add_env_vars.reset_mock() + + step2 = ResolveLicenseServer(execution_context, None, {"license_config": "license_servers.json"}) + step2.update_execution_context() + + execution_context.add_env_vars.assert_called_once_with({"TASKING_LICENSE_SERVER": "1234@license-pune.example.com"}) + + +@patch("platform.system") +@patch.object(ResolveLicenseServer, "_detect_timezone") +def test_resolved_license_flows_into_env_setup_scripts(mock_tz: Mock, mock_platform: Mock, tmp_path: Path) -> None: + mock_tz.return_value = "India Standard Time" + mock_platform.return_value = "Windows" + + config = { + "timezones": {"pune": ["India Standard Time"]}, + "servers": {"pune": {"TASKING_LICENSE_SERVER": "1234@license-pune.example.com"}}, + } + (tmp_path / "license_servers.json").write_text(json.dumps(config)) + + context = ExecutionContext(project_root_dir=tmp_path) + + resolve_step = ResolveLicenseServer(context, None, {"license_config": "license_servers.json"}) + resolve_step.run() + + env_step = GenerateEnvSetupScript(context, None) + env_step.run() + + for script_name in ["env_setup.bat", "env_setup.ps1"]: + content = (env_step.output_dir / script_name).read_text() + assert "TASKING_LICENSE_SERVER" in content + assert "1234@license-pune.example.com" in content