#!/usr/bin/python3
# Copyright Canonical 2020
#
# Based on python implementation of samba-tool gpo by
# Andrew Tridgell 2010 and Amitay Isaacs 2011-2012.
# which is based on C implementation
# by Guenther Deschner and Wilco Baan Hofman
#
# This program is free software; you can redistribute it and/or modify
# it under the terms of the GNU General Public License as published by
# the Free Software Foundation; either version 3 of the License, or
# (at your option) any later version.
#
# This program is distributed in the hope that it will be useful,
# but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
# GNU General Public License for more details.
#
# You should have received a copy of the GNU General Public License
# along with this program.  If not, see <http://www.gnu.org/licenses/>.


import argparse
import sys

from samba import dsdb, param
from samba.auth import (system_session, user_session,
                        AUTH_SESSION_INFO_DEFAULT_GROUPS, AUTH_SESSION_INFO_AUTHENTICATED, AUTH_SESSION_INFO_SIMPLE_PRIVILEGES)
from samba.credentials import MUST_USE_KERBEROS, Credentials
from samba.dcerpc import security
from samba.ndr import ndr_unpack
import samba.security
from samba.samdb import SamDB
import ldb


class ObjectClass:
    user = 'user'
    computer = 'computer'


class ReturnCode:
    NOT_FOUND = 1
    CONNECTION_FAILED = 2
    GPO_FAILED = 3


TRANSPORT_ERRORS = (
    "NT_STATUS_HOST_UNREACHABLE",      # Host does not respond
    "NT_STATUS_NETWORK_UNREACHABLE",   # Local link is down
    "NT_STATUS_CONNECTION_REFUSED",    # Service does not respond on the other end
    "NT_STATUS_CONNECTION_DISCONNECTED", # Connection was dropped
    "NT_STATUS_CONNECTION_RESET",      # Connection was reset
    "NT_STATUS_IO_TIMEOUT",            # Connection timed out
    "NT_STATUS_OBJECT_NAME_NOT_FOUND", # Host does not exist
)


# Port of the Active Directory Global Catalog LDAP service. The Global Catalog
# holds a partial replica of every object in the forest, which lets it expand a
# user's complete group membership (universal and foreign-domain groups
# included) without emitting LDAP referrals that Samba is unable to follow.
GLOBAL_CATALOG_PORT = 3268

# Well-known SIDs that the LSA injects when it builds a token (the default and
# authenticated groups that user_session() used to add) but that are absent from
# the tokenGroups attribute. We add them explicitly to the token we assemble
# ourselves so the GPO access checks match what AD would compute: World covers
# GPOs whose ACL grants read to "Everyone", Network marks the session as remote,
# and Authenticated Users covers the most common GPO read ACE.
SID_WORLD = "S-1-1-0"
SID_NETWORK = "S-1-5-2"
SID_AUTHENTICATED_USERS = "S-1-5-11"


def parse_gplink(gplink):
    ''' Parse a gPLink into an array of dn and options '''
    ret = []

    if not gplink.strip() or gplink.strip() == "b''":
        return ret

    a = gplink.split(']')
    for g in a:
        if not g:
            continue
        d = g.split(';')
        if len(d) != 2 or not d[0].startswith("[LDAP://"):
            raise RuntimeError("Badly formed gPLink '%s'" % g)
        ret.append({'dn': d[0][8:], 'options': int(d[1])})
    return ret


def attr_default(msg, attrname, default):
    ''' Get an attribute from a ldap msg with a default '''
    if attrname in msg:
        return msg[attrname][0]
    return default


def connectLDAP(url):
    ''' Connect to the directory using Kerberos '''
    c = Credentials()
    c.set_kerberos_state(MUST_USE_KERBEROS)

    lp = param.LoadParm()
    c.guess(lp)

    return SamDB(url=url,
                 session_info=system_session(),
                 credentials=c, lp=lp)


def is_transport_error(exc):
    ''' Returns whether an exception reports a transport failure. '''
    return any(status in " ".join(str(arg) for arg in exc.args) for status in TRANSPORT_ERRORS)


