fix(smbdb): properly apply filters for sqlalchemy queries

This commit is contained in:
Marshall Hallenbeck
2023-03-09 17:38:37 -05:00
parent f609c5ca92
commit 2b2124fb3e
+17 -17
View File
@@ -349,7 +349,7 @@ class database:
if not links:
asyncio.run(self.conn.execute(
insert(self.AdminRelationsTable).values(link)
insert(self.AdminRelationsTable).values(links)
))
def get_admin_relations(self, user_id=None, host_id=None):
@@ -371,12 +371,12 @@ class database:
q = delete(self.AdminRelationsTable)
if user_ids:
for user_id in user_ids:
q.filter(
q = q.filter(
self.AdminRelationsTable.c.userid == user_id
)
elif host_ids:
for host_id in host_ids:
q.filter(
q = q.filter(
self.AdminRelationsTable.c.hostid == host_id
)
asyncio.run(self.conn.execute(q))
@@ -435,7 +435,7 @@ class database:
# if we're returning a single host by ID
if self.is_computer_valid(filter_term):
q.filter(
q = q.filter(
self.ComputersTable.c.id == filter_term
)
results = asyncio.run(self.conn.execute(q)).first()
@@ -443,20 +443,20 @@ class database:
return [results]
# if we're filtering by domain controllers
elif filter_term == 'dc':
q.filter(
q = q.filter(
self.ComputersTable.c.dc == 1
)
if domain:
q.filter(
q = q.filter(
func.lower(self.ComputersTable.c.domain) == func.lower(domain)
)
# if we're filtering by ip/hostname
elif filter_term and filter_term != "":
q.filter(
logging.debug(f"In specific IP/hostname")
q = q.filter(
func.lower(self.ComputersTable.c.ip).like(func.lower(f"%{filter_term}%")) |
func.lower(self.ComputersTable.c.hostname).like(func.lower(f"%{filter_term}"))
)
results = asyncio.run(self.conn.execute(q)).all()
return results
@@ -590,11 +590,11 @@ class database:
def remove_group_relations(self, user_id=None, group_id=None):
q = delete(self.GroupRelationsTable)
if user_id:
q.filter(
q = q.filter(
self.GroupRelationsTable.c.userid == user_id
)
elif group_id:
q.filter(
q = q.filter(
self.GroupRelationsTable.c.groupid == group_id
)
asyncio.run(self.conn.execute(q))
@@ -627,12 +627,12 @@ class database:
q = select(self.UsersTable)
if self.is_user_valid(filter_term):
q.filter(
q = q.filter(
self.UsersTable.c.id == filter_term
)
# if we're filtering by username
elif filter_term and filter_term != '':
q.filter(
q = q.filter(
func.lower(self.UsersTable.c.username).like(func.lower(f"%{filter_term}%"))
)
results = asyncio.run(self.conn.execute(q)).all()
@@ -694,11 +694,11 @@ class database:
permissions = permissions.lower()
q = select(self.SharesTable)
if share_id:
q.filter(self.SharesTable.c.id == share_id)
q = q.filter(self.SharesTable.c.id == share_id)
if "r" in permissions:
q.filter(self.SharesTable.c.read == 1)
q = q.filter(self.SharesTable.c.read == 1)
if "w" in permissions:
q.filter(self.SharesTable.c.write == 1)
q = q.filter(self.SharesTable.c.write == 1)
results = asyncio.run(self.conn.execute(q)).all()
return results
@@ -709,9 +709,9 @@ class database:
self.SharesTable.c.computerid == computer_id
)
if "r" in permissions:
q.filter(self.SharesTable.c.read == 1)
q = q.filter(self.SharesTable.c.read == 1)
if "w" in permissions:
q.filter(self.SharesTable.c.write == 1)
q = q.filter(self.SharesTable.c.write == 1)
results = asyncio.run(self.conn.execute(q)).all()
return results