feat: support nested fs groups

This commit is contained in:
Ludwig Lehnert
2026-07-02 18:19:25 +00:00
parent f23aa64155
commit 618f9ebd06
3 changed files with 689 additions and 82 deletions

View File

@@ -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/<groupName>`).
- 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/<groupName>`
@@ -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/<groupName>
```
### `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.

View File

@@ -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 "<unknown>")
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()

View File

@@ -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()