def get_entity(samdb, accountname, objectClass):
    ''' Returns the entity for a given accountname and objectclass '''

    if objectClass == ObjectClass.user:
         samaccountname = accountname.split('@')[0]
         msg = samdb.search(expression='(&(|(samAccountName=%s)(samAccountName=%s$))(objectClass=%s))' %
                               (ldb.binary_encode(samaccountname), ldb.binary_encode(samaccountname), ldb.binary_encode(objectClass)),
                               attrs=['objectClass', 'objectSid'])
         if len(msg) == 0:
            msg = samdb.search(expression='(&(|(userPrincipalName=%s)(userPrincipalName=%s$))(objectClass=%s))' %
                               (ldb.binary_encode(accountname), ldb.binary_encode(accountname), ldb.binary_encode(objectClass)),
                               attrs=['objectClass', 'objectSid'])
    else:
        msg = samdb.search(expression='(&(|(samAccountName=%s)(samAccountName=%s$))(objectClass=%s))' %
                           (ldb.binary_encode(accountname), ldb.binary_encode(accountname), ldb.binary_encode(objectClass)),
                           attrs=['objectClass', 'objectSid'])

    if len(msg) == 0:
        raise Exception("Failed to find account %s" % accountname)
    current = msg[0]

    # Check that the object is really a computer or user if requested as such
    if objectClass == ObjectClass.computer and b'computer' not in current['objectClass']:
        raise Exception("Failed to find computer account %s" % accountname)
    elif objectClass == ObjectClass.user and b'computer' in current['objectClass']:
        raise Exception("Failed to find user account %s" % accountname)

    return current.dn, str(ndr_unpack(security.dom_sid, current["objectSid"][0]))


def get_group_sids(conn, dn):
    ''' Returns the object's transitive set of group SIDs as seen by conn.

    The tokenGroups constructed attribute is resolved against the given
    connection. Against the Global Catalog it expands universal and
    foreign-domain group membership across the whole forest but omits
    domain-local groups; against a domain controller it additionally returns the
    object's domain-local groups for its own domain. The caller unions both so
    the token matches what AD would compute. A plain tokenGroups search never
    emits the cross-domain LDAP referral that makes user_session() crash in a
    multi-domain forest. '''
    sids = []
    for msg in conn.search(base=str(dn), scope=ldb.SCOPE_BASE, attrs=['tokenGroups']):
        if 'tokenGroups' not in msg:
            continue
        for sid in msg['tokenGroups']:
            sids.append(ndr_unpack(security.dom_sid, sid))
    return sids


GPO_APPLY_GUID = "edacfd8f-ffb3-11d1-b41d-00a0c968f939"


def check_apply_gpo_right(secdesc, sids):
    ''' checks ntSecurityDescriptor if a GPO applies for a list of sIds '''
    # We need at least one allowed access to be applied
    applied = False
    for t in secdesc.as_sddl().split('(')[1:]:
        t = t.rstrip(')')
        access, _, _, access_right_guid, _, owner_sid = t.split(';')
        if access_right_guid != GPO_APPLY_GUID:
            continue
        for id in sids:
            if id != owner_sid:
                continue

            if access == "OA":
                applied = True

            # One denial is enough for denying the whole policy
            if access == "OD":
                return False

    return applied


def build_token(user_sid, group_sids):
    ''' Builds a security token from the user and group SIDs.

    This is the fallback for samba.auth.user_session(), whose internal group
    expansion fails in a multi-domain forest when the user belongs to a
    universal or foreign-domain group: the domain controller answers with an
    LDAP referral to the Global Catalog that Samba does not follow. Assembling
    the token directly from the SIDs already resolved through the Global Catalog
    avoids that code path entirely. It cannot reproduce everything user_session()
    injects (primary group, owner/READ_CONTROL semantics), so it is only used
    when user_session() actually raises. '''
    token = security.token()
    sids = ([security.dom_sid(user_sid)]
            + list(group_sids)
            + [security.dom_sid(s) for s in (SID_WORLD, SID_NETWORK, SID_AUTHENTICATED_USERS)])
    token.sids = sids
    # Real samba does NOT derive num_sids from the assigned list: the underlying
    # security_token keeps its own count and samba.security.access_check() only
    # iterates the first num_sids entries. Without setting it the token is seen
    # as empty and every GPO read is denied -- which silently skips every GPO
    # for any object that hits this fallback (e.g. multi-domain forest users
    # whose user_session() raised). It must be set explicitly.
    token.num_sids = len(sids)
    return token


