From d43862e41aff8183d443149a716380f6db7eaa12 Mon Sep 17 00:00:00 2001 From: Amir Husayn Panahifar Date: Sat, 12 Sep 2026 13:47:31 +0330 Subject: [PATCH] feat(core): add sshctl settings and host models --- src/sshctl/core/__init__.py | 5 ++ src/sshctl/core/_version.py | 1 + src/sshctl/core/config.py | 42 +++++++++++++ src/sshctl/core/exceptions.py | 30 ++++++++++ src/sshctl/core/logutil.py | 26 ++++++++ src/sshctl/core/models.py | 108 ++++++++++++++++++++++++++++++++++ 6 files changed, 212 insertions(+) create mode 100644 src/sshctl/core/__init__.py create mode 100644 src/sshctl/core/_version.py create mode 100644 src/sshctl/core/config.py create mode 100644 src/sshctl/core/exceptions.py create mode 100644 src/sshctl/core/logutil.py create mode 100644 src/sshctl/core/models.py diff --git a/src/sshctl/core/__init__.py b/src/sshctl/core/__init__.py new file mode 100644 index 0000000..31fac35 --- /dev/null +++ b/src/sshctl/core/__init__.py @@ -0,0 +1,5 @@ +from __future__ import annotations + +from sshctl.core._version import __version__ + +__all__ = ["__version__"] diff --git a/src/sshctl/core/_version.py b/src/sshctl/core/_version.py new file mode 100644 index 0000000..485f44a --- /dev/null +++ b/src/sshctl/core/_version.py @@ -0,0 +1 @@ +__version__ = "0.1.1" diff --git a/src/sshctl/core/config.py b/src/sshctl/core/config.py new file mode 100644 index 0000000..057f441 --- /dev/null +++ b/src/sshctl/core/config.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +from pathlib import Path + +from pydantic import Field +from pydantic_settings import BaseSettings + + +class SshctlConfig(BaseSettings): + model_config = {"env_prefix": "SSHCTL_"} + + ssh_dir: Path = Field(default=Path("~/.ssh").expanduser()) + config_file: str = "config" + verbose: bool = False + timeout: int = 10 + + def ssh_config_path(self) -> Path: + return self.ssh_dir / self.config_file + + def conf_d_path(self) -> Path: + return self.ssh_dir / "conf.d" + + def keys_d_path(self) -> Path: + return self.ssh_dir / "keys.d" + + def control_d_path(self) -> Path: + return self.ssh_dir / "control.d" + + def scripts_path(self) -> Path: + return self.ssh_dir / "scripts" + + def git_path(self) -> Path: + return self.ssh_dir / ".git" + + +def load_config(path: Path | None = None) -> SshctlConfig: + if path and path.exists(): + import yaml + + raw = yaml.safe_load(path.read_text()) or {} + return SshctlConfig(**raw) + return SshctlConfig() diff --git a/src/sshctl/core/exceptions.py b/src/sshctl/core/exceptions.py new file mode 100644 index 0000000..8cd2e38 --- /dev/null +++ b/src/sshctl/core/exceptions.py @@ -0,0 +1,30 @@ +class SshctlError(Exception): + pass + + +class ConfigError(SshctlError): + pass + + +class ParseError(SshctlError): + pass + + +class ResolverError(SshctlError): + pass + + +class CheckError(SshctlError): + pass + + +class ExecError(SshctlError): + pass + + +class PermissionFixError(SshctlError): + pass + + +class ValidationError(SshctlError): + pass diff --git a/src/sshctl/core/logutil.py b/src/sshctl/core/logutil.py new file mode 100644 index 0000000..5a931fd --- /dev/null +++ b/src/sshctl/core/logutil.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +import structlog +from structlog.stdlib import BoundLogger + +_LOG: BoundLogger | None = None + + +def setup_logging(*, verbose: bool = False) -> None: + structlog.configure( + processors=[ + structlog.stdlib.add_log_level, + structlog.dev.ConsoleRenderer() if verbose else structlog.processors.JSONRenderer(), + ], + wrapper_class=structlog.stdlib.BoundLogger, + context_class=dict, + logger_factory=structlog.PrintLoggerFactory(), + cache_logger_on_first_use=True, + ) + + +def get_logger(name: str) -> BoundLogger: + global _LOG + if _LOG is None: + _LOG = structlog.get_logger(name) + return _LOG diff --git a/src/sshctl/core/models.py b/src/sshctl/core/models.py new file mode 100644 index 0000000..d363fc7 --- /dev/null +++ b/src/sshctl/core/models.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Literal, cast + + +@dataclass +class HostOption: + name: str + value: str + + +@dataclass +class HostBlock: + alias: str + options: dict[str, str] = field(default_factory=dict) + + def get(self, key: str, default: str | None = None) -> str | None: + return self.options.get(key, default) + + def set(self, key: str, value: str) -> None: + self.options[key] = value + + @property + def hostname(self) -> str | None: + return self.get("HostName") + + @property + def user(self) -> str | None: + return self.get("User") + + @property + def port(self) -> int | None: + val = self.get("Port") + return int(val) if val else None + + @property + def proxy_jump(self) -> str | None: + return self.get("ProxyJump") + + @property + def identity_file(self) -> str | None: + return self.get("IdentityFile") + + +@dataclass +class SshConfig: + blocks: list[HostBlock] = field(default_factory=list) + raw_includes: list[str] = field(default_factory=list) + + def find(self, alias: str) -> HostBlock | None: + for block in self.blocks: + if block.alias == alias: + return block + return None + + def find_all(self, pattern: str) -> list[HostBlock]: + from fnmatch import fnmatch + + return [b for b in self.blocks if fnmatch(b.alias, pattern)] + + @property + def aliases(self) -> list[str]: + return [b.alias for b in self.blocks if not _is_pattern(b.alias)] + + +def _is_pattern(alias: str) -> bool: + return any(c in alias for c in "*?[]") + + +Universe = Literal["internal", "external"] + + +@dataclass +class ZonePath: + universe: Universe + sector: str | None = None + cluster: str | None = None + zone: str | None = None + vlan: str | None = None + + def to_rel_path(self) -> Path: + parts: list[str] = ["conf.d", self.universe] + if self.sector: + parts.append(self.sector) + if self.cluster: + parts.append(self.cluster) + if self.zone: + parts.append(self.zone) + if self.vlan: + parts.append(self.vlan) + return Path(*parts) + + @classmethod + def from_rel_path(cls, path: Path) -> ZonePath: + parts = path.parts + if len(parts) < 2 or parts[0] != "conf.d": + raise ValueError(f"Invalid zone path: {path}") + raw_universe = parts[1] + if raw_universe not in ("internal", "external"): + raise ValueError(f"Unknown universe: {raw_universe}") + universe = cast(Literal["internal", "external"], raw_universe) + sector = parts[2] if len(parts) > 2 else None + cluster = parts[3] if len(parts) > 3 else None + zone = parts[4] if len(parts) > 4 else None + vlan = parts[5] if len(parts) > 5 else None + return cls(universe=universe, sector=sector, cluster=cluster, zone=zone, vlan=vlan)