diff --git a/README.md b/README.md index 7274a70..efb11ef 100644 --- a/README.md +++ b/README.md @@ -10,6 +10,7 @@ This repository provides a production-oriented Samba file server container that - `\\server\Data` -> `/data/groups/data` - `\\server\FSLogix` -> `/data/fslogix` - FS_* groups are projected as folders inside the Data share (`/data/groups/data/`). +- Data folder ACLs expand nested AD group membership recursively and detect group cycles. - Group records are persisted in SQLite at `/state/shares.db`. - Group folders are name-based while active and moved to archive on deactivation: - active: `/data/groups/data/` @@ -167,6 +168,7 @@ Kerberos requires close time alignment. - Path: `/data/groups/data` - Contains one folder per active `FS_*` AD group. - Root is discoverable as one share, while access to each group folder is enforced via POSIX/ACL group permissions. +- `FS_*` groups may contain other AD groups; reconciliation recursively grants ACLs to nested groups and logs detected cycles. - No guest access. ### FSLogix @@ -281,6 +283,17 @@ docker compose exec samba sh -lc 'tail -n 200 /var/log/backup.log' docker compose exec samba tail -n 100 /var/log/reconcile.log ``` +### Nested Data group access fails + +- Check reconciliation logs for detected group cycles or unresolved nested members. +- Verify winbind can resolve every nested group to a local GID: + + ```bash + docker compose exec samba getent group 'EXAMPLE\NestedGroup' + docker compose exec samba id 'EXAMPLE\alice' + docker compose exec samba getfacl /data/groups/data/ + ``` + ### `acl_xattr.so` or `full_audit.so` module load error - If logs show `Error loading module .../vfs/acl_xattr.so` (or `full_audit.so`), your running image is missing Samba VFS modules. diff --git a/app/reconcile_shares.py b/app/reconcile_shares.py index 0fec2ee..f44d5bb 100755 --- a/app/reconcile_shares.py +++ b/app/reconcile_shares.py @@ -12,7 +12,8 @@ import subprocess import sys import tempfile import uuid -from typing import Dict, List, Optional, Tuple +from dataclasses import dataclass, field +from typing import Callable, Dict, Iterable, List, Optional, Set, Tuple DB_PATH = "/state/shares.db" @@ -26,6 +27,25 @@ LDAP_FILTER = "(&(objectClass=group)(sAMAccountName=FS_*))" GROUP_PREFIXES = ("FS_",) GROUP_TITLE_ATTRS = ("displayname", "name", "cn") USER_STATUS_FILTER = "(&(objectClass=user)(!(objectClass=computer))(sAMAccountName=*))" +GROUP_SEARCH_ATTRS = [ + "objectGUID", + "sAMAccountName", + "displayName", + "name", + "cn", + "distinguishedName", + "member", + "objectClass", +] +PRINCIPAL_SEARCH_ATTRS = [ + "objectGUID", + "sAMAccountName", + "distinguishedName", + "member", + "objectClass", +] +LDAP_DN_LOOKUP_BATCH_SIZE = 25 +SETFACL_ENTRY_CHUNK_SIZE = 40 REQUIRED_ENV = ["REALM", "WORKGROUP", "DOMAIN"] ATTR_RE = re.compile(r"^([^:]+)(::?)\s*(.*)$") @@ -46,6 +66,20 @@ UAC_LOCKOUT = 0x0010 AD_NEVER_EXPIRES_VALUES = {0, 9223372036854775807} MAX_GROUP_FOLDER_NAME = 120 +LdapValue = Tuple[str, bool] +LdapEntry = Dict[str, List[LdapValue]] +Principal = Dict[str, object] +PrincipalLookup = Callable[[List[str]], Dict[str, Principal]] + + +@dataclass +class MembershipExpansion: + group_sams: List[str] = field(default_factory=list) + terminal_user_count: int = 0 + unresolved_dns: List[str] = field(default_factory=list) + ignored_dns: List[str] = field(default_factory=list) + cycle_paths: List[List[str]] = field(default_factory=list) + def now_utc() -> str: return dt.datetime.now(dt.timezone.utc).isoformat(timespec="seconds") @@ -88,18 +122,28 @@ def run_command(command: List[str], check: bool = True) -> subprocess.CompletedP return result -def parse_ldap_entries(output: str) -> List[Dict[str, Tuple[str, bool]]]: - entries: List[Dict[str, Tuple[str, bool]]] = [] - current: Dict[str, Tuple[str, bool]] = {} +def unfold_ldap_lines(output: str) -> List[str]: + lines: List[str] = [] + for raw_line in output.splitlines(): + if raw_line.startswith(" ") and lines: + lines[-1] += raw_line[1:] + continue + lines.append(raw_line) + return lines - for line in output.splitlines(): + +def parse_ldap_entries(output: str) -> List[LdapEntry]: + entries: List[LdapEntry] = [] + current: LdapEntry = {} + + for line in unfold_ldap_lines(output): stripped = line.strip() if not stripped: if current: entries.append(current) current = {} continue - if stripped.startswith("#") or stripped.startswith("dn:"): + if stripped.startswith("#"): continue match = ATTR_RE.match(stripped) @@ -107,7 +151,7 @@ def parse_ldap_entries(output: str) -> List[Dict[str, Tuple[str, bool]]]: continue key, delimiter, value = match.groups() - current[key.lower()] = (value, delimiter == "::") + current.setdefault(key.lower(), []).append((value, delimiter == "::")) if current: entries.append(current) @@ -115,6 +159,40 @@ def parse_ldap_entries(output: str) -> List[Dict[str, Tuple[str, bool]]]: return entries +def decode_ldap_text(value: str, is_b64: bool) -> str: + if not is_b64: + return value.strip() + try: + return base64.b64decode(value).decode("utf-8").strip() + except (ValueError, UnicodeDecodeError): + return value.strip() + + +def ldap_values(entry: LdapEntry, attr: str) -> List[str]: + attr_key = attr.lower() + values: List[str] = [] + for entry_key, entry_values in entry.items(): + if entry_key != attr_key and not entry_key.startswith(f"{attr_key};"): + continue + values.extend( + decode_ldap_text(value, is_b64) for value, is_b64 in entry_values + ) + return values + + +def ldap_first(entry: LdapEntry, attr: str) -> Optional[str]: + values = ldap_values(entry, attr) + for value in values: + if value: + return value + return None + + +def ldap_first_raw(entry: LdapEntry, attr: str) -> Optional[LdapValue]: + values = entry.get(attr.lower(), []) + return values[0] if values else None + + def derive_share_name(sam_account_name: str) -> Optional[str]: for prefix in GROUP_PREFIXES: if sam_account_name.startswith(prefix): @@ -123,47 +201,80 @@ def derive_share_name(sam_account_name: str) -> Optional[str]: return None -def derive_group_title(entry: Dict[str, Tuple[str, bool]]) -> Optional[str]: +def derive_group_title(entry: LdapEntry) -> Optional[str]: for attr in GROUP_TITLE_ATTRS: - if attr in entry: - value = entry[attr][0].strip() - if value: - return value + value = ldap_first(entry, attr) + if value: + return value return None -def parse_groups_from_ldap_output(output: str) -> List[Dict[str, str]]: - entries = parse_ldap_entries(output) +def entry_dn(entry: LdapEntry) -> str: + return ldap_first(entry, "distinguishedName") or ldap_first(entry, "dn") or "" - groups: List[Dict[str, str]] = [] + +def parse_principal_from_entry(entry: LdapEntry) -> Optional[Principal]: + dn = entry_dn(entry) + if not dn: + return None + + guid = "" + guid_raw = ldap_first_raw(entry, "objectGUID") + if guid_raw is not None: + try: + guid = parse_guid(guid_raw[0].strip(), guid_raw[1]) + except ValueError: + guid = "" + + object_classes = {value.lower() for value in ldap_values(entry, "objectClass")} + return { + "dn": dn, + "objectGUID": guid, + "samAccountName": ldap_first(entry, "sAMAccountName") or "", + "memberDns": ldap_values(entry, "member"), + "objectClasses": object_classes, + } + + +def parse_groups_from_ldap_entries(entries: List[LdapEntry]) -> List[Dict[str, object]]: + groups: List[Dict[str, object]] = [] for entry in entries: - if "objectguid" not in entry or "samaccountname" not in entry: + guid_raw = ldap_first_raw(entry, "objectGUID") + sam = ldap_first(entry, "sAMAccountName") + if guid_raw is None or not sam: continue - sam_value, _ = entry["samaccountname"] - sam = sam_value.strip() share_name = derive_group_title(entry) or derive_share_name(sam) if not share_name: continue - guid_value, is_b64 = entry["objectguid"] - guid = parse_guid(guid_value.strip(), is_b64) + guid = parse_guid(guid_raw[0].strip(), guid_raw[1]) + object_classes = {value.lower() for value in ldap_values(entry, "objectClass")} + if not object_classes: + object_classes = {"group"} groups.append( { "objectGUID": guid, "samAccountName": sam, "shareName": share_name, + "distinguishedName": entry_dn(entry), + "memberDns": ldap_values(entry, "member"), + "objectClasses": object_classes, } ) - deduped: Dict[str, Dict[str, str]] = {} + deduped: Dict[str, Dict[str, object]] = {} for group in groups: - deduped[group["objectGUID"]] = group + deduped[str(group["objectGUID"])] = group return list(deduped.values()) +def parse_groups_from_ldap_output(output: str) -> List[Dict[str, object]]: + return parse_groups_from_ldap_entries(parse_ldap_entries(output)) + + def sanitize_group_folder_name(raw_name: str) -> str: candidate = GROUP_FOLDER_INVALID_RE.sub("_", raw_name.strip()) candidate = candidate.strip().strip(".") @@ -231,30 +342,46 @@ def next_available_path(path: str) -> str: index += 1 -def fetch_groups_via_net_ads() -> List[Dict[str, str]]: +def chunked(values: List[str], size: int) -> Iterable[List[str]]: + for index in range(0, len(values), size): + yield values[index : index + size] + + +def escape_ldap_filter_value(value: str) -> str: + replacements = { + "\\": "\\5c", + "*": "\\2a", + "(": "\\28", + ")": "\\29", + "\x00": "\\00", + } + return "".join(replacements.get(char, char) for char in value) + + +def build_distinguished_name_filter(dns: List[str]) -> str: + clauses = [f"(distinguishedName={escape_ldap_filter_value(dn)})" for dn in dns] + if len(clauses) == 1: + return clauses[0] + return f"(|{''.join(clauses)})" + + +def search_entries_via_net_ads( + filter_expr: str, attributes: List[str] +) -> List[LdapEntry]: result = run_command( - [ - "net", - "ads", - "search", - "-P", - LDAP_FILTER, - "objectGUID", - "sAMAccountName", - "displayName", - "name", - "cn", - ], + ["net", "ads", "search", "-P", filter_expr, *attributes], check=False, ) if result.returncode != 0: raise RuntimeError( result.stderr.strip() or result.stdout.strip() or "net ads search failed" ) - return parse_groups_from_ldap_output(result.stdout) + return parse_ldap_entries(result.stdout) -def fetch_groups_via_ldap_bind() -> List[Dict[str, str]]: +def search_entries_via_ldap_bind( + filter_expr: str, attributes: List[str] +) -> List[LdapEntry]: realm = os.environ["REALM"] join_user = os.getenv("JOIN_USER", "") join_password = os.getenv("JOIN_PASSWORD", "") @@ -288,26 +415,207 @@ def fetch_groups_via_ldap_bind() -> List[Dict[str, str]]: pw_file, "-b", base_dn, - LDAP_FILTER, - "objectGUID", - "sAMAccountName", - "displayName", - "name", - "cn", + filter_expr, + *attributes, ] ) - return parse_groups_from_ldap_output(result.stdout) + return parse_ldap_entries(result.stdout) finally: if pw_file and os.path.exists(pw_file): os.remove(pw_file) -def fetch_fileshare_groups() -> List[Dict[str, str]]: +def search_directory_entries(filter_expr: str, attributes: List[str]) -> List[LdapEntry]: try: - return fetch_groups_via_net_ads() + return search_entries_via_net_ads(filter_expr, attributes) except Exception as net_exc: # pylint: disable=broad-except log(f"net ads search failed, falling back to LDAP bind: {net_exc}") - return fetch_groups_via_ldap_bind() + return search_entries_via_ldap_bind(filter_expr, attributes) + + +def fetch_fileshare_groups() -> List[Dict[str, object]]: + entries = search_directory_entries(LDAP_FILTER, GROUP_SEARCH_ATTRS) + return parse_groups_from_ldap_entries(entries) + + +def normalize_dn(dn: str) -> str: + return dn.strip().casefold() + + +def group_to_principal(group: Dict[str, object]) -> Principal: + object_classes = group.get("objectClasses") or {"group"} + return { + "dn": str(group.get("distinguishedName") or ""), + "objectGUID": str(group.get("objectGUID") or ""), + "samAccountName": str(group.get("samAccountName") or ""), + "memberDns": [str(dn) for dn in group.get("memberDns", []) if str(dn).strip()], + "objectClasses": {str(value).lower() for value in object_classes}, + } + + +def principal_display_name(principal: Principal) -> str: + sam = str(principal.get("samAccountName") or "").strip() + if sam: + return sam + return str(principal.get("dn") or "") + + +def is_group_principal(principal: Principal) -> bool: + return "group" in principal.get("objectClasses", set()) + + +def is_computer_principal(principal: Principal) -> bool: + return "computer" in principal.get("objectClasses", set()) + + +def is_user_principal(principal: Principal) -> bool: + return "user" in principal.get("objectClasses", set()) and not is_computer_principal( + principal + ) + + +def lookup_principals_by_dns(dns: List[str]) -> Dict[str, Principal]: + unique_dns: List[str] = [] + seen_dns: Set[str] = set() + for dn in dns: + key = normalize_dn(dn) + if not key or key in seen_dns: + continue + seen_dns.add(key) + unique_dns.append(dn) + + principals: Dict[str, Principal] = {} + for dn_chunk in chunked(unique_dns, LDAP_DN_LOOKUP_BATCH_SIZE): + entries = search_directory_entries( + build_distinguished_name_filter(dn_chunk), PRINCIPAL_SEARCH_ATTRS + ) + for entry in entries: + principal = parse_principal_from_entry(entry) + if principal is None: + continue + key = normalize_dn(str(principal.get("dn") or "")) + if key: + principals[key] = principal + + return principals + + +def load_principals_for_dns( + dns: List[str], cache: Dict[str, Principal], lookup_func: PrincipalLookup +) -> None: + missing_dns: List[str] = [] + seen_dns: Set[str] = set() + for dn in dns: + key = normalize_dn(dn) + if not key or key in cache or key in seen_dns: + continue + seen_dns.add(key) + missing_dns.append(dn) + + if not missing_dns: + return + + for lookup_key, principal in lookup_func(missing_dns).items(): + principal_key = normalize_dn(str(principal.get("dn") or lookup_key)) + if principal_key: + cache[principal_key] = principal + + +def expand_group_membership( + root_group: Dict[str, object], + principal_cache: Optional[Dict[str, Principal]] = None, + lookup_func: PrincipalLookup = lookup_principals_by_dns, +) -> MembershipExpansion: + cache = principal_cache if principal_cache is not None else {} + root_principal = group_to_principal(root_group) + root_key = normalize_dn(str(root_principal.get("dn") or "")) + if root_key: + cache[root_key] = root_principal + + expansion = MembershipExpansion() + group_sams_seen: Set[str] = set() + unresolved_seen: Set[str] = set() + ignored_seen: Set[str] = set() + user_dns_seen: Set[str] = set() + cycle_paths_seen: Set[str] = set() + visiting: Set[str] = set() + visited: Set[str] = set() + + def add_unresolved(dn: str) -> None: + key = normalize_dn(dn) + if key and key not in unresolved_seen: + unresolved_seen.add(key) + expansion.unresolved_dns.append(dn) + + def add_ignored(dn: str) -> None: + key = normalize_dn(dn) + if key and key not in ignored_seen: + ignored_seen.add(key) + expansion.ignored_dns.append(dn) + + def add_cycle(path: List[str]) -> None: + key = " -> ".join(path).casefold() + if key not in cycle_paths_seen: + cycle_paths_seen.add(key) + expansion.cycle_paths.append(path) + + def visit_group(principal: Principal, path: List[str]) -> None: + current_key = normalize_dn(str(principal.get("dn") or "")) + if current_key: + visiting.add(current_key) + + member_dns = [ + str(dn) for dn in principal.get("memberDns", []) if str(dn).strip() + ] + load_principals_for_dns(member_dns, cache, lookup_func) + + for member_dn in member_dns: + member_key = normalize_dn(member_dn) + if not member_key: + continue + + member = cache.get(member_key) + if member is None: + add_unresolved(member_dn) + continue + + member_label = principal_display_name(member) + if member_key in visiting: + add_cycle(path + [member_label]) + continue + + if is_group_principal(member): + member_sam = str(member.get("samAccountName") or "").strip() + if not member_sam: + add_unresolved(str(member.get("dn") or member_dn)) + else: + sam_key = member_sam.casefold() + if sam_key not in group_sams_seen: + group_sams_seen.add(sam_key) + expansion.group_sams.append(member_sam) + + if member_key not in visited: + visit_group(member, path + [member_label]) + continue + + if is_computer_principal(member): + add_ignored(str(member.get("dn") or member_dn)) + continue + + if is_user_principal(member): + if member_key not in user_dns_seen: + user_dns_seen.add(member_key) + expansion.terminal_user_count += 1 + continue + + add_ignored(str(member.get("dn") or member_dn)) + + if current_key: + visiting.discard(current_key) + visited.add(current_key) + + visit_group(root_principal, [principal_display_name(root_principal)]) + return expansion def windows_filetime_now() -> int: @@ -345,16 +653,13 @@ def fetch_non_login_users() -> set: now_filetime = windows_filetime_now() for entry in parse_ldap_entries(result.stdout): - if "samaccountname" not in entry: - continue - - username = entry["samaccountname"][0].strip().lower() + username = (ldap_first(entry, "sAMAccountName") or "").lower() if not username: continue - uac = parse_int(entry.get("useraccountcontrol", ("0", False))[0], 0) - account_expires = parse_int(entry.get("accountexpires", ("0", False))[0], 0) - lockout_time = parse_int(entry.get("lockouttime", ("0", False))[0], 0) + uac = parse_int(ldap_first(entry, "userAccountControl") or "0", 0) + account_expires = parse_int(ldap_first(entry, "accountExpires") or "0", 0) + lockout_time = parse_int(ldap_first(entry, "lockoutTime") or "0", 0) is_disabled = bool(uac & UAC_ACCOUNTDISABLE) is_locked = bool(uac & UAC_LOCKOUT) or lockout_time > 0 @@ -395,7 +700,7 @@ def ensure_group_path(path: str) -> None: os.chmod(path, 0o2770) -def reconcile_db(conn: sqlite3.Connection, ad_groups: List[Dict[str, str]]) -> None: +def reconcile_db(conn: sqlite3.Connection, ad_groups: List[Dict[str, object]]) -> None: timestamp = now_utc() seen = set() @@ -566,33 +871,78 @@ def resolve_gid_from_sid(sid: str) -> Optional[int]: return None +def resolve_group_gids_for_acl( + workgroup: str, + group_names: List[str], + resolver: Callable[[str, str], Optional[int]] = resolve_group_gid_flexible, +) -> Tuple[List[int], List[str]]: + gids: List[int] = [] + unresolved: List[str] = [] + seen_gids: Set[int] = set() + seen_names: Set[str] = set() + + for group_name in group_names: + name = group_name.strip() + name_key = name.casefold() + if not name or name_key in seen_names: + continue + seen_names.add(name_key) + + gid = resolver(workgroup, name) + if gid is None: + unresolved.append(name) + continue + if gid not in seen_gids: + seen_gids.add(gid) + gids.append(gid) + + return gids, unresolved + + +def apply_setfacl_entries(path: str, acl_entries: List[str]) -> None: + for acl_chunk in chunked(acl_entries, SETFACL_ENTRY_CHUNK_SIZE): + result = run_command(["setfacl", "-m", ",".join(acl_chunk), path], check=False) + if result.returncode != 0: + log( + f"setfacl failed for {path}: " + f"{result.stderr.strip() or result.stdout.strip()}" + ) + return + + def apply_group_permissions( - path: str, group_gid: int, admin_gid: Optional[int], is_dir: bool + path: str, + owner_group_gid: int, + acl_group_gids: List[int], + admin_gid: Optional[int], + is_dir: bool, ) -> None: if os.path.islink(path): return mode = 0o2770 if is_dir else 0o660 group_perms = "rwx" if is_dir else "rw-" + acl_gids: List[int] = [] + seen_gids: Set[int] = set() + for gid in [owner_group_gid, *acl_group_gids]: + if gid not in seen_gids: + seen_gids.add(gid) + acl_gids.append(gid) + if admin_gid is not None and admin_gid not in seen_gids: + acl_gids.append(admin_gid) - os.chown(path, 0, group_gid) + os.chown(path, 0, owner_group_gid) os.chmod(path, mode) run_command(["setfacl", "-b", path], check=False) - acl_entries = [f"g:{group_gid}:{group_perms}"] - if admin_gid is not None: - acl_entries.append(f"g:{admin_gid}:{group_perms}") + acl_entries = [f"g:{gid}:{group_perms}" for gid in acl_gids] + acl_entries.append(f"m:{group_perms}") if is_dir: - acl_entries.append(f"d:g:{group_gid}:rwx") - if admin_gid is not None: - acl_entries.append(f"d:g:{admin_gid}:rwx") + acl_entries.extend(f"d:g:{gid}:rwx" for gid in acl_gids) + acl_entries.append("d:m:rwx") - result = run_command(["setfacl", "-m", ",".join(acl_entries), path], check=False) - if result.returncode != 0: - log( - f"setfacl failed for {path}: {result.stderr.strip() or result.stdout.strip()}" - ) + apply_setfacl_entries(path, acl_entries) def apply_private_permissions( @@ -625,17 +975,30 @@ def apply_private_permissions( def enforce_group_tree_permissions( - root_path: str, group_gid: int, admin_gid: Optional[int] + root_path: str, + owner_group_gid: int, + acl_group_gids: List[int], + admin_gid: Optional[int], ) -> None: - apply_group_permissions(root_path, group_gid, admin_gid, is_dir=True) + apply_group_permissions( + root_path, owner_group_gid, acl_group_gids, admin_gid, is_dir=True + ) for current_root, dirnames, filenames in os.walk(root_path): for dirname in dirnames: apply_group_permissions( - os.path.join(current_root, dirname), group_gid, admin_gid, is_dir=True + os.path.join(current_root, dirname), + owner_group_gid, + acl_group_gids, + admin_gid, + is_dir=True, ) for filename in filenames: apply_group_permissions( - os.path.join(current_root, filename), group_gid, admin_gid, is_dir=False + os.path.join(current_root, filename), + owner_group_gid, + acl_group_gids, + admin_gid, + is_dir=False, ) @@ -799,7 +1162,15 @@ def sync_private_directories() -> None: enforce_private_tree_permissions(user_path, uid, user_gid, admin_gid) -def sync_dynamic_directory_permissions(conn: sqlite3.Connection) -> None: +def format_dn_list(dns: List[str], limit: int = 3) -> str: + shown = dns[:limit] + suffix = "" if len(dns) <= limit else f", ... +{len(dns) - limit} more" + return ", ".join(shown) + suffix + + +def sync_dynamic_directory_permissions( + conn: sqlite3.Connection, ad_groups: List[Dict[str, object]] +) -> None: workgroup = os.environ["WORKGROUP"] admin_group = os.getenv("DOMAIN_ADMINS_GROUP", "") admin_gid = None @@ -808,21 +1179,66 @@ def sync_dynamic_directory_permissions(conn: sqlite3.Connection) -> None: if admin_gid is None: admin_gid = resolve_gid_from_sid(os.getenv("DOMAIN_ADMINS_SID", "")) + ad_groups_by_guid = {str(group["objectGUID"]): group for group in ad_groups} + principal_cache: Dict[str, Principal] = {} + rows = conn.execute( - "SELECT samAccountName, path FROM shares WHERE isActive = 1" + "SELECT objectGUID, samAccountName, path FROM shares WHERE isActive = 1" ).fetchall() for row in rows: + guid = row["objectGUID"] sam = row["samAccountName"] path = row["path"] - os.makedirs(path, exist_ok=True) - os.chmod(path, 0o2770) - - gid = resolve_group_gid_flexible(workgroup, sam) - if gid is None: - log(f"Unable to resolve GID for {sam}; leaving existing ACLs") + ad_group = ad_groups_by_guid.get(guid) + if ad_group is None: + log(f"No AD data available for {sam}; leaving existing ACLs") continue - enforce_group_tree_permissions(path, gid, admin_gid) + try: + expansion = expand_group_membership(ad_group, principal_cache) + except Exception as exc: # pylint: disable=broad-except + log( + f"Unable to expand nested members for {sam}: " + f"{exc}; leaving existing ACLs" + ) + continue + + for cycle_path in expansion.cycle_paths: + log(f"Detected nested group cycle for {sam}: {' -> '.join(cycle_path)}") + + if expansion.ignored_dns: + log( + f"Ignoring computer/unsupported nested member(s) for {sam}: " + f"{format_dn_list(expansion.ignored_dns)}" + ) + + if expansion.unresolved_dns: + log( + f"Unable to resolve nested member(s) for {sam}: " + f"{format_dn_list(expansion.unresolved_dns)}; leaving existing ACLs" + ) + continue + + acl_group_names = [sam, *expansion.group_sams] + acl_group_gids, unresolved_groups = resolve_group_gids_for_acl( + workgroup, acl_group_names + ) + if unresolved_groups: + log( + f"Unable to resolve GID(s) for {sam}: " + f"{', '.join(unresolved_groups)}; leaving existing ACLs" + ) + continue + + owner_group_gid = acl_group_gids[0] + if expansion.group_sams or expansion.cycle_paths: + log( + f"Expanded {sam}: {len(expansion.group_sams)} nested group(s), " + f"{expansion.terminal_user_count} terminal user(s)" + ) + + os.makedirs(path, exist_ok=True) + enforce_group_tree_permissions(path, owner_group_gid, acl_group_gids, admin_gid) os.makedirs(GROUP_ROOT, exist_ok=True) os.chown(GROUP_ROOT, 0, 0) @@ -850,7 +1266,7 @@ def with_lock() -> bool: groups = fetch_fileshare_groups() log(f"Discovered {len(groups)} data folder group(s) from AD") reconcile_db(conn, groups) - sync_dynamic_directory_permissions(conn) + sync_dynamic_directory_permissions(conn, groups) finally: conn.close() diff --git a/tests/test_reconcile_shares.py b/tests/test_reconcile_shares.py new file mode 100644 index 0000000..c839089 --- /dev/null +++ b/tests/test_reconcile_shares.py @@ -0,0 +1,178 @@ +import base64 +import unittest + +from app import reconcile_shares as rs + + +ROOT_DN = "CN=FS_Data,OU=Groups,DC=example,DC=com" +GROUP_A_DN = "CN=GroupA,OU=Groups,DC=example,DC=com" +GROUP_B_DN = "CN=GroupB,OU=Groups,DC=example,DC=com" +GROUP_C_DN = "CN=GroupC,OU=Groups,DC=example,DC=com" +USER_1_DN = "CN=Alice,OU=Users,DC=example,DC=com" +USER_2_DN = "CN=Bob,OU=Users,DC=example,DC=com" +USER_3_DN = "CN=Carol,OU=Users,DC=example,DC=com" +COMPUTER_DN = "CN=PC01,OU=Computers,DC=example,DC=com" + + +def principal(dn, sam, classes, members=()): + return { + "dn": dn, + "objectGUID": "", + "samAccountName": sam, + "objectClasses": set(classes), + "memberDns": list(members), + } + + +class LdapParsingTests(unittest.TestCase): + def test_parser_keeps_repeated_members_and_unfolds_lines(self): + display_name = base64.b64encode(b"Data Folder").decode("ascii") + output = f""" +dn: {ROOT_DN} +objectGUID: 550e8400-e29b-41d4-a716-446655440000 +objectClass: top +objectClass: group +sAMAccountName: FS_Data +displayName:: {display_name} +member: {USER_1_DN} +member: CN=Nested,OU=Groups,DC=example, + DC=com +member;range=2-2: CN=Ranged,OU=Groups,DC=example,DC=com + +""" + + entries = rs.parse_ldap_entries(output) + + self.assertEqual(len(entries), 1) + self.assertEqual(rs.ldap_first(entries[0], "displayName"), "Data Folder") + self.assertEqual( + rs.ldap_values(entries[0], "member"), + [ + USER_1_DN, + "CN=Nested,OU=Groups,DC=example,DC=com", + "CN=Ranged,OU=Groups,DC=example,DC=com", + ], + ) + self.assertEqual(rs.ldap_values(entries[0], "objectClass"), ["top", "group"]) + + def test_parse_groups_includes_dn_members_and_classes(self): + output = f""" +dn: {ROOT_DN} +objectGUID: 550e8400-e29b-41d4-a716-446655440000 +objectClass: group +sAMAccountName: FS_Data +displayName: Data Folder +distinguishedName: {ROOT_DN} +member: {GROUP_A_DN} + +""" + + groups = rs.parse_groups_from_ldap_output(output) + + self.assertEqual(len(groups), 1) + self.assertEqual(groups[0]["samAccountName"], "FS_Data") + self.assertEqual(groups[0]["shareName"], "Data Folder") + self.assertEqual(groups[0]["distinguishedName"], ROOT_DN) + self.assertEqual(groups[0]["memberDns"], [GROUP_A_DN]) + self.assertEqual(groups[0]["objectClasses"], {"group"}) + + def test_distinguished_name_filter_escapes_rfc4515_specials(self): + dn = r"CN=A*B(C)\\Name,DC=example,DC=com" + + self.assertEqual( + rs.build_distinguished_name_filter([dn]), + r"(distinguishedName=CN=A\2aB\28C\29\5c\5cName,DC=example,DC=com)", + ) + + +class MembershipExpansionTests(unittest.TestCase): + def lookup_from(self, principals): + indexed = {rs.normalize_dn(value["dn"]): value for value in principals} + + def lookup(dns): + return { + rs.normalize_dn(dn): indexed[rs.normalize_dn(dn)] + for dn in dns + if rs.normalize_dn(dn) in indexed + } + + return lookup + + def test_recursive_expansion_dedupes_groups_users_and_detects_cycle(self): + root_group = { + "objectGUID": "root-guid", + "samAccountName": "FS_Data", + "distinguishedName": ROOT_DN, + "memberDns": [USER_1_DN, GROUP_A_DN, GROUP_B_DN, COMPUTER_DN], + "objectClasses": {"group"}, + } + lookup = self.lookup_from( + [ + principal(USER_1_DN, "alice", {"user"}), + principal(USER_2_DN, "bob", {"user"}), + principal(USER_3_DN, "carol", {"user"}), + principal(COMPUTER_DN, "PC01$", {"user", "computer"}), + principal(GROUP_A_DN, "GroupA", {"group"}, [USER_2_DN, GROUP_B_DN]), + principal(GROUP_B_DN, "GroupB", {"group"}, [GROUP_A_DN, USER_3_DN]), + ] + ) + + expansion = rs.expand_group_membership(root_group, lookup_func=lookup) + + self.assertEqual(expansion.group_sams, ["GroupA", "GroupB"]) + self.assertEqual(expansion.terminal_user_count, 3) + self.assertEqual(expansion.unresolved_dns, []) + self.assertEqual(expansion.ignored_dns, [COMPUTER_DN]) + self.assertEqual(expansion.cycle_paths, [["FS_Data", "GroupA", "GroupB", "GroupA"]]) + + def test_self_cycle_is_reported_once(self): + root_group = { + "objectGUID": "root-guid", + "samAccountName": "FS_Data", + "distinguishedName": ROOT_DN, + "memberDns": [GROUP_C_DN], + "objectClasses": {"group"}, + } + lookup = self.lookup_from( + [principal(GROUP_C_DN, "GroupC", {"group"}, [GROUP_C_DN])] + ) + + expansion = rs.expand_group_membership(root_group, lookup_func=lookup) + + self.assertEqual(expansion.group_sams, ["GroupC"]) + self.assertEqual(expansion.cycle_paths, [["FS_Data", "GroupC", "GroupC"]]) + + def test_unresolved_member_is_reported(self): + missing_dn = "CN=Missing,OU=Groups,DC=example,DC=com" + root_group = { + "objectGUID": "root-guid", + "samAccountName": "FS_Data", + "distinguishedName": ROOT_DN, + "memberDns": [missing_dn], + "objectClasses": {"group"}, + } + + expansion = rs.expand_group_membership(root_group, lookup_func=lambda dns: {}) + + self.assertEqual(expansion.unresolved_dns, [missing_dn]) + self.assertEqual(expansion.group_sams, []) + + +class AclResolutionTests(unittest.TestCase): + def test_gid_resolution_preserves_order_dedupes_and_reports_missing(self): + mapping = {"FS_Data": 1001, "Nested": 1002} + + def resolver(workgroup, group_name): + self.assertEqual(workgroup, "EXAMPLE") + return mapping.get(group_name) + + gids, unresolved = rs.resolve_group_gids_for_acl( + "EXAMPLE", ["FS_Data", "Nested", "nested", "Missing"], resolver + ) + + self.assertEqual(gids, [1001, 1002]) + self.assertEqual(unresolved, ["Missing"]) + + +if __name__ == "__main__": + unittest.main()