def get_token(samdb, dn):
    ''' Returns the security token AD itself would build for the given DN.

    samba.auth.user_session() assembles the same token the LSA computes at
    logon: the primary group, the full transitive group membership, the default
    well-known SIDs with their correct attributes, and the owner/READ_CONTROL
    semantics a hand-built token cannot reproduce. It is the authoritative
    source and matches what the domain controller grants, so we prefer it and
    only fall back when it raises in a multi-domain forest. '''
    session_info_flags = (AUTH_SESSION_INFO_DEFAULT_GROUPS
                          | AUTH_SESSION_INFO_AUTHENTICATED
                          | AUTH_SESSION_INFO_SIMPLE_PRIVILEGES)
    session = user_session(samdb, lp_ctx=samdb.lp, dn=dn,
                           session_info_flags=session_info_flags)
    return session.security_token


def get_primary_group_sid(samdb, dn, object_sid):
    ''' Returns the object's primary group SID (Domain Computers / Domain Users).

    tokenGroups omits the primary group, yet GPO security descriptors very
    commonly grant read to it (a machine reads computer GPOs through Domain
    Computers), so the tokenGroups fallback must add it back. The SID is the
    object's domain SID -- its own SID with the last RID stripped -- combined
    with the primaryGroupID RID. '''
    msg = samdb.search(base=str(dn), scope=ldb.SCOPE_BASE, attrs=['primaryGroupID'])
    raw = msg[0]['primaryGroupID'][0]
    primary_group_id = int(raw.decode() if isinstance(raw, bytes) else raw)
    domain_sid = str(object_sid).rsplit('-', 1)[0]
    return security.dom_sid("%s-%d" % (domain_sid, primary_group_id))


def build_token_from_tokengroups(samdb, fqdn, dn, object_sid, accountname, debug=False):
    ''' Assemble the security token from tokenGroups when user_session() fails.

    Resolves group membership from two complementary sources and unions them:
    the domain controller (which includes the object's domain-local groups) and
    the Global Catalog (which expands universal and foreign-domain membership
    forest-wide but omits domain-local groups). Returns (token, sids), or
    (None, None) if neither source could be queried. '''
    group_sids = {}
    dc_failure = None
    gc_failure = None

    try:
        dc_group_sids = get_group_sids(samdb, dn)
        for sid in dc_group_sids:
            group_sids[str(sid)] = sid
        if debug:
            print("DEBUG: domain controller tokenGroups (%d): %s" % (len(dc_group_sids), ", ".join(str(s) for s in dc_group_sids)), file=sys.stderr)
    except Exception as exc:
        dc_failure = exc

    try:
        gc = connectLDAP("ldap://%s:%d" % (fqdn, GLOBAL_CATALOG_PORT))
        gc_group_sids = get_group_sids(gc, dn)
        for sid in gc_group_sids:
            group_sids[str(sid)] = sid
        if debug:
            print("DEBUG: Global Catalog tokenGroups (%d): %s" % (len(gc_group_sids), ", ".join(str(s) for s in gc_group_sids)), file=sys.stderr)
    except Exception as exc:
        # The Global Catalog is optional: the selected DC may not be a GC, or
        # port 3268 may be blocked. We can still proceed with the domain
        # controller's groups -- we just cannot expand cross-domain membership.
        gc_failure = exc
        print("WARNING: Global Catalog unreachable (%s); resolving group membership against the domain controller only. Cross-domain group membership is NOT expanded, so GPOs scoped to universal or foreign-domain groups may be skipped" % exc, file=sys.stderr)

    # Only a total failure -- neither source could be queried -- is fatal. An
    # object that legitimately belongs to no extra group yields an empty set,
    # which is not an error.
    if dc_failure is not None and gc_failure is not None:
        print("Couldn't resolve the group membership for %s: the domain controller lookup failed (%s) and the Global Catalog lookup failed (%s)" % (accountname, dc_failure, gc_failure), file=sys.stderr)
        return None, None, is_transport_error(dc_failure) or is_transport_error(gc_failure)

    # tokenGroups omits the primary group (Domain Computers for a machine,
    # Domain Users for a user), which GPO ACLs frequently grant read to. Add it
    # explicitly so the fallback token matches what user_session() would build.
    try:
        primary_sid = get_primary_group_sid(samdb, dn, object_sid)
        group_sids[str(primary_sid)] = primary_sid
        if debug:
            print("DEBUG: primary group SID: %s" % primary_sid, file=sys.stderr)
    except Exception as exc:
        if is_transport_error(exc):
            print("Couldn't resolve the primary group for %s: %s" % (accountname, exc), file=sys.stderr)
            return None, None, True
        if debug:
            print("DEBUG: could not resolve the primary group: %s" % exc, file=sys.stderr)

    group_sids = list(group_sids.values())
    try:
        token = build_token(object_sid, group_sids)
    except Exception as exc:
        print("Couldn't build the security token for %s: %s" % (accountname, exc), file=sys.stderr)
        return None, None, is_transport_error(exc)

    return token, [str(sid) for sid in group_sids], False


