diff --git a/cme/protocols/smb.py b/cme/protocols/smb.py index 7e62a6b6..0b223b45 100755 --- a/cme/protocols/smb.py +++ b/cme/protocols/smb.py @@ -818,7 +818,7 @@ class smb(connection): domain, name = group.name.split('/') self.logger.highlight('{}\\{}'.format(domain.upper(), name)) try: - group_id = self.db.get_groups(groupName=self.args.local_groups, groupDomain=domain)[0][0] + group_id = self.db.get_groups(group_name=self.args.local_groups, group_domain=domain)[0][0] except IndexError: group_id = self.db.add_group(domain, self.args.local_groups) @@ -867,7 +867,7 @@ class smb(connection): self.logger.highlight('{}\\{}'.format(group.memberdomain, group.membername)) try: - group_id = self.db.get_groups(groupName=self.args.groups, groupDomain=group.groupdomain)[0][0] + group_id = self.db.get_groups(group_name=self.args.groups, group_domain=group.groupdomain)[0][0] except IndexError: group_id = self.db.add_group(group.groupdomain, self.args.groups) diff --git a/cme/protocols/smb/database.py b/cme/protocols/smb/database.py index bd44409f..484c48eb 100755 --- a/cme/protocols/smb/database.py +++ b/cme/protocols/smb/database.py @@ -167,7 +167,9 @@ class database: """ domain = domain.split('.')[0].upper() - results = self.conn.query(self.computers_table).filter(self.computers_table.c.ip == ip).all() + results = self.conn.query(self.computers_table).filter( + self.computers_table.c.ip == ip + ).all() host = { "ip": ip, "hostname": hostname, @@ -350,10 +352,13 @@ class database: def get_admin_relations(self, user_id=None, host_id=None): if user_id: - results = self.conn.query(self.admin_relations_table).filter(self.admin_relations_table.c.userid == user_id).all() + results = self.conn.query(self.admin_relations_table).filter( + self.admin_relations_table.c.userid == user_id + ).all() elif host_id: results = self.conn.query(self.admin_relations_table).filter( - self.admin_relations_table.c.computerid == host_id).all() + self.admin_relations_table.c.computerid == host_id + ).all() else: results = self.conn.query(self.admin_relations_table).all() @@ -406,7 +411,10 @@ class database: """ Check if this credential ID is valid. """ - results = self.conn.query(self.users_table).filter(self.users_table.c.id == credential_id, self.users_table.c.password is not None).all() + results = self.conn.query(self.users_table).filter( + self.users_table.c.id == credential_id, + self.users_table.c.password is not None + ).all() self.conn.commit() self.conn.close() return len(results) > 0 @@ -428,12 +436,18 @@ class database: """ # if we're returning a single credential by ID if self.is_credential_valid(filter_term): - results = self.conn.query(self.users_table).filter(self.users_table.c.id == filter_term).all() + results = self.conn.query(self.users_table).filter( + self.users_table.c.id == filter_term + ).all() elif cred_type: - results = self.conn.query(self.users_table).filter(self.users_table.c.credtype == cred_type).all() + results = self.conn.query(self.users_table).filter( + self.users_table.c.credtype == cred_type + ).all() # if we're filtering by username elif filter_term and filter_term != '': - results = self.conn.query(self.users_table).filter(func.lower(self.users_table.c.username).like(func.lower(f"%{filter_term}%"))).all() + results = self.conn.query(self.users_table).filter( + func.lower(self.users_table.c.username).like(func.lower(f"%{filter_term}%")) + ).all() # otherwise return all credentials else: results = self.conn.query(self.users_table).all() @@ -479,7 +493,9 @@ class database: """ Check if this host ID is valid. """ - results = self.conn.query(self.computers_table).filter(self.computers_table.c.id == hostID).all() + results = self.conn.query(self.computers_table).filter( + self.computers_table.c.id == hostID + ).all() self.conn.commit() self.conn.close() return len(results) > 0 @@ -511,39 +527,45 @@ class database: def get_domain_controllers(self, domain=None): return self.get_computers(filterTerm='dc', domain=domain) - def is_group_valid(self, groupID): + def is_group_valid(self, group_id): """ Check if this group ID is valid. """ - self.conn.execute('SELECT * FROM groups WHERE id=? LIMIT 1', [groupID]) - results = self.conn.fetchall() + results = self.conn.query(self.groups_table).filter( + self.groups_table.c.id == group_id + ).first() self.conn.commit() self.conn.close() - logging.debug('is_group_valid(groupID={}) => {}'.format(groupID, True if len(results) else False)) - return len(results) > 0 + valid = True if results else False + logging.debug(f"is_group_valid(groupID={group_id}) => {valid}") - def get_groups(self, filterTerm=None, groupName=None, groupDomain=None): + return valid + + def get_groups(self, filter_term=None, group_name=None, group_domain=None): """ Return groups from the database """ - if groupDomain: - groupDomain = groupDomain.split('.')[0].upper() - - if self.is_group_valid(filterTerm): - self.conn.execute("SELECT * FROM groups WHERE id=? LIMIT 1", [filterTerm]) - - elif groupName and groupDomain: - self.conn.execute("SELECT * FROM groups WHERE LOWER(name)=LOWER(?) AND LOWER(domain)=LOWER(?)", [groupName, groupDomain]) - - elif filterTerm and filterTerm !="": - self.conn.execute("SELECT * FROM groups WHERE LOWER(name) LIKE LOWER(?)", ['%{}%'.format(filterTerm)]) + if group_domain: + group_domain = group_domain.split('.')[0].upper() + if self.is_group_valid(filter_term): + results = self.conn.query(self.groups_table).filter( + self.groups_table.c.id == filter_term + ).first() + elif group_name and group_domain: + results = self.conn.query(self.groups_table).filter( + func.lower(self.groups_table.c.username) == func.lower(group_name), + func.lower(self.groups_table.c.domain) == func.lower(group_domain) + ).all() + elif filter_term and filter_term != "": + results = self.conn.query(self.groups_table).filter( + func.lower(self.groups_table.c.name).like(func.lower(f"%{filter_term}%")) + ).all() else: - self.conn.execute("SELECT * FROM groups") + results = self.conn.query(self.groups_table).all() - results = self.conn.fetchall() self.conn.commit() self.conn.close() - logging.debug('get_groups(filterTerm={}, groupName={}, groupDomain={}) => {}'.format(filterTerm, groupName, groupDomain, results)) + logging.debug(f"get_groups(filterTerm={filter_term}, groupName={group_name}, groupDomain={group_domain}) => {results}") return results diff --git a/cme/protocols/smb/db_navigator.py b/cme/protocols/smb/db_navigator.py index 8a93232c..b7954448 100644 --- a/cme/protocols/smb/db_navigator.py +++ b/cme/protocols/smb/db_navigator.py @@ -197,7 +197,7 @@ class navigator(DatabaseNavigator): self.display_groups(groups) else: - groups = self.db.get_groups(filterTerm=filterTerm) + groups = self.db.get_groups(filter_term=filterTerm) if len(groups) > 1: self.display_groups(groups)