import os import shlex import stat from pathlib import Path from core.base import BaseWorker from core.schemas.results import AuditFindings, AuditResults from core.schemas.status import AuditSeverity, AuditStatus from .config import SSHModuleConfig class SSHSecurityWorker(BaseWorker): config_model = SSHModuleConfig def __init__(self, config: SSHModuleConfig | None = None) -> None: super().__init__(config) self.module_config = config or SSHModuleConfig() def run(self) -> AuditResults: if not self.module_config.enabled: return AuditResults( status=AuditStatus.SKIPPED, findings=[], risk_level=0.0, ) if self.module_config.test_failure: findings = [ AuditFindings( name="im fucking dumbass rragh", description="im testing!.", severity=AuditSeverity.CRITICAL, ) ] return AuditResults( status=AuditStatus.PASS if not findings else AuditStatus.FAIL, findings=findings, risk_level=self._calculate_risk(findings), ) if os.name != "posix": return AuditResults( status=AuditStatus.SKIPPED, findings=[], risk_level=0.0, ) sshd_config = Path("/etc/ssh/sshd_config") if not sshd_config.exists(): return AuditResults( status=AuditStatus.SKIPPED, findings=[ AuditFindings( name="sshd_config_not_found", description="SSH server config /etc/ssh/sshd_config was not found.", severity=AuditSeverity.LOW, ) ], risk_level=0.1, ) findings: list[AuditFindings] = [] resolved_config = self._read_merged_config(sshd_config) findings.extend(self._check_config_directives(resolved_config)) findings.extend(self._check_config_permissions(sshd_config)) findings.extend(self._check_host_key_permissions(Path("/etc/ssh"))) return AuditResults( status=AuditStatus.PASS if not findings else AuditStatus.FAIL, findings=findings, risk_level=self._calculate_risk(findings), ) def _read_merged_config(self, root_config: Path) -> dict[str, str]: config_files = [root_config] include_dir = root_config.parent / "sshd_config.d" if include_dir.exists() and include_dir.is_dir(): config_files.extend(sorted(include_dir.glob("*.conf"))) directives: dict[str, str] = {} for config_file in config_files: for raw_line in config_file.read_text(encoding="utf-8").splitlines(): line = raw_line.strip() if not line or line.startswith("#"): continue line = line.split("#", 1)[0].strip() if not line: continue parts = shlex.split(line, comments=False, posix=True) if len(parts) < 2: # noqa: PLR2004 continue directive = parts[0].lower() value = " ".join(parts[1:]) directives[directive] = value return directives def _check_config_directives(self, config: dict[str, str]) -> list[AuditFindings]: findings: list[AuditFindings] = [] if config.get("permitrootlogin", "prohibit-password").lower() not in { "no", "forced-commands-only", }: findings.append( AuditFindings( name="permit_root_login_enabled", description="Directive PermitRootLogin should be set to 'no' or 'forced-commands-only'.", severity=AuditSeverity.CRITICAL, ) ) if config.get("passwordauthentication", "yes").lower() != "no": findings.append( AuditFindings( name="password_authentication_enabled", description="Directive PasswordAuthentication should be disabled to prefer key-based access.", severity=AuditSeverity.HIGH, ) ) if config.get("pubkeyauthentication", "yes").lower() != "yes": findings.append( AuditFindings( name="pubkey_authentication_disabled", description="Directive PubkeyAuthentication should be enabled.", severity=AuditSeverity.HIGH, ) ) if config.get("x11forwarding", "no").lower() != "no": findings.append( AuditFindings( name="x11_forwarding_enabled", description="Directive X11Forwarding should usually be disabled on hardened servers.", severity=AuditSeverity.MEDIUM, ) ) max_auth_tries = config.get("maxauthtries", "6") if max_auth_tries.isdigit() and int(max_auth_tries) > 4: # noqa: PLR2004 findings.append( AuditFindings( name="max_auth_tries_too_high", description="Directive MaxAuthTries should be 4 or less.", severity=AuditSeverity.MEDIUM, ) ) if config.get("permitemptypasswords", "no").lower() != "no": findings.append( AuditFindings( name="empty_passwords_permitted", description="Directive PermitEmptyPasswords must be disabled.", severity=AuditSeverity.CRITICAL, ) ) return findings def _check_config_permissions(self, config_path: Path) -> list[AuditFindings]: findings: list[AuditFindings] = [] file_stat = config_path.stat() mode = stat.S_IMODE(file_stat.st_mode) if file_stat.st_uid != 0: findings.append( AuditFindings( name="sshd_config_not_owned_by_root", description="File /etc/ssh/sshd_config should be owned by root.", severity=AuditSeverity.HIGH, ) ) if mode & stat.S_IWGRP or mode & stat.S_IWOTH: findings.append( AuditFindings( name="sshd_config_writable_by_non_root", description="File /etc/ssh/sshd_config must not be writable by group or others.", severity=AuditSeverity.HIGH, ) ) return findings def _check_host_key_permissions(self, ssh_dir: Path) -> list[AuditFindings]: findings: list[AuditFindings] = [] for key_path in sorted(ssh_dir.glob("ssh_host_*_key")): file_stat = key_path.stat() mode = stat.S_IMODE(file_stat.st_mode) if file_stat.st_uid != 0: findings.append( AuditFindings( name=f"{key_path.name}_not_owned_by_root", description=f"Private host key {key_path} should be owned by root.", severity=AuditSeverity.HIGH, ) ) if mode & (stat.S_IRWXG | stat.S_IRWXO): findings.append( AuditFindings( name=f"{key_path.name}_permissions_too_open", description=f"Private host key {key_path} should not be accessible to group or others.", severity=AuditSeverity.CRITICAL, ) ) return findings def _calculate_risk(self, findings: list[AuditFindings]) -> float: if not findings: return 0.0 severity_scores = { AuditSeverity.CRITICAL: 0.45, AuditSeverity.HIGH: 0.3, AuditSeverity.MEDIUM: 0.2, AuditSeverity.LOW: 0.1, } risk = sum(severity_scores[finding.severity] for finding in findings) return min(1.0, risk)