def get_gpos_for_dn(samdb, dn, token, sids, is_computer, debug=False):
    ''' List gpos for given dn, considering inheritance and enforced GPOs '''
    gpos = []
    inherit = True
    dn = ldb.Dn(samdb, str(dn)).parent()

    while True:
        msg = samdb.search(base=dn, scope=ldb.SCOPE_BASE, attrs=['gPLink', 'gPOptions'])[0]
        if 'gPLink' in msg:
            glist = parse_gplink(str(msg['gPLink'][0]))
            for g in glist:
                if not inherit and not (g['options'] & dsdb.GPLINK_OPT_ENFORCE):
                    continue
                if g['options'] & dsdb.GPLINK_OPT_DISABLE:
                    continue

                try:
                    sd_flags = (security.SECINFO_OWNER
                                | security.SECINFO_GROUP
                                | security.SECINFO_DACL)
                    gmsg = samdb.search(base=g['dn'], scope=ldb.SCOPE_BASE,
                                        attrs=['name', 'displayName', 'flags',
                                               'nTSecurityDescriptor', 'gPCFileSysPath'],
                                        controls=['sd_flags:1:%d' % sd_flags])
                    secdesc_ndr = gmsg[0]['nTSecurityDescriptor'][0]
                    secdesc = ndr_unpack(security.descriptor, secdesc_ndr)
                except Exception:
                    print("Failed to fetch gpo object with nTSecurityDescriptor %s" % g['dn'], file=sys.stderr)
                    print(file=sys.stderr) # Empty line (no escaped EOL as we need to echo -E the script when using integration tests coverage)
                    # GPOs that are unreadable are just skipped by AD
                    continue

                if debug:
                    try:
                        print("DEBUG: security descriptor for %s: %s" % (g['dn'], secdesc.as_sddl()), file=sys.stderr)
                    except Exception as sddl_exc:
                        print("DEBUG: security descriptor for %s unavailable: %s" % (g['dn'], sddl_exc), file=sys.stderr)

                try:
                    samba.security.access_check(secdesc, token,
                                                security.SEC_STD_READ_CONTROL
                                                | security.SEC_ADS_LIST
                                                | security.SEC_ADS_READ_PROP)
                except Exception:
                    # The object's token is not granted read on this GPO. AD
                    # silently skips GPOs it cannot read -- exactly like the
                    # unreadable-descriptor case above -- so this must never
                    # abort the whole refresh: otherwise a single hardened or
                    # security-filtered GPO denies the machine every policy.
                    print("Skipping GPO %s: the object has no read access to it" % g['dn'], file=sys.stderr)
                    print(file=sys.stderr) # Empty line (no escaped EOL as we need to echo -E the script when using integration tests coverage)
                    continue

                if not check_apply_gpo_right(secdesc, sids):
                    continue

                # check the flags on the GPO
                flags = int(attr_default(gmsg[0], 'flags', 0))
                if is_computer and (flags & dsdb.GPO_FLAG_MACHINE_DISABLE):
                    continue
                if not is_computer and (flags & dsdb.GPO_FLAG_USER_DISABLE):
                    continue

                # Enforced policy (higher wins)
                if g['options'] & dsdb.GPLINK_OPT_ENFORCE:
                    gpos.insert(0, (gmsg[0]['displayName'][0], gmsg[0]['gPCFileSysPath'][0]))
                # Others (higher have less weight)
                else:
                    gpos.append((gmsg[0]['displayName'][0], gmsg[0]['gPCFileSysPath'][0]))

        # check if this blocks inheritance
        gpoptions = int(attr_default(msg, 'gPOptions', 0))
        if gpoptions & dsdb.GPO_BLOCK_INHERITANCE:
            inherit = False

        if dn == samdb.get_default_basedn():
            break
        dn = dn.parent()
    return gpos


def main():
    parser = argparse.ArgumentParser(description='List GPOs for a user or computer.')
    parser.add_argument('fqdn', metavar='FQDN', type=str,
                        help='FQDN of the domain controller (without ldap:// prefix). \
                        e.g. dc.example.com')
    parser.add_argument('accountname', help='Name of the object to search for.')
    parser.add_argument('--objectclass', type=str,
                        choices=(ObjectClass.user, ObjectClass.computer), default=ObjectClass.user,
                        help='Class of the object to search for.')
    parser.add_argument('--debug', action='store_true',
                        help='Print the resolved security token and each GPO security descriptor to stderr to troubleshoot access checks.')

    args = parser.parse_args()

    accountname = args.accountname
    fqdn = args.fqdn

    try:
        samdb = connectLDAP("ldap://" + fqdn)
    except Exception as exc:
        if is_transport_error(exc):
            # samba/ldb prints the error message on stderr
            return ReturnCode.CONNECTION_FAILED
        print("Failed to open session: %s" % exc, file=sys.stderr)
        return ReturnCode.NOT_FOUND

    accountnames = [accountname]
    # Some AD limits computer names to 15 characters
    if args.objectclass == ObjectClass.computer and len(accountname) > 15:
        accountnames.append(accountname[:15])
    i = 0
    for accountname in accountnames:
        i += 1
        try:
            dn, object_sid = get_entity(samdb, accountname, args.objectclass)
            break
        except Exception as exc:
            print("Searching for account failed with: %s" % exc, file=sys.stderr)
            if is_transport_error(exc):
                return ReturnCode.CONNECTION_FAILED
            # We still have some candidates, don’t error out right away
            if i < len(accountnames):
                continue
            return ReturnCode.NOT_FOUND

    # Build the security token the same way AD itself does. user_session() is
    # authoritative: it reproduces the primary group, the default well-known
    # SIDs and the owner/READ_CONTROL semantics that a hand-assembled token
    # cannot, so the GPO access checks match exactly what the domain controller
    # grants. We use it whenever it succeeds.
    #
    # In a multi-domain forest it can fail: when the object belongs to a
    # universal or foreign-domain group defined in another domain, the DC
    # answers with an LDAP referral to the Global Catalog that Samba does not
    # follow and user_session() raises (#1358, #1256). Only then do we fall back
    # to assembling the token ourselves from the tokenGroups the domain
    # controller and the Global Catalog report -- which never chase a referral.
    try:
        token = get_token(samdb, dn)
        sids = [str(sid) for sid in token.sids]
    except Exception as session_exc:
        if is_transport_error(session_exc):
            print("Couldn't build the security token for %s: %s" % (accountname, session_exc), file=sys.stderr)
            return ReturnCode.CONNECTION_FAILED
        print("WARNING: user_session() could not build the token (%s); falling back to tokenGroups resolved from the domain controller and the Global Catalog" % session_exc, file=sys.stderr)
        token, sids, transport_error = build_token_from_tokengroups(samdb, fqdn, dn, object_sid, accountname, args.debug)
        if token is None:
            if transport_error:
                return ReturnCode.CONNECTION_FAILED
            return ReturnCode.GPO_FAILED

    sids.append(object_sid)
    sids.extend(('WD', 'NW', 'AU'))

    if args.debug:
        try:
            token_sids = ", ".join(str(s) for s in token.sids)
        except Exception as exc:
            token_sids = "<unavailable: %s>" % exc
        print("DEBUG: object SID: %s" % object_sid, file=sys.stderr)
        print("DEBUG: access-check token SIDs: %s" % token_sids, file=sys.stderr)
        print("DEBUG: apply-check SIDs: %s" % ", ".join(str(s) for s in sids), file=sys.stderr)

    try:
        gpos = get_gpos_for_dn(samdb, dn, token, sids, args.objectclass == ObjectClass.computer, debug=args.debug)
    except Exception as exc:
        print("Couldn't get GPOs: %s" % exc, file=sys.stderr)
        if is_transport_error(exc):
            return ReturnCode.CONNECTION_FAILED
        return ReturnCode.GPO_FAILED

    for g in gpos:
        gpo_name = g[0]
        gpo_path = parse_gpo_path(g[1], fqdn)
        print("%s\t%s" % (gpo_name, gpo_path))

def parse_gpo_path(gpo_path, dc_fqdn):
    ''' Parse a GPO path to a SMB path with the appropriate DC FQDN '''
    path = str(gpo_path).replace("\\", "/")
    parts = path[2:].split("/")
    parts[0] = dc_fqdn

    return "smb://" +"/".join(parts)

if __name__ == "__main__":
    exit(main())
