diff --git a/.gitignore b/.gitignore index 9cd26e9..5ce0a43 100644 --- a/.gitignore +++ b/.gitignore @@ -3,12 +3,12 @@ *.csv *__pycache__* *.parquet -chunkinator.json +*chunkinator.json jobs.json *.xl* *.exe -securitytest.py *.toml system_config.json -Devel_unused/ -AirlockTools_client*/ \ No newline at end of file +Development_Stubs/ +AirlockTools_client*/ +.vscode/ \ No newline at end of file diff --git a/.vscode/settings.json b/.vscode/settings.json new file mode 100644 index 0000000..642ff51 --- /dev/null +++ b/.vscode/settings.json @@ -0,0 +1,3 @@ +{ + "python.REPL.enableREPLSmartSend": false +} \ No newline at end of file diff --git a/AirlockTools_Client.py b/AirlockTools_Client.py index 08a46d9..e881a2d 100644 --- a/AirlockTools_Client.py +++ b/AirlockTools_Client.py @@ -21,8 +21,13 @@ #TODO Fix Requirements.txt #TODO Create Generic system_config.json for gitea -import logging import os +import sys + +project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', "..")) +sys.path.append(project_root) +import asyncio +import logging import dotenv import urllib3 @@ -30,21 +35,26 @@ import urllib3 import utils.menus as menus from services.API import AirlockAPIWrapper from services.security import getAPI +from services.TaskQueue import AsyncTaskQueue from utils.setup import setup urllib3.disable_warnings( urllib3.exceptions.InsecureRequestWarning ) -def main(): + +async def main(): #Determine working directory, setup directory, configure logging, sent env, get API and URL if not already stored - - working_dir = setup() + working_dir = await setup() logger = logging.getLogger(__name__) logger.debug("šŸ” Logging test: this should appear in both console and file.") dotenv.load_dotenv(dotenv_path=working_dir / ".env") + queue = AsyncTaskQueue(worker_count = 3) + await queue.start_workers() + + try: url = os.getenv("URL") username = os.getenv("USERNAME") @@ -61,15 +71,13 @@ def main(): logger.error(f"Configuration error: {e}", exc_info=True) raise - - api = AirlockAPIWrapper( - base_url=str(os.getenv("URL")), - api_key = getAPI(username, "AirlockTools"), - ) - - - menus.menu_main(api) + if username: + api = AirlockAPIWrapper( + base_url=str(os.getenv("URL")), + api_key = await getAPI(username, "AirlockTools"), # pyright: ignore[reportArgumentType] + ) + await menus.menu_main(api, queue) if __name__ == "__main__": - main() + asyncio.run(main()) diff --git a/AirlockTools_Server.py b/AirlockTools_Server.py index 56a58a0..df82380 100644 --- a/AirlockTools_Server.py +++ b/AirlockTools_Server.py @@ -25,11 +25,16 @@ import logging import os +import Development_Stubs.WIP.localApproval as la import dotenv import urllib3 -import flows.localApproval as la -from Server.scheduler_async import recurring_job, register_function, reload_jobs, start_scheduler +from Development_Stubs.WIP.Server.scheduler_async import ( + recurring_job, + register_function, + reload_jobs, + start_scheduler, +) from services.API import AirlockAPIWrapper from services.policyhandler import updateAuditPoliciesFromEnforcementPolices from services.security import getAPI diff --git a/Server/scheduler_async.py b/Server/scheduler_async.py deleted file mode 100644 index e20c123..0000000 --- a/Server/scheduler_async.py +++ /dev/null @@ -1,177 +0,0 @@ -import asyncio -import json -import logging -import os -from typing import Any, Callable, Dict, List - -logger = logging.getLogger(__name__) - -# Registry of functions that can be scheduled -FUNCTION_MAP: Dict[str, Callable] = {} - -# Dictionary to manually track scheduled jobs by ID -scheduled_jobs: Dict[str, asyncio.TimerHandle] = {} - -# Path to the JSON file for job persistence TODO - pin this to the correct place -JOBS_FILE = os.path.join(os.getcwd(), "jobs.json") - -def register_function(name: str, func: Callable): - """ - Register a function so it can be called by name later. - Example: - register_function("say_hello", say_hello) - """ - FUNCTION_MAP[name] = func - -def load_jobs() -> List[Dict[str, Any]]: - """ - Load jobs from the JSON file, or return [] if none exist. - """ - if not os.path.exists(JOBS_FILE): - return [] - with open(JOBS_FILE, "r") as f: - return json.load(f) - -def save_jobs(jobs: List[Dict[str, Any]]): - """ - Save jobs to the JSON file (overwrite). - """ - with open(JOBS_FILE, "w") as f: - json.dump(jobs, f, indent=4) - -def cancel_job(job_id: str): - """ - Cancel a scheduled job by ID and remove it from the registry and persistence. - """ - handle = scheduled_jobs.pop(job_id, None) - if handle: - handle.cancel() - logger.info(f"Cancelled job '{job_id}'") - - jobs = [j for j in load_jobs() if j.get("id") != job_id] - save_jobs(jobs) - -def run_once_job(job_id: str, func_name: str, delay_seconds: float, args=None, kwargs=None, persist=True): - """ - Schedule a job to run once after a delay (in seconds). - """ - args = args or [] - kwargs = kwargs or {} - - def job_wrapper(): - func = FUNCTION_MAP.get(func_name) - if func is None: - logger.error(f"Function '{func_name}' is not registered.") - return - func(*args, **kwargs) - cancel_job(job_id) - - loop = asyncio.get_event_loop() - handle = loop.call_later(delay_seconds, job_wrapper) - scheduled_jobs[job_id] = handle - - if persist: - jobs = [j for j in load_jobs() if j.get("id") != job_id] - jobs.append({ - "id": job_id, - "type": "once", - "delay": delay_seconds, - "function": func_name, - "args": args, - "kwargs": kwargs - }) - save_jobs(jobs) - logger.info(f"Scheduled one-time job '{job_id}' to run in {delay_seconds} seconds.") - -def recurring_job(job_id: str, func_name: str, interval: float, args=None, kwargs=None, persist=True): - """ - Schedule a recurring job. - """ - args = args or [] - kwargs = kwargs or {} - - def job_wrapper(): - func = FUNCTION_MAP.get(func_name) - if func is None: - logger.error(f"Function '{func_name}' is not registered.") - return - func(*args, **kwargs) - # Reschedule the job - handle = asyncio.get_event_loop().call_later(interval, job_wrapper) - scheduled_jobs[job_id] = handle - - cancel_job(job_id) - handle = asyncio.get_event_loop().call_later(interval, job_wrapper) - scheduled_jobs[job_id] = handle - - if persist: - jobs = [j for j in load_jobs() if j.get("id") != job_id] - jobs.append({ - "id": job_id, - "type": "recurring", - "interval": interval, - "function": func_name, - "args": args, - "kwargs": kwargs - }) - save_jobs(jobs) - logger.info(f"Scheduled recurring job '{job_id}' every {interval} seconds.") - -def reload_jobs(): - """ - Reload jobs from JSON and reschedule them. - """ - jobs = load_jobs() - for job in jobs: - if job["type"] == "once": - run_once_job( - job["id"], - job["function"], - job["delay"], - job.get("args"), - job.get("kwargs"), - persist=False - ) - elif job["type"] == "recurring": - recurring_job( - job["id"], - job["function"], - job["interval"], - job.get("args"), - job.get("kwargs"), - persist=False - ) - -async def start_scheduler(): - """ - Start the asynchronous scheduler loop. - - This function is a placeholder to keep the event loop alive. - Jobs are scheduled using asyncio.call_later and do not require polling. - """ - try: - await asyncio.Event().wait() - except asyncio.CancelledError: - logger.critical("Scheduler stopped.") - - """ - Start the asynchronous scheduler loop. - - This function is a placeholder for compatibility. Since we use asyncio.call_later, - jobs are scheduled directly on the event loop and no polling is required. - - Usage: - # In an async app (e.g., Textual) - asyncio.create_task(start_scheduler()) - - # Or in a standalone script - async def main(): - await start_scheduler() - - asyncio.run(main()) - """ - try: - while True: - await asyncio.sleep(3600) # Sleep indefinitely; jobs run via call_later - except asyncio.CancelledError: - logger.critical("Scheduler stopped.") \ No newline at end of file diff --git a/flows/localApproval.py b/flows/localApproval.py deleted file mode 100644 index 789b917..0000000 --- a/flows/localApproval.py +++ /dev/null @@ -1,291 +0,0 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - - -import datetime -import logging -import os -import re -import time - -import dotenv -import numpy as np -import pandas as pd - -from models.agent import Agent -from services.agenthandler import findAllAgents, moveAgentToRelatedPolicy, selectAgents -from services.API import AirlockAPIWrapper -from utils.configmanager import get_protected_json, load_env, load_env_json -from utils.setup import get_base_directory -from utils.utils import colorText, get_sanitized_input - -logger = logging.getLogger(__name__) - -dotenv.load_dotenv() - - -def getLocalApprovals(api: AirlockAPIWrapper): - base_dir = get_base_directory - result = api.otp_find_awaiting() - local_approval = pd.DataFrame(result["response"]["otpusage"]) - if os.path.exists(f"{base_dir}\\cache\\newest_local_approval.parquet"): - previous_run = pd.read_parquet(f"{base_dir}\\cache\\newest_local_approval.parquet") - previous_run.to_parquet( - f"{base_dir}\\cache\\last_local_approval.parquet", index=False - ) - os.remove(f"{base_dir}\\cache\\newest_local_approval.parquet") - - # Only keep rows presumably created by the generate local approval function - local_approval = local_approval[ - local_approval["purpose"].str.startswith("šŸŽ« Local Approval šŸŽ«") - ] - - local_approval["batchid"] = local_approval["purpose"].apply( - lambda x: (match := re.search(r"batch:(\S+)", str(x))) and match.group(1) - ) - - if not local_approval.empty: - local_approval.to_parquet( - f"{base_dir}\\cache\\newest_local_approval.parquet", index=False - ) - - return local_approval - - -def scheduleAddingLAHashes(api: AirlockAPIWrapper): - - policy_relationship_map = get_protected_json("POLICY_MAP_ENF_AUD","{}") - bad_publisher_list = load_env_json("BAD_PUBLISHER", "[]") - pups = load_env_json("PUPS", "[]") - threat_tolerance_constant = load_env("VT_THREAT_TOLERANCE", cast_type = int) - - try: - register_function("add_hash", returnFromLocalApproval) - register_function("move_device", moveAgentToRelatedPolicy) - except Exception as e: - logger.warning(f"Failed to register functions: {e}") - return - - try: - approvals_df = getNewLocalApprovals(api) - if approvals_df.empty: - logger.debug("No new local approvals found. Nothing to schedule.") - return - batches = approvals_df.groupby("batchid") - except Exception as e: - logger.warning(f"Failed to retrieve or group local approvals: {e}") - return - - for batchid, batch_df in batches: - try: - duration_minutes = int(batch_df["duration"].iloc[0]) - start_time = datetime.datetime.now() - run_time = start_time + datetime.timedelta(minutes=duration_minutes) - early_time = start_time + datetime.timedelta(minutes=np.floor(duration_minutes * 0.95)) - - early_timestamp = early_time.timestamp() - run_timestamp = run_time.timestamp() - - # Schedule add_hash job - try: - run_once_job( - f"add_hash_{batchid}", - "add_hash", - early_timestamp, - [ - api, - batch_df, - policy_relationship_map, - bad_publisher_list, - pups, - threat_tolerance_constant, - ], - None, - ) - logger.debug(f"Scheduled add_hash for batch {batchid} at {early_time}") - except Exception: - logger.debug("Failed to schedule add_hash for batch {batchid}: {e}") - - # Schedule move_device jobs - devices = batch_df["agentid"].drop_duplicates().tolist() - agents = [] - - for device in devices: - rows = api.agent_find_by_hostname(device).iterrows() - agents += [Agent(**row["data"]) for _, row in rows] - - for agent in agents: - try: - run_once_job( - f"move_device_{agent.hostame}_{batchid}", - "move_device", - run_timestamp, - [api, agent, policy_relationship_map], - "enforcement", - ) - - print( - f"Scheduled move_device for device {agent.hostname} in batch {batchid} at {run_time}" - ) - except Exception as e: - print( - f"Failed to schedule move_device for device {agent.hostname} in batch {batchid}: {e}" - ) - - except Exception as e: - logger.warning(f"Failed to process batch {batchid}: {e}") - - -def returnFromLocalApproval(api, device_df, policy_relationship_map, bad_publisher_list, pups, threat_tolerance_constant -): - """ - # Get unique policy names from device list - policies_in_devicelist = sorted(device_df['policy_name'].unique().tolist()) - - # Create inverse map to go from Audit to Enforcement - inverse_map = {v: k for k, v in policy_relationship_map.items()} - - # Fetch all policies - all_policies = [Policy(row['groupid'], row['hidden'], row['name'], row['parent']) for _, row in api.policy_find_all().iterrows()] - - # Define policy types - policy_types = [1, 2, 6, 7] - - #TODO finish logic for adding hashes - """ - working_dir = load_env("WORKING_DIR") - policy_relationship_map = get_protected_json("POLICY_MAP_ENF_AUD","{}") - bad_publisher_list = load_env_json("BAD_PUBLISHER", "[]") - pups = load_env_json("PUPS", "[]") - threat_tolerance_constant = load_env("VT_THREAT_TOLERANCE") - print(f"{working_dir}, {policy_relationship_map}, {bad_publisher_list}, {pups}, {threat_tolerance_constant}") - -def moveToLocalApproval(api: AirlockAPIWrapper): - possible_durations = [15, 60, 360, 1440, 10080] - duration_selected = None - - print(colorText("Please select a duration:", "white")) - for i, option in enumerate(possible_durations, start=1): - print(f"{i}. {option}") - - try: - - choice = int(get_sanitized_input("Enter the number of your choice:")) - if 1 <= choice <= len(possible_durations): - duration_selected = possible_durations[choice - 1] - print(colorText(f"You selected: {duration_selected}", "yellow")) - logger.debug(f"You selected: {duration_selected}") - else: - print(colorText("āŒ Invalid choice.", "red")) - logger.debug("Invalid Input") - return - except ValueError: - print(colorText("āŒ Invalid input. Please enter a number.", "red")) - logger.debug("Invalid Input") - return - - agents = selectAgents(api) - batch = int(time.time()) - - if not agents: - print(colorText("āŒ No agents found or error retrieving agents.", "red")) - logger.debug("No agents found or error retrieving agents") - return - - for agent in agents: - try: - addLocalApproval(api, batch, duration_selected, agent.agentid) - moveAgentToRelatedPolicy(api, agent, "audit") - except Exception as e: - print(colorText(f"āŒ Error processing agent {agent.hostname}: {e}", "red")) - - -def addLocalApproval(api: AirlockAPIWrapper, batchid, duration_selected, agentid): - - purpose = f"šŸŽ« Local Approval šŸŽ« - {duration_selected} mins - batch:{batchid} Client:{agentid}" - api.otp_generate(agentid, duration_selected, purpose) - - - -def monitorAuditStatus(api: AirlockAPIWrapper): - current_agents = findAllAgents(api) - last_agents = [] - if not last_agents: - last_agents = current_agents - policy_relationship_map = get_protected_json("POLICY_MAP_ENF_AUD","{}") - - # Reverse map for audit → enforcement - reverse_policy_map = {v: k for k, v in policy_relationship_map.items()} - known_transitions = set(policy_relationship_map.items()) | set(reverse_policy_map.items()) - - # Index last_agents by hostname for quick lookup - last_agent_map = {agent.hostname: agent for agent in last_agents} - - # Result buckets - newly_added = [] - same_policy = [] - moved_to_audit = [] - moved_to_enforcement = [] - unusual_move = [] - - for current in current_agents: - previous = last_agent_map.get(current.hostname) - - if not previous: - newly_added.append(current) - continue - - if current.groupid == previous.groupid: - same_policy.append(current) - elif (previous.groupid, current.groupid) in known_transitions: - moved_to_audit.append(current) - elif (current.groupid, previous.groupid) in known_transitions: - moved_to_enforcement.append(current) - else: - unusual_move.append(current) - - # Return all five DataFrames - return newly_added, same_policy, moved_to_audit, moved_to_enforcement, unusual_move - - -def getNewLocalApprovals(api: AirlockAPIWrapper): - - working_dir = load_env("WORKING_DIR") - current_la = getLocalApprovals(api) - - # Load old approval list - old_la_path = f"{working_dir}\\Scheduling\\last_local_approval.parquet" - if os.path.exists(old_la_path): - old_la = pd.read_parquet(old_la_path) - else: - old_la = pd.DataFrame(columns=current_la.columns) - - # Create composite keys - current_la["key"] = current_la["clientid"].astype(str) + "_" + current_la["granted"].astype(str) - old_la["key"] = old_la["clientid"].astype(str) + "_" + old_la["granted"].astype(str) - - # Find new entries - new_entries = current_la[~current_la["key"].isin(old_la["key"])] - - # Convert 'granted' to datetime and filter by last 10 minutes - new_entries["granted"] = pd.to_datetime(new_entries["granted"], utc=True, errors="coerce") - ten_minutes_ago = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta(minutes=10) - recent_entries = new_entries[new_entries["granted"] > ten_minutes_ago] - - # Save current approvals for next run - current_la.drop(columns=["key"], inplace=True) - current_la.to_parquet(old_la_path, index=False) - - return recent_entries diff --git a/flows/otp.py b/flows/otp.py index 4f5d52b..1fe48b1 100644 --- a/flows/otp.py +++ b/flows/otp.py @@ -1,69 +1,61 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - - - +import asyncio import logging from services.agenthandler import selectAgents from services.API import AirlockAPIWrapper -from utils.selector import Selector +from utils.Selector import Selector from utils.utils import colorText, get_sanitized_input logger = logging.getLogger(__name__) - - - -def generate(api: AirlockAPIWrapper): +async def generate(api: AirlockAPIWrapper): otp_dict = {} - agents = selectAgents(api) - print(colorText("Would you like to continue with these devices?","white")) + agents = await selectAgents(api) + + print(colorText("Would you like to continue with these devices?", "white")) for agent in agents: print(agent.hostname) - confirm = Selector.confirm() + + confirm = await Selector.confirm() if agents and confirm: - requester = get_sanitized_input("Who is requesting the OTP: ") - because = get_sanitized_input("Why/What work are they doing?: ") - + requester = await get_sanitized_input("Who is requesting the OTP: ") + because = await get_sanitized_input("Why/What work are they doing?: ") + purpose = f"Requester: {requester} - for : {because}" possible_durations = [15, 60, 360, 1440, 10080] print(colorText("Please select a duration in minutes: ", "white")) print(colorText("15 mins, 60 mins, 360 mins(6 Hours), 1440 mins (24 Hours), 10080 mins (7 Days):", "white")) - duration_selected = Selector.select_int(possible_durations) + duration_selected = await Selector.select_int(possible_durations) + if duration_selected: - for agent in agents: - otp_code = api.otp_generate(agent.agentid, duration_selected, purpose) + async def generate_otp(agent): + otp_code = await api.otp_generate(agent.agentid, duration_selected, purpose) # pyright: ignore[reportArgumentType] logger.info(f"Generated OTP for {agent.hostname}: {otp_code}") - otp_dict[agent.hostname] = otp_code + return agent.hostname, otp_code + + results = await asyncio.gather(*(generate_otp(agent) for agent in agents)) + otp_dict = dict(results) return otp_dict -def otp_activities_by_agent(api: AirlockAPIWrapper): - agents = selectAgents(api) + +async def otp_activities_by_agent(api: AirlockAPIWrapper): + agents = await selectAgents(api) otp_dict = {} + for agent in agents: - otp_info = api.otp_find_by_agent(agent.agentid) + otp_info = await api.otp_find_by_agent(agent.agentid) otp_dict[agent.hostname] = otp_info return otp_dict -def revoke(api: AirlockAPIWrapper): - otp_dict = otp_activities_by_agent(api) - list_to_revoke = [entry["otpid"] for entry in otp_dict] +async def revoke(api: AirlockAPIWrapper): + otp_dict = await otp_activities_by_agent(api) + list_to_revoke = [entry["otpid"] for entry in otp_dict.values() if entry] + if otp_dict and list_to_revoke: - for revokee in list_to_revoke: - api.otp_revoke(revokee) \ No newline at end of file + async def revoke_otp(otpid): + await api.otp_revoke(otpid) + + await asyncio.gather(*(revoke_otp(otpid) for otpid in list_to_revoke)) diff --git a/flows/prepPolicy.py b/flows/prepPolicy.py index d79d242..2d4bab8 100644 --- a/flows/prepPolicy.py +++ b/flows/prepPolicy.py @@ -13,10 +13,11 @@ # You should have received a copy of the GNU Affero General Public License # along with this program. If not, see . +import asyncio import logging import os import os.path -from typing import List +from typing import List, Optional import dotenv import pandas as pd @@ -24,12 +25,12 @@ import pandas as pd from models.execution import ExecutionHistoryRecord, Hash from models.policy import Allowlist, Policy from services.API import AirlockAPIWrapper +from services.TaskQueue import AsyncTaskQueue, run_sync_task_in_thread from utils.configmanager import get_protected_value, load_env, load_env_json -from utils.selector import Selector +from utils.Selector import Selector from utils.utils import ( colorText, formatHTML, - import_to_dataframe, regulator, ) @@ -37,79 +38,86 @@ logger = logging.getLogger(__name__) dotenv.load_dotenv() +async def selectPolicies(api: AirlockAPIWrapper, allow_multiple=True) -> List[Policy]: + df = await api.policy_find_all() + policies = [Policy(**row.to_dict()) for _, row in df.iterrows()] - -def selectPolicies(api: AirlockAPIWrapper, allow_multiple=True) -> List[Policy]: - - policies = [Policy(**row.to_dict()) for _, row in api.policy_find_all().iterrows()] logger.debug("Prompting for Policies") print(colorText("Please select policy/policies", "white")) - selected = Selector.select_objects(policies, allow_multiple, prompt_each=True) + + selected = await Selector.select_objects(policies, allow_multiple, prompt_each=True) if selected is None: return [] - # Normalize to always return a list - logger.debug("Returning {selected.dict}") + logger.debug("Returning selected policies") return selected if isinstance(selected, list) else [selected] -def selectAllowlists(api: AirlockAPIWrapper, policy = all, allow_multiple=True) -> List[Allowlist]: - if policy == "all": allowlists = [Allowlist(**row.to_dict()) for _, row in api.allowlist_find_all().iterrows()] - else: allowlists = [Allowlist(**row.to_dict()) for _, row in api.policy_list_allowlists(policy[0].groupid).iterrows()] +async def selectAllowlists(api: AirlockAPIWrapper, policy="all", allow_multiple=True) -> List[Allowlist]: + if policy == "all": + df = await api.allowlist_find_all() + else: + df = await api.policy_list_allowlists(policy[0].groupid) # pyright: ignore[reportAttributeAccessIssue] + + allowlists = [Allowlist(**row.to_dict()) for _, row in df.iterrows()] + logger.debug("Prompting for Allowlist(s)") print(colorText("Please select allowlist(s)", "white")) - selected = Selector.select_objects(allowlists, allow_multiple, prompt_each=True) + + selected = await Selector.select_objects(allowlists, allow_multiple, prompt_each=True) if selected is None: return [] - # Normalize to always return a list logger.debug(f"Returning {selected}") return selected if isinstance(selected, list) else [selected] -def sortHashes( +async def sortHashes( api: AirlockAPIWrapper, + queue: AsyncTaskQueue, selected_policies: List[Policy], - type=[1, 2, 6, 7] + type=[1, 2, 6, 7], + history_days: Optional[int] = None ): - working_dir = load_env("WORKING_DIR") - history_days = Selector.select_value( - prompt="Enter how many days of history to pull (1–150): ", - value_type=int, - valid_range=(1, 150), - ) - + if history_days is None: + history_days = await Selector.select_value( + prompt="Enter how many days of history to pull (1–150): ", + value_type=int, + valid_range=(1, 150), + ) + logger.debug(f"{history_days} day selected for history") - + if history_days is None: logging.warning("No history range selected. Aborting.") return executions = [] hashes = [] + working_dir = await load_env("WORKING_DIR") # Pull execution histories for each policy - - policy_executions = ExecutionHistoryRecord.from_policies( + policy_executions = await ExecutionHistoryRecord.from_policies( api, selected_policies, type_=type, history_days=history_days ) - - logger.debug(f"Policy_executions is {policy_executions}") + logger.debug(f"Policy_executions is {policy_executions}") executions.extend(policy_executions) logger.debug(f"Executions contains {executions}") + if executions: - hashes = [Hash(sha256=row["sha256"], **row["data"]) for _, row in api.hash_query([record.sha256 for record in executions]).iterrows() + sha_list = [record.sha256 for record in executions] + hash_df = await api.hash_query(sha_list) + hashes = [ + Hash(sha256=row["sha256"], **row["data"]) + for _, row in hash_df.iterrows() ] if hashes: unique_hashes = Hash.deduplicate(hashes) - - needs_review, approved, unapproved = Hash.categorize_hashes( - hashes=unique_hashes - ) + needs_review, approved, unapproved = await Hash.categorize_hashes(hashes=unique_hashes) categories = { "needs_review": needs_review, @@ -118,157 +126,158 @@ def sortHashes( } for label, category in categories.items(): - csv_path = f"{working_dir}\\Needs_Review\\Review_First\\{label}_executions.csv" - html_path = f"{working_dir}\\Needs_Review\\HTML\\{label}.html" + csv_path = f"{working_dir}/Needs_Review/Review_First/{selected_policies[0].name}_{label}_executions.csv" + html_path = f"{working_dir}/Needs_Review/HTML/{selected_policies[0].name}_{label}.html" - ExecutionHistoryRecord.enrich_with_hashes_and_export( - executions, category, f"{working_dir}\\Needs_Review\\Review_First", label=label - ) - df = import_to_dataframe(csv_path) - formatHTML(df, html_path) + df = await ExecutionHistoryRecord.enrich_with_hashes(executions, category) + + asyncio.create_task(queue.enqueue( + f"DF TO CSV {selected_policies[0].name}_{label}", + run_sync_task_in_thread, + df.to_csv, + csv_path, + index=False, + encoding='utf-8' + )) -def buildPathsandPublishers(split): - working_dir = load_env("WORKING_DIR") - df1 = pd.DataFrame() - df2 = pd.DataFrame() - all_approved_hashes = pd.DataFrame() - path1 = f"{working_dir}\\Approved\\approved_executions.csv" - path2 = f"{working_dir}\\Approved\\needs_review_executions.csv" - if os.path.exists(path1): - df1 = pd.read_csv(path1) - else: - logger.warning(f"File not found: {path1}") + asyncio.create_task(queue.enqueue( + f"DF TO HTML {selected_policies[0].name}_{label}", + run_sync_task_in_thread, + formatHTML, + df, + html_path + )) - if os.path.exists(path2): - df2 = pd.read_csv(path2) - else: - logger.warning(f"File not found: {path2}") + print("sortHashes completed successfully.") + + +async def buildPathsandPublishers(queue: AsyncTaskQueue, split): + working_dir = await load_env("WORKING_DIR") + path1 = f"{working_dir}/Approved/approved_executions.csv" + path2 = f"{working_dir}/Approved/needs_review_executions.csv" + + df1 = await asyncio.to_thread(pd.read_csv, path1) if os.path.exists(path1) else pd.DataFrame() + if df1.empty: + logger.warning(f"File not found or empty: {path1}") + + df2 = await asyncio.to_thread(pd.read_csv, path2) if os.path.exists(path2) else pd.DataFrame() + if df2.empty: + logger.warning(f"File not found or empty: {path2}") if df1.empty and df2.empty: logger.warning("Both DataFrames are empty. Skipping sort.") all_approved_hashes = pd.DataFrame() - logger.debug(all_approved_hashes.head) + logger.debug(all_approved_hashes.head()) + return + + all_approved_hashes = pd.concat([df1, df2], ignore_index=True) + if "filename_exec" in all_approved_hashes.columns: + all_approved_hashes = all_approved_hashes.sort_values(by="filename_exec") else: - all_approved_hashes = pd.concat([df1, df2], ignore_index=True) - if "filename_exec" in all_approved_hashes.columns: - all_approved_hashes = all_approved_hashes.sort_values(by="filename_exec") - else: - logger.warning("Warning: 'filename_exec' column not found in concatenated DataFrame.") + logger.warning("'filename_exec' column not found in concatenated DataFrame.") - if not all_approved_hashes.empty: - primary_path_exclusions = calculatePath( - all_approved_hashes, - split, - ) - remaining_hashes = all_approved_hashes[ - ~all_approved_hashes["sha256"].isin(primary_path_exclusions["sha256"]) - ] - secondary_path_exclusions = calculatePath( - remaining_hashes, split - ) - remaining_hashes = remaining_hashes[ - ~remaining_hashes["sha256"].isin(secondary_path_exclusions["sha256"]) - ] - dataframes = { - "primary_Paths": primary_path_exclusions, - "secondary_Paths": secondary_path_exclusions, - "hashes_to_add": remaining_hashes, - } - logger.debug("Preparing to sort dataframes") - for name, df in dataframes.items(): - logger.debug(f" DataFrame headers: {list(df.columns)}") - if name == "hashes_to_add": df.sort_values(by="filename_exec", inplace=True) - else: df.sort_values(by="longestcfp", inplace=True) - - df.to_csv(f"{working_dir}\\Needs_Review\\Review_Second\\{name}.csv", index=False) - formatHTML(df, f"{working_dir}\\Needs_Review\\HTML\\{name}.html") + primary_path_exclusions = await calculatePath(all_approved_hashes, split) + remaining_hashes = all_approved_hashes[~all_approved_hashes["sha256"].isin(primary_path_exclusions["sha256"])] + secondary_path_exclusions = await calculatePath(remaining_hashes, split) + remaining_hashes = remaining_hashes[~remaining_hashes["sha256"].isin(secondary_path_exclusions["sha256"])] - if not all_approved_hashes.empty: - # Drop all not signed, only keep unique values - publist = all_approved_hashes[ - all_approved_hashes["publisher_hash"] != "Not Signed" - ].drop_duplicates(subset=["publisher_hash"]) - # Remove Bad publisher if somehow they made it this far - pattern = regulator(load_env_json("BAD_PUBLISHERS","[]")) - publist = publist[~publist["publisher_hash"].str.contains(pattern, na=False)] - publist = publist[["publisher_hash"]] - publist.sort_values(by="publisher_hash", inplace=True) - publist.to_csv(f"{working_dir}\\Needs_Review\\Review_Second\\publishers.csv", index=False) - -def buildPreflights(): - working_dir = load_env("WORKING_DIR") - - df1 = pd.DataFrame() - df2 = pd.DataFrame() - approved_hashes = pd.DataFrame() - approved_publishers = pd.DataFrame() - - hash = f"{working_dir}\\Approved\\hashes_to_add.csv" - path1 = f"{working_dir}\\Approved\\primary_Paths.csv" - path2 = f"{working_dir}\\Approved\\secondary_Paths.csv" - publishers = f"{working_dir}\\Approved\\publishers.csv" - - if os.path.exists(hash): - approved_hashes = pd.read_csv(hash) - - else: - logger.warning(f"File not found: {hash}") - - if os.path.exists(path1): - df1 = pd.read_csv(path1) - else: - logger.warning(f"File not found: {path1}") - - if os.path.exists(path2): - df2 = pd.read_csv(path2) - else: - logger.warning(f"File not found: {path2}") - - if df1.empty and df2.empty: - logger.warning("Both DataFrames are empty. Skipping sort.") - approved_paths = pd.DataFrame() - else: - approved_paths = pd.concat([df1, df2], ignore_index=True) - - if os.path.exists(publishers): - approved_publishers = pd.read_csv(publishers) - - else: - logger.warning(f"File not found: {publishers}") - - dataframes = {"approved_paths": approved_paths, "approved_hashes": approved_hashes, "approved_publishers": approved_publishers} + dataframes = { + "primary_Paths": primary_path_exclusions, + "secondary_Paths": secondary_path_exclusions, + "hashes_to_add": remaining_hashes, + } for name, df in dataframes.items(): - logger.debug(f" DataFrame headers: {list(df.columns)}") - if name == "approved_paths":df.sort_values(by="longestcfp", inplace=True) - elif name == "approved_hashes":df.sort_values(by="filename_exec", inplace=True) - elif name == "approved_publishers" : df.sort_values(by="publisher_hash", inplace=True) - - df.to_csv(f"{working_dir}\\Preflight\\{name}.csv", index=False) - formatHTML(df, f"{working_dir}\\Preflight\\HTML\\{name}.html") + logger.debug(f"DataFrame headers for {name}: {list(df.columns)}") + sort_column = "filename_exec" if name == "hashes_to_add" else "longestcfp" + df.sort_values(by=sort_column, inplace=True) + csv_path = f"{working_dir}/Needs_Review/Review_Second/{name}.csv" + html_path = f"{working_dir}/Needs_Review/HTML/{name}.html" + await asyncio.to_thread(df.to_csv, csv_path, index=False) + await asyncio.to_thread(formatHTML, df, html_path) -def splitFilepathsGrouped(df, col="filename"): - path_exclusion_constant = get_protected_value("PATH_EXCLUSION_CONST", cast_type= int) - min_files_for_path = get_protected_value("MIN_FILES_FOR_PATH", cast_type= int) + publist = all_approved_hashes[all_approved_hashes["publisher_hash"] != "Not Signed"].drop_duplicates(subset=["publisher_hash"]) + pattern = regulator(await load_env_json("BAD_PUBLISHERS", "[]")) + publist = publist[~publist["publisher_hash"].str.contains(pattern, na=False)] + publist = publist[["publisher_hash"]] + publist.sort_values(by="publisher_hash", inplace=True) + pub_csv_path = f"{working_dir}/Needs_Review/Review_Second/publishers.csv" + await asyncio.to_thread(publist.to_csv, pub_csv_path, index=False) + print("buildPathsandPublishers completed asynchronously.") + + +async def buildPreflights(): + working_dir = await load_env("WORKING_DIR") + + hash_path = f"{working_dir}/Approved/hashes_to_add.csv" + path1 = f"{working_dir}/Approved/primary_Paths.csv" + path2 = f"{working_dir}/Approved/secondary_Paths.csv" + publishers_path = f"{working_dir}/Approved/publishers.csv" + + df1 = await asyncio.to_thread(pd.read_csv, path1) if os.path.exists(path1) else pd.DataFrame() + if df1.empty: + logger.warning(f"File not found or empty: {path1}") + + df2 = await asyncio.to_thread(pd.read_csv, path2) if os.path.exists(path2) else pd.DataFrame() + if df2.empty: + logger.warning(f"File not found or empty: {path2}") + + approved_hashes = await asyncio.to_thread(pd.read_csv, hash_path) if os.path.exists(hash_path) else pd.DataFrame() + if approved_hashes.empty: + logger.warning(f"File not found or empty: {hash_path}") + + approved_publishers = await asyncio.to_thread(pd.read_csv, publishers_path) if os.path.exists(publishers_path) else pd.DataFrame() + if approved_publishers.empty: + logger.warning(f"File not found or empty: {publishers_path}") + + approved_paths = pd.concat([df1, df2], ignore_index=True) if not (df1.empty and df2.empty) else pd.DataFrame() + + dataframes = { + "approved_paths": approved_paths, + "approved_hashes": approved_hashes, + "approved_publishers": approved_publishers + } + + for name, df in dataframes.items(): + logger.debug(f"DataFrame headers for {name}: {list(df.columns)}") + if name == "approved_paths": + df.sort_values(by="longestcfp", inplace=True) + elif name == "approved_hashes": + df.sort_values(by="filename_exec", inplace=True) + elif name == "approved_publishers": + df.sort_values(by="publisher_hash", inplace=True) + + csv_path = f"{working_dir}/Preflight/{name}.csv" + html_path = f"{working_dir}/Preflight/HTML/{name}.html" + + await asyncio.to_thread(df.to_csv, csv_path, index=False) + await asyncio.to_thread(formatHTML, df, html_path) + + print("buildPreflights completed asynchronously.") + + +async def splitFilepathsGrouped(df, col="filename"): + path_task = asyncio.create_task(get_protected_value("PATH_EXCLUSION_CONST", int)) + min_files_task = asyncio.create_task(get_protected_value("MIN_FILES_FOR_PATH", int)) + + path_exclusion_constant = await path_task + min_files_for_path = await min_files_task + def clean_split(path): if not isinstance(path, (str, bytes, os.PathLike)): return [] parts = os.path.normpath(path).split(os.sep) - parts = [p for p in parts if p] # Remove empty strings - return parts + return [p for p in parts if p] - # Diagnostic: log any non-string entries non_string_entries = df[~df[col].apply(lambda x: isinstance(x, (str, bytes, os.PathLike)))] if not non_string_entries.empty: - print(f"[WARNING] Non-string entries found in column '{col}':") - print(non_string_entries) + logger.warning(f"Non-string entries found in column '{col}':") + logger.debug(non_string_entries) df = df.copy() split_paths = df[col].apply(clean_split) - - # Filter out paths with fewer than `min_files_for_path` components df = df[split_paths.apply(lambda parts: len(parts) >= min_files_for_path)].copy() split_paths = split_paths[df.index] @@ -295,11 +304,7 @@ def splitFilepathsGrouped(df, col="filename"): for i, parts in enumerate(split_parts): filename = parts[-1] - middle = ( - os.sep.join(parts[len(common_prefix):-1]) - if len(parts) > len(common_prefix) + 1 - else "" - ) + middle = os.sep.join(parts[len(common_prefix):-1]) if len(parts) > len(common_prefix) + 1 else "" row = group_df.iloc[i].copy() row["longestcfp"] = prefix_str row["middle"] = middle @@ -309,19 +314,16 @@ def splitFilepathsGrouped(df, col="filename"): return pd.DataFrame(new_rows).drop(columns=["group_key"]) -def calculatePath(approved_hashes, split): - if split: - dfs_by_policy = [group for _, group in approved_hashes.groupby("policy")] - else: - dfs_by_policy = [approved_hashes] - badpathparts = load_env_json("BAD_PATH_PARTS", "[]") - min_files_for_path = get_protected_value("MIN_FILES_FOR_PATH", cast_type = int) +async def calculatePath(approved_hashes, split): + dfs_by_policy = [group for _, group in approved_hashes.groupby("policy")] if split else [approved_hashes] + badpathparts = await asyncio.to_thread(load_env_json, "BAD_PATH_PARTS", "[]") + min_files_for_path = await asyncio.to_thread(get_protected_value, "MIN_FILES_FOR_PATH", int) processed_dfs = [] for df in dfs_by_policy: - haslcp = splitFilepathsGrouped(df, "filename_exec") + haslcp = await splitFilepathsGrouped(df, "filename_exec") haslcp = haslcp.drop_duplicates() forbidden = regulator(badpathparts, True) @@ -329,18 +331,12 @@ def calculatePath(approved_hashes, split): logger.debug("Removing forbidden filepaths for path exceptions") print(colorText("Removing forbidden filepaths for path exceptions", "green")) + lcp_not_forbidden = haslcp[~forbidden_lcfp].copy() - lcp_not_forbidden_review = lcp_not_forbidden[ - [ - "policyname", - "longestcfp", - "middle", - "filename_only", - "file_extension", - "sha256", - ] - ] + lcp_not_forbidden_review = lcp_not_forbidden[[ + "policyname", "longestcfp", "middle", "filename_only", "file_extension", "sha256" + ]] unique_sha_counts = ( lcp_not_forbidden_review.groupby("longestcfp")["sha256"].nunique().reset_index() @@ -353,8 +349,8 @@ def calculatePath(approved_hashes, split): lcp_not_forbidden_review = lcp_not_forbidden_review[ lcp_not_forbidden_review["unique_sha256_count"] >= min_files_for_path ] + processed_dfs.append(lcp_not_forbidden_review) pathExclusions = pd.concat(processed_dfs, ignore_index=True) - - return pathExclusions + return pathExclusions \ No newline at end of file diff --git a/flows/quietAgent.py b/flows/quietAgent.py index 6e8a642..9a85bea 100644 --- a/flows/quietAgent.py +++ b/flows/quietAgent.py @@ -1,124 +1,84 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - import datetime import logging -import dotenv import pandas as pd from flows.prepPolicy import selectPolicies from services.API import AirlockAPIWrapper -from services.policyhandler import getPolicyInfo -from utils.selector import Selector +from services.PolicyHandler import getPolicyInfo +from utils.Selector import Selector from utils.utils import colorText, load_env logger = logging.getLogger(__name__) -dotenv.load_dotenv() - - -def findQuietAgents(api: AirlockAPIWrapper): +async def findQuietAgents(api: AirlockAPIWrapper): working_dir = load_env("WORKING_DIR") - # Get policy selection and agent list - selected_policy = selectPolicies(api, False) - if selected_policy: - agents = api.agents_find_by_group(selected_policy[0].groupid) + selected_policy = await selectPolicies(api, False) - # Prompt user for history range - history_days = Selector.select_value( + if selected_policy: + agents = await api.agents_find_by_group(selected_policy[0].groupid) + + history_days = await Selector.select_value( prompt="Enter how many days of history to pull (1–150): ", value_type=int, valid_range=(1, 150), ) - required_quiet = Selector.select_value( - prompt="Enter how many days without an untrusted execution before these are considered ready for enforcement? (1–365): ", + required_quiet = await Selector.select_value( + prompt="Enter how many days without an untrusted execution before these are considered ready for enforcement? (1–150): ", value_type=int, valid_range=(1, 150), ) - # Get execution history as a DataFrame - policy_exec_history = getPolicyInfo( + policy_exec_history = await getPolicyInfo( api, selected_policy[0], [1, 2, 6, 7], history_days ) - - + if policy_exec_history.empty: logging.info("No execution history found for the selected policy and time range.") return - - # Convert 'datetime' column to timezone-aware datetime objects policy_exec_history["datetime"] = pd.to_datetime( policy_exec_history["datetime"], format="%Y-%m-%dT%H:%M:%SZ", utc=True ) - # Get current UTC time now = datetime.datetime.now(datetime.timezone.utc) - - # Calculate days ago policy_exec_history["days_ago"] = policy_exec_history["datetime"].apply( lambda dt: (now - dt).days ) - # Count total executions per hostname hostname_counts = policy_exec_history["hostname"].value_counts() - - # Map execution counts to agents agents["execution_count"] = agents["hostname"].map(hostname_counts).fillna(0).astype(int) - # Find most recent execution per hostname most_recent_exec = policy_exec_history.sort_values(by="days_ago").drop_duplicates( subset="hostname", keep="first" ) - - # Map most recent execution age to agents agents["days_since"] = agents["hostname"].map( most_recent_exec.set_index("hostname")["days_ago"] ) - # Check for enforcement readiness agents["required_quiet"] = required_quiet agents["enforce_ready"] = agents["days_since"].apply( lambda x: True if pd.isna(x) or x > required_quiet else False ) - # Sort agents by execution count and hostname agents = agents.sort_values(by=["execution_count", "hostname"], ascending=[True, True]) - # Save to CSV filename = f"{working_dir}\\{selected_policy[0].name}_agents_last_{history_days}_days.csv" logging.debug(f"Saving CSV to {filename}") print(colorText(f"Saving CSV to {filename}", "green")) agents.to_csv(filename, index=False) - # Summary statistics total_agents = len(agents) ready_agents = agents["enforce_ready"].sum() not_ready_agents = total_agents - ready_agents ready_percentage = (ready_agents / total_agents) * 100 - # Print results - - message = ( f"Total agents: {total_agents}\n" f"Agents marked as 'enforce_ready': {ready_agents}\n" f"Agents not ready: {not_ready_agents}\n" f"Percentage ready for enforcement: {ready_percentage:.2f}%" ) - logger.debug(message) - colorText(message,"green") + logger.info(message) + colorText(message, "green") \ No newline at end of file diff --git a/flows/recursivePolicyEnforce.py b/flows/recursivePolicyEnforce.py new file mode 100644 index 0000000..f29d6d3 --- /dev/null +++ b/flows/recursivePolicyEnforce.py @@ -0,0 +1,108 @@ +""" +MAJOR WORK IN PROGRESS + +from flows.prepPolicy import ( + buildPathsandPublishers, + buildPreflights, + selectAllowlists, + selectPolicies, + sortHashes, + ) + + + +#Pull History for last 24 hours, Make a list of unique policy names that had exectutions + +sortHashes(api,selected_policies, type=[1, 2, 6, 7], history_days=150) + + + + + + +buildPathsandPublishers(False) + +buildPreflights() + else: + print("File not found. Please make sure it's saved correctly and try again.") + + + + pathexclusions = pd.read_csv(f"{working_dir}\\Preflight\\approved_paths.csv") + hashes = pd.read_csv(f"{working_dir}\\Preflight\\approved_hashes.csv") + + unique_combinations = pathexclusions[["longestcfp", "file_extension"]].drop_duplicates() + + drive_letter_pattern = re.compile(r"^[a-zA-Z]:\\") + processed_paths = [ + (path if drive_letter_pattern.match(path) else f"\\\\{path}") + f"\\**{ext}" + for path, ext in unique_combinations.itertuples(index=False, name=None) + ] + + print(processed_paths) + print(colorText("These publishers would added", "yellow")) + + if os.path.exists(f"{working_dir}\\Preflight\\approved_publishers.csv"): + publishers = pd.read_csv(f"{working_dir}\\Preflight\\approved_hashes.csv") + if publishers.empty: + print(colorText("The publishers list is empty.", "red")) + else: + processed_publishers = ( + publishers[publishers["publisher_hash"] != "Not Signed"] + ["publisher_hash"] + .drop_duplicates() + .tolist() + ) + print(processed_publishers) + + print(colorText("These hashes would be added to:", "yellow")) + print(destination_allowlist) + + processed_hashes = hashes["sha256"].unique().tolist() + print(processed_hashes) + + if processed_paths and processed_hashes: + tested = True + else: + # Log which condition(s) failed + missing_items = [] + if not os.path.exists(f"{working_dir}\\Preflight\\approved_paths.csv"): + missing_items.append("approved_paths.csv not found") + if not os.path.exists(f"{working_dir}\\Preflight\\approved_hashes.csv"): + missing_items.append("approved_hashes.csv not found") + if not destination_policy: + missing_items.append("destination_policy is empty or None") + if not destination_allowlist: + missing_items.append("destination_allowlist is empty or None") + + logger.error("Preflight check failed due to the following:") + for item in missing_items: + logger.error(f" - {item}") + + + elif choice == "7": + areYouSure() + confirmation = get_sanitized_input("Type 'I AGREE' to continue: ") + if ( + tested + and destination_policy + and destination_allowlist + and confirmation.strip() == "I AGREE" + ): + print(colorText("Proceeding with the code...", "yellow")) + api.hash_add_to_allowlist(destination_allowlist[0].applicationid, processed_hashes) + api.policy_add_path_exclusions(destination_policy[0].groupid, processed_paths) + if processed_publishers: + api.policy_add_publishers(destination_policy[0].groupid, processed_publishers) + else: + logger.error("Confirmation block failed. Reasons:") + if not tested: + logger.error(" - Preflight checks were not completed successfully (`tested` is False).") + if not destination_policy: + logger.error(" - `destination_policy` is missing or invalid.") + if not destination_allowlist: + logger.error(" - `destination_allowlist` is missing or invalid.") + if confirmation.strip() != "I AGREE": + logger.error(" - User did not confirm with 'I AGREE'. Received: '%s'", confirmation.strip()) + + """ \ No newline at end of file diff --git a/models/execution.py b/models/execution.py index a7466f6..27f1130 100644 --- a/models/execution.py +++ b/models/execution.py @@ -1,18 +1,4 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - +import asyncio import inspect import json import logging @@ -22,46 +8,20 @@ from dataclasses import dataclass from datetime import datetime from typing import List, Optional -import dotenv +import aiofiles import pandas as pd -from services.policyhandler import pullPolicyExechistories +from services.PolicyHandler import pullPolicyExechistories from utils.configmanager import get_protected_value, load_env_json -from utils.utils import colorText, regulator +from utils.utils import regulator logger = logging.getLogger(__name__) -dotenv.load_dotenv() - - class Hash: - """ - Hash model representing Hash data - """ - - def __init__( - self, - sha256, - applications=None, - baselines=None, - blocklists=None, - createtime=None, - datetime=None, - description=None, - filename=None, - filepath=None, - filesize=None, - md5=None, - modtime=None, - origname=None, - productname=None, - productversion=None, - publisher=None, - reputation=None, - sha128=None, - sha384=None, - sha512=None, - ): + def __init__(self, sha256, applications=None, baselines=None, blocklists=None, createtime=None, + datetime=None, description=None, filename=None, filepath=None, filesize=None, + md5=None, modtime=None, origname=None, productname=None, productversion=None, + publisher=None, reputation=None, sha128=None, sha384=None, sha512=None): self.sha256 = sha256 self.applications = applications self.baselines = baselines @@ -88,35 +48,23 @@ class Hash: return f"" def __eq__(self, other): - if isinstance(other, Hash): - return self.sha256 == other.sha256 - return False + return isinstance(other, Hash) and self.sha256 == other.sha256 def __hash__(self): return hash(self.sha256) def to_dict(self): - """Returns a dictionary representation of the hash.""" return self.__dict__ @staticmethod def safe_int(value, default=0): - """Safely convert a value to int, returning default on failure.""" try: return int(value) except (TypeError, ValueError): return default - @classmethod def deduplicate(cls, hash_list): - """ - Deduplicates a list of Hash objects based on sha256. - Args: - hash_list (list): List of Hash instances. - Returns: - list: Deduplicated list of Hash instances. - """ seen = set() deduped = [] for h in hash_list: @@ -126,14 +74,11 @@ class Hash: return deduped @classmethod - def categorize_hashes(cls, hashes): - threat_tolerance = get_protected_value("VT_THREAT_TOLERANCE", cast_type=int) - bad_publishers_pattern = regulator(load_env_json("BAD_PUBLISHERS", "[]")) - pups_pattern = regulator(load_env_json("PUPS", "[]")) - - needs_review = [] - approved = [] - unapproved = [] + async def categorize_hashes(cls, hashes): + threat_tolerance = await get_protected_value("VT_THREAT_TOLERANCE", cast_type=int) + bad_publishers_pattern = regulator(await load_env_json("BAD_PUBLISHERS", "[]")) + pups_pattern = regulator(await load_env_json("PUPS", "[]")) + needs_review, approved, unapproved = [], [], [] for hash_obj in hashes: publisher = hash_obj.publisher or "" @@ -141,57 +86,32 @@ class Hash: reputation = hash_obj.reputation if isinstance(hash_obj.reputation, dict) else {} scannermatch = reputation.get("scannermatch") - logger.debug(f"Evaluating hash: {hash_obj}") - logger.debug(f"Publisher: {publisher}, Description: {description}, Scannermatch: {scannermatch}") - - # 1. Unapproved: bad publisher or PUP if re.search(bad_publishers_pattern, publisher, re.IGNORECASE): - logger.debug("Unapproved: Publisher matches bad publisher pattern.") unapproved.append(hash_obj) continue - if re.search(pups_pattern, description, re.IGNORECASE): - logger.debug("Unapproved: Description matches PUP pattern.") unapproved.append(hash_obj) continue - - # 2. Approved: signed if publisher != "Not Signed": - logger.debug("Approved: File is signed and not flagged.") approved.append(hash_obj) continue - - # 3. Approved or Unapproved based on threat level try: score = int(scannermatch) # pyright: ignore[reportArgumentType] - logger.debug(f"Parsed scannermatch score: {score}") - if score > threat_tolerance: # pyright: ignore[reportOperatorIssue] - logger.debug("Unapproved: Unsigned file with high threat score.") + if score > threat_tolerance: # type: ignore unapproved.append(hash_obj) else: - logger.debug("Approved: Unsigned file with low threat score.") approved.append(hash_obj) except (ValueError, TypeError): - logger.debug("Needs Review: Scannermatch score is missing or invalid.") needs_review.append(hash_obj) - - logger.debug(f"Final counts — Needs Review: {len(needs_review)}, Approved: {len(approved)}, Unapproved: {len(unapproved)}") return needs_review, approved, unapproved - @classmethod - def export_to_csv(cls, hash_list, directory_path): - """ - Exports a list of Hash objects to a CSV file in the specified directory. - The filename is derived from the variable name of the list if possible, - and includes a timestamp to ensure uniqueness. - """ + async def export_to_csv(cls, hash_list, directory_path): filename = "hashes_export.csv" frame = inspect.currentframe() if frame is not None and frame.f_back is not None: - callers_local_vars = frame.f_back.f_locals.items() - for var_name, var_val in callers_local_vars: + for var_name, var_val in frame.f_back.f_locals.items(): if var_val is hash_list: timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") filename = f"{var_name}_{timestamp}.csv" @@ -203,52 +123,16 @@ class Hash: os.makedirs(directory_path, exist_ok=True) file_path = os.path.join(directory_path, filename) - df = pd.DataFrame([h.to_dict() for h in hash_list]) - df.to_csv(file_path, index=False) + df = await asyncio.to_thread(pd.DataFrame, [h.to_dict() for h in hash_list]) + await asyncio.to_thread(df.to_csv, file_path, index=False) - logger.info(f"CSV file saved to: {file_path}") - - -""" -#Example - Convert Dataframe returned by hash query into hash objects -hash_objects = [] -for _, row in df.iterrows(): - try: - parsed_data = ast.literal_eval(row['data']) - hash_obj = Hash(sha256=row['sha256'], **parsed_data) - hash_objects.append(hash_obj) - except Exception as e: - print(f"Error parsing row: {e}") - -# Display the created Hash objects -for obj in hash_objects: - print(obj) - - -# Categorize hashes -needs_review, approved, unapproved = Hash.categorize_hashes( - hashes=hash_objects, - threat_tolerance=3, - untrusted_pattern=untrusted_pattern, - pups_pattern=pups_pattern -) - - # Deduplicate - deduped_hashes = Hash.deduplicate(hash_list) - -# Specify the directory where you want to save the CSV -output_directory = "C:/Users/Brandon/Documents/HashExports" - -# Call the export method -Hash.export_to_csv(hashes_for_export, output_directory) - - -""" + async with aiofiles.open(file_path, mode='r') as f: + preview = await f.read() + print(f"CSV file saved to: {file_path}\nPreview:\n{preview[:500]}") @dataclass class ExecutionHistoryRecord: - # Mandatory fields username: str hostname: str netdomain: str @@ -260,8 +144,6 @@ class ExecutionHistoryRecord: publisher: str sha256: str datetime: str - - # Optional fields type: Optional[int] = None pprocess: Optional[str] = None gprocess: Optional[str] = None @@ -273,62 +155,69 @@ class ExecutionHistoryRecord: localip: Optional[str] = None extid: Optional[str] = None extname: Optional[str] = None - exttype: Optional[int] = None # 1 = CRX Chromium Extension, 2 = XPI Firefox Extension - extbrowser: Optional[int] = None # 1 = Chrome, 2 = Firefox, 3 = Edge + exttype: Optional[int] = None + extbrowser: Optional[int] = None + + @classmethod + async def from_policies(cls, api, selected_policies, type_: list, history_days: int) -> List["ExecutionHistoryRecord"]: + async def fetch_and_parse(policy): + execs = await pullPolicyExechistories(api, policy, type_, history_days, True) + if not execs: + return [] + data = json.loads(execs) + exechistories = data.get("response", {}).get("exechistories", []) + if not exechistories: + return [] + df = await asyncio.to_thread(pd.DataFrame, exechistories) + df = await asyncio.to_thread(df.drop_duplicates, subset=["sha256", "filename", "hostname"]) + df = await asyncio.to_thread(df.sort_values, by=["sha256", "filename"]) + return [cls.from_dict(row.to_dict()) for _, row in df.iterrows()] + + tasks = [fetch_and_parse(policy) for policy in selected_policies] + results = await asyncio.gather(*tasks) + return [record for sublist in results for record in sublist] @staticmethod - def enrich_with_hashes_and_export( - executions: list, hashes: list, directory_path: str, label: str = "enriched" - ): - exec_df = pd.DataFrame([e.__dict__ for e in executions]) - hash_df = pd.DataFrame([h.to_dict() for h in hashes]) + async def enrich_with_hashes(executions: list, hashes: list): - logger.debug(f"Execution DataFrame columns: {exec_df.columns}") - logger.debug(f"Hash DataFrame columns: {hash_df.columns}") + exec_task = asyncio.to_thread(pd.DataFrame, [e.__dict__ for e in executions]) + hash_task = asyncio.to_thread(pd.DataFrame, [h.to_dict() for h in hashes]) + + exec_df, hash_df = await asyncio.gather(exec_task, hash_task) if hash_df.empty: - logger.warning(f"hash_df is empty for label: {label}. Skipping merge.") merged_df = exec_df.copy() + logger.debug("Hash dataframe appears empty") else: - merged_df = pd.merge( + merged_df = await asyncio.to_thread( + pd.merge, exec_df, hash_df, on="sha256", - how="left", # Preserve all executions, enrich where possible + how="left", suffixes=("_exec", "_hash") ) - merged_df.sort_values(by="filename_exec", inplace=True) - logger.info(f"Merged {len(merged_df)} rows. Non-null hash matches: {merged_df['sha256'].notna().sum()}") - filename = f"{label}_executions.csv" - os.makedirs(directory_path, exist_ok=True) - file_path = os.path.join(directory_path, filename) - merged_df.to_csv(file_path, index=False) - - logger.info(f"CSV file saved to: {file_path}") + # Log available columns for debugging + logger.debug(f"Merged DataFrame columns: {merged_df.columns.tolist()}") + # Only sort if the column exists + if "filename_exec" in merged_df.columns: + merged_df = await asyncio.to_thread(merged_df.sort_values, by="filename_exec") + else: + merged_df = await asyncio.to_thread(merged_df.sort_values, by="filename") + return merged_df + @classmethod def from_dict(cls, data: dict): mandatory_fields = [ - "username", - "hostname", - "netdomain", - "filename", - "ppolicy", - "policyname", - "policyver", - "commandline", - "publisher", - "sha256", - "datetime", - ] - missing_fields = [ - field for field in mandatory_fields if field not in data or data[field] is None + "username", "hostname", "netdomain", "filename", "ppolicy", + "policyname", "policyver", "commandline", "publisher", "sha256", "datetime" ] + missing_fields = [field for field in mandatory_fields if field not in data or data[field] is None] if missing_fields: raise ValueError(f"Missing mandatory fields: {missing_fields}") - return cls( username=data["username"], hostname=data["hostname"], @@ -354,45 +243,4 @@ class ExecutionHistoryRecord: extname=data.get("extname"), exttype=data.get("exttype"), extbrowser=data.get("extbrowser"), - ) - - @classmethod - def from_policies( - cls, api, selected_policies, type_: list, history_days: int - ) -> List["ExecutionHistoryRecord"]: - executions = [] - for policy in selected_policies: - execs = pullPolicyExechistories( - api, policy, type_, history_days, True - ) - if execs: - data = json.loads(execs) - exechistories = data.get("response", {}).get("exechistories", []) - if not exechistories: - continue - - df = pd.DataFrame(exechistories) - df = df.drop_duplicates(subset=["sha256", "filename", "hostname"]) - df = df.sort_values(by=["sha256", "filename"]) - - executions.extend([cls.from_dict(row.to_dict()) for _, row in df.iterrows()]) - logger.debug(f"Staging of Execution history for policy: {policy.name} is complete") - print( - colorText( - f"Staging of Execution history for policy: {policy.name} is complete", - "green", - ) - ) - - return executions - - def __repr__(self): - attrs = ", ".join(f"{key}={repr(value)}" for key, value in self.__dict__.items()) - return f"" - - -""" -executions = ExecutionHistoryRecord.from_policies(api, selected_policies, type_=[0,1,3], history_days=30) - -ExecutionHistoryRecord.enrich_with_hashes_and_export(executions, hash_objects, "C:/Users/Brandon/Documents/EnrichedExports") -""" + ) \ No newline at end of file diff --git a/services/API.py b/services/API.py index 5742d0a..0866493 100644 --- a/services/API.py +++ b/services/API.py @@ -1,274 +1,175 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - import json import logging from typing import Dict, List, Optional +import httpx import pandas as pd -import requests logger = logging.getLogger(__name__) - - class AirlockAPIWrapper: - """ - A wrapper class for interacting with the Airlock API. - Provides methods for managing agents, policies, hashes, OTPs, and execution history. - """ - def __init__(self, base_url: str, api_key: str): - """ - Initialize the API wrapper. - - Parameters: - - base_url (str): Base URL of the Airlock API. - - api_key (str): API key for authentication. - """ self.base_url = base_url.rstrip("/") self.api_key = api_key self.headers = {"X-APIKey": self.api_key} - def _post(self, endpoint: str, payload: Optional[dict] = None) -> dict: - """ - Internal method to send POST requests to the API. - - Parameters: - - endpoint (str): API endpoint. - - payload (dict, optional): Request payload. - - Returns: - - dict: JSON response from the API. - """ + async def _post(self, endpoint: str, payload: Optional[dict] = None) -> dict: url = f"{self.base_url}{endpoint}" data = json.dumps(payload or {}) - try: - logger.debug(f"POST Request to {url} with payload: {payload}") - response = requests.post(url, headers=self.headers, data=data, verify=False) - response.raise_for_status() - logger.debug(f"Response received from {url}") - return response.json() - except requests.exceptions.RequestException as e: - logger.error(f"API request failed: {e}") - raise + timeout = httpx.Timeout(300.0) + async with httpx.AsyncClient(verify=False, timeout=timeout) as client: + try: + logger.debug(f"POST Request to {url} with payload: {payload}") + response = await client.post(url, headers=self.headers, data=data) # pyright: ignore[reportArgumentType] + response.raise_for_status() + logger.debug(f"Response received from {url}") + return response.json() + except httpx.RequestError as e: + logger.error(f"API request failed: {e}") + raise # Allowlist Management - def allowlist_find_all(self) -> pd.DataFrame: - """ - Retrieve all applications in the allowlist. - - Returns: - - pd.DataFrame: DataFrame containing allowlisted applications. - """ - result = self._post("/v1/application", {}) + async def allowlist_find_all(self) -> pd.DataFrame: + result = await self._post("/v1/application", {}) return pd.DataFrame(result["response"]["applications"]) # Agent Management - def agent_find_all(self) -> pd.DataFrame: - """Retrieve all agents.""" - result = self._post("/v1/agent/find", {}) + async def agent_find_all(self) -> pd.DataFrame: + result = await self._post("/v1/agent/find", {}) return pd.DataFrame(result["response"]["agents"]) - def agent_find_by_hostname(self, hostname: str) -> pd.DataFrame: - """Find agents by hostname.""" + async def agent_find_by_hostname(self, hostname: str) -> pd.DataFrame: payload = {"hostname": hostname} - result = self._post("/v1/agent/find", payload) + result = await self._post("/v1/agent/find", payload) return pd.DataFrame(result["response"]["agents"]) - def agent_find_by_id(self, agentid: str) -> pd.DataFrame: - """Find agents by agent ID.""" + async def agent_find_by_id(self, agentid: str) -> pd.DataFrame: payload = {"agentid": agentid} - result = self._post("/v1/agent/find", payload) + result = await self._post("/v1/agent/find", payload) return pd.DataFrame(result["response"]["agents"]) - def agent_find_by_status(self, status: int) -> pd.DataFrame: - """Find agents by status (0 = Offline, 1 = Online, 3 = Safemode).""" + async def agent_find_by_status(self, status: int) -> pd.DataFrame: payload = {"status": status} - result = self._post("/v1/agent/find", payload) + result = await self._post("/v1/agent/find", payload) return pd.DataFrame(result["response"]["agents"]) - def agent_find_by_username(self, username: str) -> pd.DataFrame: - """Find agents by username.""" + async def agent_find_by_username(self, username: str) -> pd.DataFrame: payload = {"username": username} - result = self._post("/v1/agent/find", payload) + result = await self._post("/v1/agent/find", payload) return pd.DataFrame(result["response"]["agents"]) - def agent_move(self, agentid: str, groupid: str) -> dict: - """Move an agent to a different group.""" + async def agent_move(self, agentid: str, groupid: str) -> dict: payload = {"agentid": agentid, "groupid": groupid} - return self._post("/v1/agent/move", payload) + return await self._post("/v1/agent/move", payload) - def agents_find_by_group(self, groupid: str) -> pd.DataFrame: - """Find agents by group ID.""" + async def agents_find_by_group(self, groupid: str) -> pd.DataFrame: payload = {"groupid": groupid} - result = self._post("/v1/agent/find", payload) + result = await self._post("/v1/agent/find", payload) return pd.DataFrame(result["response"]["agents"]) # Hash Management - def hash_add_to_allowlist(self, applicationid: str, hashes: List[str]) -> dict: - """Add hashes to the allowlist for a specific application.""" + async def hash_add_to_allowlist(self, applicationid: str, hashes: List[str]) -> dict: payload = {"applicationid": applicationid, "hashes": hashes} - return self._post("/v1/hash/application/add", payload) + return await self._post("/v1/hash/application/add", payload) - def hash_query(self, hashes: List[str]) -> pd.DataFrame: - """Query information about specific hashes.""" + async def hash_query(self, hashes: List[str]) -> pd.DataFrame: payload = {"hashes": hashes} - result = self._post("/v1/hash/query", payload) + result = await self._post("/v1/hash/query", payload) return pd.DataFrame(result["response"]["results"]) # OTP Management - def otp_find_active(self) -> pd.DataFrame: - """Find active OTPs.""" + async def otp_find_active(self) -> pd.DataFrame: payload = {"status": "1"} - result = self._post("/v1/otp/usage", payload) + result = await self._post("/v1/otp/usage", payload) return pd.DataFrame(result["response"]["otpusage"]) - def otp_find_awaiting(self) -> pd.DataFrame: - """Find OTPs that are awaiting activation.""" + async def otp_find_awaiting(self) -> pd.DataFrame: payload = {"status": "0"} - result = self._post("/v1/otp/usage", payload) + result = await self._post("/v1/otp/usage", payload) return pd.DataFrame(result["response"]["otpusage"]) - - def otp_find_by_agent(self, agentid) -> pd.DataFrame: - """Find OTP by agent.""" + + async def otp_find_by_agent(self, agentid) -> pd.DataFrame: payload = {"agentid": agentid} - result = self._post("/v1/otp/usage", payload) + result = await self._post("/v1/otp/usage", payload) return pd.DataFrame(result["response"]["otpusage"]) - - def otp_generate(self, agentid: str, duration: int, purpose: str) -> str: - """Generate a new OTP for an agent.""" + + async def otp_generate(self, agentid: str, duration: int, purpose: str) -> str: payload = { "duration": str(duration), "agentid": str(agentid), "purpose": purpose, } - result = self._post("/v1/otp/retrieve", payload) + result = await self._post("/v1/otp/retrieve", payload) return result["response"]["otpcode"] - def otp_get_activities(self, otpid: str) -> pd.DataFrame: - """Retrieve activities associated with a specific OTP.""" + async def otp_get_activities(self, otpid: str) -> pd.DataFrame: payload = {"otpid": otpid} - result = self._post("/v1/otp/activities", payload) + result = await self._post("/v1/otp/activities", payload) return pd.DataFrame(result["response"]["otpactivities"]) - - def otp_revoke(self, otpid: str) -> dict: - """ - Revoke an active OTP. - Parameters: - - otpid (str): The ID of the OTP to revoke. - Returns: - - dict: JSON response from the API. - """ - payload = {"otpid": otpid} - return self._post("/v1/otp/revoke", payload) - - def otp_validate(self, otpcode: str) -> dict: - """ - Validate an OTP code. - Parameters: - - otpcode (str): The OTP code to validate. - Returns: - - dict: JSON response indicating validity. - """ - payload = {"otpcode": otpcode} - return self._post("/v1/otp/validate", payload) + async def otp_revoke(self, otpid: str) -> dict: + payload = {"otpid": otpid} + return await self._post("/v1/otp/revoke", payload) + + async def otp_validate(self, otpcode: str) -> dict: + payload = {"otpcode": otpcode} + return await self._post("/v1/otp/validate", payload) # Policy Management - def policy_add_path_exclusions(self, groupid: str, paths: List[str]) -> dict: - """Add path exclusions to a policy group.""" + async def policy_add_path_exclusions(self, groupid: str, paths: List[str]) -> dict: payload = {"groupid": groupid, "path": paths} - return self._post("/v1/group/path/add", payload) + return await self._post("/v1/group/path/add", payload) - def policy_add_publishers(self, groupid: str, publishers: List[str]) -> dict: - """Add publishers to a policy group.""" + async def policy_add_publishers(self, groupid: str, publishers: List[str]) -> dict: payload = {"groupid": groupid, "publisher": publishers} - return self._post("/v1/group/publisher/add", payload) + return await self._post("/v1/group/publisher/add", payload) - def policy_clone(self, source_groupid: str, target_groupid: str) -> dict: - """Clone a policy from one group to another.""" + async def policy_clone(self, source_groupid: str, target_groupid: str) -> dict: payload = {"groupid": source_groupid, "targetgroupid": target_groupid} - return self._post("/v1/group/assign", payload) + return await self._post("/v1/group/assign", payload) - def policy_find_all(self) -> pd.DataFrame: - """Retrieve all policy groups.""" - result = self._post("/v1/group") + async def policy_find_all(self) -> pd.DataFrame: + result = await self._post("/v1/group") return pd.DataFrame(result["response"]["groups"]) - def policy_list_agents(self, groupid: str) -> pd.DataFrame: - """List agents assigned to a specific policy group.""" + async def policy_list_agents(self, groupid: str) -> pd.DataFrame: payload = {"groupid": groupid} - result = self._post("/v1/group/agents", payload) + result = await self._post("/v1/group/agents", payload) return pd.DataFrame(result["response"]["agents"]) - - def policy_list_allowlists(self, groupid: str) -> pd.DataFrame: - """List allowlists assigned to a specific policy group.""" + + async def policy_list_allowlists(self, groupid: str) -> pd.DataFrame: payload = {"groupid": groupid} - result = self._post("/v1/group/policies", payload) + result = await self._post("/v1/group/policies", payload) return pd.DataFrame(result["response"]["applications"]) - def policy_set_auditmode(self, groupid: str, auditmode: str) -> dict: - """Set audit mode for a policy group. 1=Audit, 0=Enforcement""" + async def policy_set_auditmode(self, groupid: str, auditmode: str) -> dict: payload = {"groupid": groupid, "auditmode": auditmode} - return self._post("/v1/group/settings/auditmode", payload) + return await self._post("/v1/group/settings/auditmode", payload) + + async def policy_set_script_custom( + self, + groupid: str, + script_custom: int, + scripts_audit: Optional[List[str]] = None, + scripts_disabled: Optional[List[str]] = None, + scripts_respect: Optional[List[str]] = None + ) -> dict: + payload = { + "groupid": groupid, + "script_custom": script_custom, + "scripts_audit": scripts_audit or [], + "scripts_disabled": scripts_disabled or [], + "scripts_respect": scripts_respect or [] + } + return await self._post("/v1/group/settings/script_custom", payload) # Execution History - def history_logging(self, type: List[str], checkpoint: str, policy: List[str]) -> str: - """Retrieve execution history logs.""" + async def history_logging(self, type: List[str], checkpoint: str, policy: List[str]) -> List[Dict]: payload = {"type": type, "checkpoint": checkpoint, "policy": policy} - result = self._post("/v1/logging/exechistories", payload) + result = await self._post("/v1/logging/exechistories", payload) return result["response"]["exechistories"] - def history_execution(self, today: str, date_selected: str, agent_name: str) -> List[Dict]: - """ - Retrieve execution history logs. - - "datefrom":"", //(Optional) Datefrom is for date range search, formatted as "YYYY-MM-DD" - "dateto":"", //(Optional) Dateto is for date range search, formatted as "YYYY-MM-DD" - "category":"", //(Optional) Category for filtering type - "hostname":"", //(Optional) Hostname to filter - "username":"admin", //(Optional) Username to filter - "netdomain":"", //(Optional) Domain (or group) to filter - "filename":"", //(Optional) Filename to filter - "ppolicy":"", //(Optional) Parent Policy name to filter - "policyname":"", //(Optional) Policy name to filter - "policyver":"", //(Optional) Policy version to filter (e.g. "v95") - "commandline":"", //(Optional) Commandline to filter - "publisher":"", //(Optional) Publisher to filter - "pprocess":"", //(Optional) Parent Process to filter - "sha256":"", //(Optional) SHA256 hash to filter - "contains":["hostname"], //(Optional) Contains is an array for wildcard searches on a filter - "limit":"5" //(Optional) Limit the amount of results returned, default set to 50 - - """ - + async def history_execution(self, today: str, date_selected: str, agent_name: str) -> List[Dict]: payload = {"datefrom": date_selected, "dateto": today, "hostname": agent_name} - result = self._post("/v1/getexechistory", payload) - return result["response"]["exechistory"] - - -""" -from services.API import AirlockAPIWrapper - - -api = AirlockAPIWrapper(base_url="https://airlock.example.com/api", api_key="your_api_key_here") - -#Example: Get all agents - -agents_df = api.agent_find_all() -print("All Agents:") -print(agents_df) -""" + result = await self._post("/v1/getexechistory", payload) + return result["response"]["exechistory"] \ No newline at end of file diff --git a/services/TaskQueue.py b/services/TaskQueue.py new file mode 100644 index 0000000..0260fd6 --- /dev/null +++ b/services/TaskQueue.py @@ -0,0 +1,108 @@ +import asyncio +import logging +from typing import Any, Callable + +logger = logging.getLogger(__name__) + + +class AsyncTaskQueue: + def __init__(self, worker_count: int = 3): + self.queue = asyncio.Queue() + self.worker_count = worker_count + self.workers = [] + self._stop_event = asyncio.Event() + + async def start_workers(self): + """Start the worker pool.""" + logger.debug(f"Starting {self.worker_count} workers...") + for i in range(self.worker_count): + worker = asyncio.create_task(self.worker_loop(f"Worker-{i+1}")) + self.workers.append(worker) + logger.debug("All workers started.") + + async def stop_workers(self): + """Stop the worker pool and wait for all tasks to complete.""" + logger.debug("Stopping workers...") + self._stop_event.set() # Signal workers to stop + await self.queue.join() # Wait for all tasks to be processed + for worker in self.workers: + worker.cancel() + await asyncio.gather(*self.workers, return_exceptions=True) + logger.debug("All workers stopped.") + + + async def worker_loop(self, name: str): + """Worker loop: Process tasks from the queue.""" + logger.debug(f"{name} started.") + while not self._stop_event.is_set(): + task = None + try: + task = await self.queue.get() + logger.info(f"{name} processing: {task['name']}") + await task['func'](*task['args'], **task['kwargs']) + except asyncio.CancelledError: + logger.debug(f"{name} received cancellation.") + break + except Exception as e: + logger.error(f"Error in {name}: {e}", exc_info=True) + finally: + if task is not None: + self.queue.task_done() + logger.debug(f"{name} exited.") + + + + async def enqueue(self, name: str, func: Callable, *args: Any, **kwargs: Any): + """Add a task to the queue.""" + logger.debug(f"Enqueuing task: {name}") + await self.queue.put({'name': name, 'func': func, 'args': args, 'kwargs': kwargs}) + + + +async def run_sync_task_in_thread(func: Callable, *args: Any, **kwargs: Any): + """Run a synchronous function in a separate thread. + + Args: + func: The synchronous function to run. + *args: Arguments to pass to the function. + """ + await asyncio.to_thread(func, *args, **kwargs) + + """ +from asyncTaskQueue import AsyncTaskQueue, run_sync_task_in_thread +import asyncio + +# Example async task +async def your_async_function(name: str, duration: int): + print(f"{name} started, will sleep for {duration} seconds") + await asyncio.sleep(duration) + print(f"{name} finished") + +# Example sync task +def your_sync_function(name: str, duration: int): + print(f"{name} started, will sleep for {duration} seconds") + import time + time.sleep(duration) + print(f"{name} finished") + +async def main(): + queue = AsyncTaskQueue(worker_count=2) + await queue.start_workers() + + # Enqueue async tasks directly + await queue.enqueue("AsyncTask1", your_async_function, "AsyncTask1", 2) + await queue.enqueue("AsyncTask2", your_async_function, "AsyncTask2", 1) + + # Enqueue sync tasks using the wrapper + await queue.enqueue("SyncTask1", run_sync_task_in_thread, your_sync_function, "SyncTask1", 2) + + # Let tasks run for a while + await asyncio.sleep(5) + + # Stop workers + await queue.stop_workers() + +asyncio.run(main()) + + +""" \ No newline at end of file diff --git a/services/agenthandler.py b/services/agenthandler.py index 290baba..4682413 100644 --- a/services/agenthandler.py +++ b/services/agenthandler.py @@ -1,19 +1,4 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - - +import asyncio import json import logging import os @@ -27,16 +12,16 @@ import pandas as pd from models.agent import Agent from models.policy import Policy from services.API import AirlockAPIWrapper +from services.TaskQueue import AsyncTaskQueue, run_sync_task_in_thread from utils.configmanager import get_protected_json, load_env -from utils.selector import Selector +from utils.Selector import Selector from utils.utils import colorText, get_sanitized_input logger = logging.getLogger(__name__) - -def devicehistory(api: AirlockAPIWrapper, outputjson: bool): - agents = selectAgents(api) - history_days = Selector.select_value( +async def devicehistory(api: AirlockAPIWrapper, outputjson: bool): + agents = await selectAgents(api) + history_days = await Selector.select_value( prompt="Enter how many days of history to pull (1–150): ", value_type=int, valid_range=(1, 150), @@ -48,63 +33,68 @@ def devicehistory(api: AirlockAPIWrapper, outputjson: bool): historical_date = (datetime.now() - timedelta(days=history_days)).strftime("%Y-%m-%d") today = datetime.now().strftime("%Y-%m-%d") - all_history = [] - for agent in agents: + async def fetch_history(agent): try: - exechistory = api.history_execution(today, historical_date, agent.hostname) + exechistory = await api.history_execution(today, historical_date, agent.hostname) + if isinstance(exechistory, list): + for block in exechistory: + record = { + "Command": block.get("commandline", "N/A"), + "Date": block.get("datetime", "N/A"), + "Filename": block.get("filename", "N/A"), + "Policy Name": block.get("policyname", "N/A"), + "Hostname": block.get("hostname", "N/A"), + "Hash": block.get("sha256", "N/A"), + } + all_history.append(record) + if not outputjson: + for key, value in record.items(): + print(colorText(f"{key}: {value}", "green")) + print("\n") + else: + print(colorText(f"No execution history found for {agent.hostname}.", "yellow")) except Exception as e: print(colorText(f"āŒ Error retrieving history for {agent.hostname}: {e}", "red")) - continue - if isinstance(exechistory, list): - for block in exechistory: - record = { - "Command": block.get("commandline", "N/A"), - "Date": block.get("datetime", "N/A"), - "Filename": block.get("filename", "N/A"), - "Policy Name": block.get("policyname", "N/A"), - "Hostname": block.get("hostname", "N/A"), - "Hash": block.get("sha256", "N/A"), - } - all_history.append(record) - - if not outputjson: - for key, value in record.items(): - print(colorText(f"{key}: {value}", "green")) - print("\n") - else: - print(colorText(f"No execution history found for {agent.hostname}.", "yellow")) + await asyncio.gather(*(fetch_history(agent) for agent in agents)) if outputjson: print(json.dumps(all_history, indent=2)) +async def findAllAgents(api: AirlockAPIWrapper): + policies_df = await api.policy_find_all() + agents_df = await api.agent_find_all() -def findAllAgents(api): - # Step 1: Load data from API - policies = [Policy(**row["data"]) for _, row in api.policy_find_all().iterrows()] - agents = [Agent(**row["data"]) for _, row in api.agent_find_all().iterrows()] + policies = [Policy(**row["data"]) for _, row in policies_df.iterrows()] + agents = [Agent(**row["data"]) for _, row in agents_df.iterrows()] - # Step 2: Create groupid → groupname map groupid_to_name = {policy.groupid: policy.name for policy in policies} - # Step 3: Enrich agents - for agent in agents: + queue = AsyncTaskQueue() + await queue.start_workers() + + async def enrich_agent(agent): agent.enrich(groupid_to_name) + for agent in agents: + await queue.enqueue(f"enrich_{agent.hostname}", enrich_agent, agent) + + await asyncio.sleep(1) + await queue.stop_workers() + return agents -def findAgents(api, return_dataframe): - agents = selectAgents(api) - working_dir = load_env("WORKING_DIR") +async def findAgents(api: AirlockAPIWrapper, return_dataframe: bool): + agents = await selectAgents(api) + working_dir = await load_env("WORKING_DIR") if not agents: logging.warning("No agents or policies found.") print("No agents matched the criteria.") return - # Convert enriched agents to DataFrame agent_dicts = [asdict(agent) for agent in agents] agent_df = pd.DataFrame(agent_dicts) @@ -112,30 +102,24 @@ def findAgents(api, return_dataframe): logging.debug("Returning DataFrame to caller.") return agent_df - # Otherwise, print and optionally export print(agent_df) logging.debug("Displayed DataFrame to console.") - user_input = get_sanitized_input("\nWould you like to export the results to a CSV file? (y/n): ").strip().lower() + user_input = await get_sanitized_input("\nWould you like to export the results to a CSV file? (y/n): ") + user_input = user_input.strip().lower() if user_input == 'y': timestamp = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") filename = f"agentsearch_{timestamp}.csv" file_path = os.path.join(working_dir, filename) - agent_df.to_csv(file_path, index=False) - logging.info(f"Exported DataFrame to {file_path}") + await run_sync_task_in_thread(agent_df.to_csv, file_path, index=False) - print( - colorText( - f"\nāœ… Matched devices exported to: {working_dir}\\{filename}", - "green", - ) - ) + logging.info(f"Exported DataFrame to {file_path}") + print(colorText(f"\nāœ… Matched devices exported to: {working_dir}\\{filename}", "green")) else: logging.debug("User declined to export the DataFrame.") - -def selectAgents(api: AirlockAPIWrapper) -> List[Agent]: +async def selectAgents(api: AirlockAPIWrapper) -> List[Agent]: print(colorText("šŸ” Device Search", "cyan")) print(colorText("Enter the device hostnames you'd like to search for, one per line.", "cyan")) print(colorText("When you're done, press Enter twice.\n", "cyan")) @@ -144,50 +128,44 @@ def selectAgents(api: AirlockAPIWrapper) -> List[Agent]: print(colorText("UTN00000", "cyan")) print(colorText("i-hSuperSecretServer", "cyan")) print(colorText("u-hVenderBroke\n", "cyan")) - print(colorText("Paste or type your device names below:", "white")) - policies = [Policy(**row.to_dict()) for _, row in api.policy_find_all().iterrows()] + + policies_df = await api.policy_find_all() + policies = [Policy(**row.to_dict()) for _, row in policies_df.iterrows()] device_input_lines = [] empty_line_count = 0 - - # Regex to validate each line valid_line_pattern = re.compile(r'^[a-zA-Z0-9_\- ]+$') while True: - line = get_sanitized_input("") + line = await get_sanitized_input("") stripped_line = line.strip() - if stripped_line == "": empty_line_count += 1 if empty_line_count == 2: break - continue # Don't validate empty lines + continue else: empty_line_count = 0 + if valid_line_pattern.match(stripped_line): + device_input_lines.append(stripped_line) + else: + print(colorText(f"āš ļø Invalid input: '{stripped_line}' — only letters, numbers, underscores, spaces, and hyphens are allowed.", "yellow")) - # Validate only non-empty lines - if valid_line_pattern.match(stripped_line): - device_input_lines.append(stripped_line) - else: - print(colorText(f"āš ļø Invalid input: '{stripped_line}' — only letters, numbers, underscores, spaces, and hyphens are allowed.", "yellow")) - device_names = [name for name in device_input_lines if name] if not device_names: logger.debug("No device names entered") print(colorText("āš ļø No device names entered.", "red")) return [] - # Build regex pattern to match hostnames - pattern = "|".join(map(re.escape, device_names)) + pattern = "\n".join(map(re.escape, device_names)) regex = re.compile(pattern, re.IGNORECASE) - # Fetch agents - agents = [Agent(**row.to_dict()) for _, row in api.agent_find_all().iterrows()] + agents_df = await api.agent_find_all() + agents = [Agent(**row.to_dict()) for _, row in agents_df.iterrows()] matched_agents = [agent for agent in agents if regex.search(agent.hostname)] matched_agents.sort(key=lambda agent: agent.hostname.lower()) - # Show unmatched unmatched = [name for name in device_names if not any(regex.search(agent.hostname) for agent in agents)] if unmatched: logger.debug(f"āš ļø No matches for: {', '.join(unmatched)}") @@ -197,30 +175,17 @@ def selectAgents(api: AirlockAPIWrapper) -> List[Agent]: logger.debug("āŒ No matching devices found.") print(colorText("āŒ No matching devices found.", "red")) else: - logger.debug(f"āœ… Found {len(matched_agents)} matching device(s).") + logger.debug(f"āœ… Found {len(matched_agents)} matching device(s).") print(colorText(f"āœ… Found {len(matched_agents)} matching device(s).", "green")) - # Enrich each agent using its class method for agent in matched_agents: agent.enrich_with_policies(policies) return matched_agents -def moveAgentToRelatedPolicy( - api: AirlockAPIWrapper, - agent: Agent, - mode: str = "audit", -): - """ - Moves an agent between audit and enforcement policies based on the mode. - - Args: - api: AirlockAPIWrapper instance. - agent: Agent object. - policy_relationship_map: Dict mapping enforcement → audit. - mode: 'audit' to move to audit, 'enforcement' to move to enforcement. - """ - policy_relationship_map = get_protected_json("POLICY_MAP_ENF_AUD", "{}") +async def moveAgentToRelatedPolicy(api: AirlockAPIWrapper, agent: Agent, mode: str = "audit"): + policy_relationship_map = await get_protected_json("POLICY_MAP_ENF_AUD", "{}") + #TODO - Have this return the policy name it was moved to instead of the groupid if mode == "audit": if agent.groupid in policy_relationship_map: target_policy = policy_relationship_map[agent.groupid] @@ -231,7 +196,6 @@ def moveAgentToRelatedPolicy( else: logger.warning(f"Error: No corresponding audit policy found for groupid: {agent.groupid}.") return - elif mode == "enforcement": inverse_map = {v: k for k, v in policy_relationship_map.items()} if agent.groupid in inverse_map: @@ -242,9 +206,11 @@ def moveAgentToRelatedPolicy( else: logger.warning(f"Error: No corresponding enforcement policy found for groupid: {agent.groupid}.") return - else: logger.error(f"Unknown mode '{mode}'. Use 'audit' or 'enforcement'.") return - api.agent_move(agent.agentid, target_policy) + result = await api.agent_move(agent.agentid, target_policy) + if result == {'error': 'Success'}: logger.info(f"{agent.hostname} has been moved to {target_policy}") + else: logger.debug(result) + \ No newline at end of file diff --git a/services/policyhandler.py b/services/policyhandler.py index f76b969..2c84422 100644 --- a/services/policyhandler.py +++ b/services/policyhandler.py @@ -1,29 +1,15 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - - +import asyncio import datetime import gc import json import logging import os -import sys +import uuid +import aiofiles import pandas as pd -import tqdm from bson import ObjectId +from tqdm.asyncio import tqdm_asyncio from models.policy import Policy from services.API import AirlockAPIWrapper @@ -35,181 +21,101 @@ logger = logging.getLogger(__name__) -def pullPolicyExechistories( - api: AirlockAPIWrapper, - policy: Policy, - type: list, - days, - outputjson: bool, -): - - file_path = f"{get_base_directory()}\\cache\\chunkinator.json" - - # Ensure the file exists - if not os.path.exists(file_path): - with open(file_path, "w") as file: - json.dump({"error": "Success", "response": {"exechistories": []}}, file) - logger.debug(f"File '{file_path}' has been created.") - else: - logger.debug(f"File '{file_path}' already exists.") - +async def pullPolicyExechistories(api: AirlockAPIWrapper, policy: Policy, type, days, outputjson): + file_path = f"{get_base_directory()}\\cache\\chunkinator_{policy.name}_{uuid.uuid4().hex}.json" checkpoint = str(skipback(days)) json_output = {"error": "Success", "response": {"exechistories": []}} - with tqdm.tqdm( - file=sys.stdout, - leave=True, - total=10000, - desc=f"Checkpoint Progress: {checkpoint}", - colour="blue", - initial=1, - ) as filebar: - with tqdm.tqdm( - file=sys.stdout, - leave=True, - total=100, - desc=f"Total of {policy} Complete: ", - ) as pbar: - while True: - histories = api.history_logging( - type=type, checkpoint=checkpoint, policy= [policy.name] - ) + if not os.path.exists(file_path): + async with aiofiles.open(file_path, "w") as file: + await file.write(json.dumps(json_output)) - # Ensure histories is a list of dictionaries - if not isinstance(histories, list) or not all( - isinstance(h, dict) for h in histories - ): - logger.error( - "Unexpected response format from API. Expected list of dictionaries." - ) - break + filebar = tqdm_asyncio(total=10000, desc=f"Checkpoint Progress: {checkpoint}", colour="blue") + pbar = tqdm_asyncio(total=100, desc=f"Total of {policy.name} Complete: ") - filebar.total = len(histories) + while True: + histories = await api.history_logging(type=type, checkpoint=checkpoint, policy=[policy.name]) + if not histories: + break - if not histories: - break + for index, history_item in enumerate(histories): + if "checkpoint" not in history_item or "datetime" not in history_item: + continue - for index, history_item in enumerate(histories): - if ( - "checkpoint" not in history_item - or "datetime" not in history_item - ): - continue # Skip malformed entries + if index == len(histories) - 1: + checkpoint = history_item["checkpoint"] + filebar.set_description(f"Checkpoint Progress: {checkpoint}") + break - # Update checkpoint on last item - if index == len(histories) - 1: - checkpoint = history_item["checkpoint"] # pyright: ignore[reportArgumentType] - filebar.desc = f"Checkpoint Progress: {checkpoint}" - break + try: + history_date = datetime.datetime.strptime( + history_item["datetime"].replace(" +0000 UTC", ""), + "%Y-%m-%dT%H:%M:%SZ" + ).date() + except ValueError: + continue - try: - history_date = datetime.datetime.strptime( - history_item["datetime"].replace(" +0000 UTC", ""), # pyright: ignore[reportArgumentType] - "%Y-%m-%dT%H:%M:%SZ", - ).date() - except ValueError: - continue # Skip if date format is invalid + if datetime.date.today() - datetime.timedelta(days=days) <= history_date: + json_output["response"]["exechistories"].append(history_item) - if ( - datetime.date.today() - datetime.timedelta(days=days) - ) <= history_date: - json_output["response"]["exechistories"].append(history_item) + filebar.update(1) + await asyncio.sleep(0) - filebar.update(1) - filebar.refresh() + # Deduplication + seen = {} + if os.path.exists(file_path): + async with aiofiles.open(file_path, "r") as file: + content = await file.read() + existing_data = json.loads(content) + combined = existing_data["response"]["exechistories"] + json_output["response"]["exechistories"] + else: + combined = json_output["response"]["exechistories"] - # Deduplicate entries - seen = {} - if os.path.exists(file_path): - with open(file_path, "r") as file: - existing_data = json.load(file) - combined = ( - existing_data["response"]["exechistories"] - + json_output["response"]["exechistories"] - ) - else: - combined = json_output["response"]["exechistories"] + for entry in combined: + key = (entry.get("sha256"), entry.get("filename"), entry.get("hostname")) + seen[key] = entry - for entry in combined: - key = ( - entry.get("sha256"), - entry.get("filename"), - entry.get("hostname"), - ) - seen[key] = entry + deduplicated = list(seen.values()) + async with aiofiles.open(file_path, "w") as file: + await file.write(json.dumps({"error": "Success", "response": {"exechistories": deduplicated}})) - deduplicated = list(seen.values()) - with open(file_path, "w") as file: - json.dump( - { - "error": "Success", - "response": {"exechistories": deduplicated}, - }, - file, - ) + json_output["response"]["exechistories"].clear() - json_output["response"]["exechistories"].clear() + try: + last_date = datetime.datetime.strptime( + history_item["datetime"].replace(" +0000 UTC", ""), # pyright: ignore[reportPossiblyUnboundVariable] + "%Y-%m-%dT%H:%M:%SZ" + ).date() + date_diff = datetime.date.today() - last_date + percentage_diff = (((days + 10) - date_diff.days) / (days + 10)) * 100 + pbar.n = round(percentage_diff) + pbar.set_description(f"Total of {policy.name} Complete: ") + except Exception: + pass - # Update progress bar based on last valid item - try: - last_date = datetime.datetime.strptime( - history_item["datetime"].replace(" +0000 UTC", ""), # type: ignore - "%Y-%m-%dT%H:%M:%SZ", - ).date() - date_diff = datetime.date.today() - last_date - percentage_diff = ( - ((days + 10) - date_diff.days) / (days + 10) - ) * 100 - pbar.n = round(percentage_diff) - pbar.set_description_str(f"Total of {policy} Complete: ") - pbar.refresh() - except Exception: - pass + filebar.n = 1 - filebar.n = 1 - - # Final output - with open(file_path, "r") as file: - final_output = json.load(file) + async with aiofiles.open(file_path, "r") as file: + final_output = await file.read() os.remove(file_path) - return json.dumps(final_output) if outputjson else None + return final_output if outputjson else None -def getPolicyInfo(api: AirlockAPIWrapper, policy, type, days): +async def getPolicyInfo(api, policy, type, days): executionhist_policy = pd.DataFrame() - exehist = pullPolicyExechistories(api, policy, type, days, True) + exehist = await pullPolicyExechistories(api, policy, type, days, True) if exehist is not None: data = json.loads(exehist) executionhist_policy = pd.DataFrame(data["response"]["exechistories"]) if not executionhist_policy.empty: executionhist_policy = executionhist_policy[ - [ - "datetime", - "sha256", - "publisher", - "filename", - "hostname", - "username", - "pprocess", - "gprocess", - "commandline", - ] + ["datetime", "sha256", "publisher", "filename", "hostname", "username", "pprocess", "gprocess", "commandline"] ] - executionhist_policy["policy"] = policy # Add policy column here - executionhist_policy = executionhist_policy.drop_duplicates( - subset=["sha256", "filename", "hostname"] - ) - executionhist_policy = executionhist_policy.sort_values( - by=["sha256", "filename"] - ) - logger.debug( f"Staging of Execution history for policy: {policy} is complete") - print( - colorText( - f"Staging of Execution history for policy: {policy} is complete", - "green", - ) - ) + executionhist_policy["policy"] = policy.name + executionhist_policy = executionhist_policy.drop_duplicates(subset=["sha256", "filename", "hostname"]) + executionhist_policy = executionhist_policy.sort_values(by=["sha256", "filename"]) + print(colorText(f"Staging of Execution history for policy: {policy.name} is complete", "green")) del data del exehist gc.collect() @@ -230,8 +136,8 @@ def skipback(days): return ObjectId(objectid_hex) -def updateAuditPoliciesFromEnforcementPolices(api: AirlockAPIWrapper): - policy_relationship_map = get_protected_json("POLICY_MAP_ENF_AUD", "{}") +async def updateAuditPoliciesFromEnforcementPolices(api): + policy_relationship_map = await get_protected_json("POLICY_MAP_ENF_AUD", "{}") for enforcement_policy, audit_policy in policy_relationship_map.items(): - api.policy_clone(enforcement_policy, audit_policy) - api.policy_set_auditmode(audit_policy, "1") + await api.policy_clone(enforcement_policy, audit_policy) + await api.policy_set_auditmode(audit_policy, "1") \ No newline at end of file diff --git a/services/security.py b/services/security.py index 3bb7e05..aa0cea6 100644 --- a/services/security.py +++ b/services/security.py @@ -1,18 +1,3 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - import base64 import logging import os @@ -25,10 +10,12 @@ from cryptography.hazmat.primitives import hashes from cryptography.hazmat.primitives.ciphers.aead import AESGCM from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC +from utils.utils import colorText + # Constants KDF_ITERATIONS = 200_000 -SALT_SIZE = 16 # 128-bit Salt -NONCE_SIZE = 12 # AES-GCM +SALT_SIZE = 16 # 128-bit Salt +NONCE_SIZE = 12 # AES-GCM KEY_SIZE = 32 # AES-256 @@ -54,7 +41,7 @@ def configure_keyring_backend(): raise EnvironmentError(f"Unsupported OS: {system}") -def store_api_key(service: str, username: str, api_key: str, password: str): +async def store_api_key(service: str, username: str, api_key: str, password: str): configure_keyring_backend() salt = os.urandom(SALT_SIZE) key = _derive_key(password.encode(), salt) @@ -66,7 +53,7 @@ def store_api_key(service: str, username: str, api_key: str, password: str): keyring.set_password(service, username, b64) -def retrieve_api_key(service: str, username: str, password: str) -> str: +async def retrieve_api_key(service: str, username: str, password: str) -> str: configure_keyring_backend() b64 = keyring.get_password(service, username) if b64 is None: @@ -81,41 +68,38 @@ def retrieve_api_key(service: str, username: str, password: str) -> str: return pt.decode() -def api_key_exists(service: str, username: str) -> bool: +async def api_key_exists(service: str, username: str) -> bool: configure_keyring_backend() return keyring.get_password(service, username) is not None def check_password_complexity(password: str) -> bool: - if len(password) < 12: - return False - if not re.search(r"[A-Z]", password): - return False - if not re.search(r"[a-z]", password): - return False - if not re.search(r"[0-9]", password): - return False - if not re.search(r"[^A-Za-z0-9]", password): - return False - return True + return ( + len(password) >= 12 + and bool(re.search(r"[A-Z]", password)) + and bool(re.search(r"[a-z]", password)) + and bool(re.search(r"[0-9]", password)) + and bool(re.search(r"[^A-Za-z0-9]", password)) + ) -def getAPI(USERNAME, SERVICE_NAME): +async def getAPI(USERNAME, SERVICE_NAME): logging.debug( f"Checking for stored API key for user '{USERNAME}' in service '{SERVICE_NAME}'..." ) - if api_key_exists(SERVICE_NAME, USERNAME): + if await api_key_exists(SERVICE_NAME, USERNAME): for attempt in range(1, 4): password = getpass(f"Attempt {attempt}/3 - Enter password to unlock your API key: ") try: - apikey = retrieve_api_key(SERVICE_NAME, USERNAME, password) + apikey = await retrieve_api_key(SERVICE_NAME, USERNAME, password) logging.debug("API key successfully retrieved.") return apikey except Exception as e: logging.warning(f"Attempt {attempt} failed: {str(e)}") logging.error("Failed to retrieve API key after 3 incorrect attempts.") - raise ValueError("Failed to retrieve API key after 3 incorrect attempts.") + print(colorText("āŒ Authentication failed. Exiting.", "red")) + exit(1) else: logging.warning(f"No API key found for user '{USERNAME}' in service '{SERVICE_NAME}'.") api_key = getpass(f"šŸ—ļø No API key found. Please enter your API key for '{SERVICE_NAME}': ").strip() @@ -131,7 +115,7 @@ def getAPI(USERNAME, SERVICE_NAME): if check_password_complexity(password): try: - store_api_key(SERVICE_NAME, USERNAME, api_key, password) + await store_api_key(SERVICE_NAME, USERNAME, api_key, password) logging.info("API key stored securely.") break except Exception as e: @@ -145,8 +129,8 @@ class APIKeyManager: _api_key = None @classmethod - def load(cls, service: str, username: str, password: str): - cls._api_key = retrieve_api_key(service, username, password) + async def load(cls, service: str, username: str, password: str): + cls._api_key = await retrieve_api_key(service, username, password) @classmethod def get(cls) -> str: diff --git a/utils/configmanager.py b/utils/configmanager.py index 56c8f3c..1a26999 100644 --- a/utils/configmanager.py +++ b/utils/configmanager.py @@ -5,6 +5,8 @@ import sys from pathlib import Path from typing import Callable, Optional, TypeVar +import aiofiles + T = TypeVar("T") logger = logging.getLogger(__name__) @@ -19,21 +21,20 @@ PROTECTED_KEYS = [ _protected_config = {} -def get_system_config_path() -> Path: - # Check inside bundled EXE directory first +async def get_system_config_path() -> Path: bundled_dir = Path(getattr(sys, '_MEIPASS', '')) bundled_path = bundled_dir / "system_config.json" if bundled_path.exists(): return bundled_path - - # Fallback to external location return Path(__file__).parent.parent / "system_config.json" -def load_protected_config() -> dict: +async def load_protected_config() -> dict: global _protected_config try: - with open(get_system_config_path(), "r") as f: - system_config = json.load(f) + config_path = await get_system_config_path() + async with aiofiles.open(config_path, "r") as f: + content = await f.read() + system_config = json.loads(content) except FileNotFoundError: logging.warning("āš ļø system_config.json not found. Using built-in defaults.") system_config = { @@ -46,10 +47,10 @@ def load_protected_config() -> dict: } } - _protected_config = {key: system_config[key] for key in PROTECTED_KEYS} + _protected_config = {key: system_config[key] for key in PROTECTED_KEYS if key in system_config} return _protected_config -def get_protected_value(key: str, cast_type: Callable[[str], T] = str, default: Optional[T] = None) -> Optional[T]: +async def get_protected_value(key: str, cast_type: Callable[[str], T] = str, default: Optional[T] = None) -> Optional[T]: value = _protected_config.get(key) if value is None: logging.warning(f"Protected config key '{key}' not found.") @@ -62,7 +63,7 @@ def get_protected_value(key: str, cast_type: Callable[[str], T] = str, default: logging.warning(f"Invalid value for protected key '{key}': {value}. Expected type {cast_type.__name__}.") return default -def get_protected_json(key: str, default: str = "{}") -> dict: +async def get_protected_json(key: str, default: str = "{}") -> dict: raw = _protected_config.get(key, default) if isinstance(raw, dict): return raw @@ -75,11 +76,8 @@ def get_protected_json(key: str, default: str = "{}") -> dict: except Exception as e: logging.error(f"Failed to parse protected JSON key '{key}': {e}") return json.loads(default) - - - -def load_env_json(key: str, default: str): +async def load_env_json(key: str, default: str): raw = os.getenv(key, default) try: return json.loads(raw) @@ -91,24 +89,13 @@ def load_env_json(key: str, default: str): logging.error(f"Failed to parse {key}: {e}") return json.loads(default) -def load_env(key: str, cast_type: Callable[[str], T] = str, default: Optional[T] = None) -> Optional[T]: - """ - Safely retrieves an environment variable and casts it to the desired type. - - Parameters: - key (str): The name of the environment variable. - cast_type (Callable[[str], T], optional): Function to cast the value. Defaults to str. - default (Optional[T], optional): Default value if the variable is not set or invalid. - - Returns: - Optional[T]: The casted value or the default. - """ +async def load_env(key: str, cast_type: Callable[[str], T] = str, default: Optional[T] = None) -> Optional[T]: value = os.getenv(key) if value is None: logger.warning(f"Environment variable '{key}' not set.") return default try: - value = value.strip("'\"") # Strip surrounding quotes + value = value.strip("'\"") return cast_type(value) except (ValueError, TypeError): logger.warning(f"Invalid value for env var '{key}': {value}. Expected type {cast_type.__name__}.") diff --git a/utils/menus.py b/utils/menus.py index 01295a6..53f8c82 100644 --- a/utils/menus.py +++ b/utils/menus.py @@ -20,7 +20,7 @@ import re import dotenv import pandas as pd -import services.policyhandler as policyh +import services.PolicyHandler as policyh from flows.otp import generate, otp_activities_by_agent, revoke from flows.prepPolicy import ( buildPathsandPublishers, @@ -32,8 +32,9 @@ from flows.prepPolicy import ( from flows.quietAgent import findQuietAgents from services.agenthandler import findAgents, moveAgentToRelatedPolicy, selectAgents from services.API import AirlockAPIWrapper +from services.TaskQueue import AsyncTaskQueue from utils.configmanager import load_env -from utils.selector import Selector +from utils.Selector import Selector from utils.utils import ( areYouSure, colorText, @@ -47,9 +48,9 @@ logger = logging.getLogger(__name__) dotenv.load_dotenv() -def menu_main(api: AirlockAPIWrapper): - working_dir = load_env("WORKING_DIR") - extras = load_env("EXTRAS") +async def menu_main(api: AirlockAPIWrapper, queue: AsyncTaskQueue): + working_dir = await load_env("WORKING_DIR") + extras = await load_env("EXTRAS") while True: displayIntro() # Add Settings, and give option to change working dir @@ -63,34 +64,34 @@ def menu_main(api: AirlockAPIWrapper): print(colorText("S. šŸ› ļø - Settings", "yellow")) print(colorText("Q. šŸ”š - Quit", "yellow")) - choice = get_sanitized_input("\nEnter Menu Item: ") + choice = await get_sanitized_input("\nEnter Menu Item: ") if choice == "1": print("This Feature is still in development") - get_sanitized_input("Press enter to continue") + await get_sanitized_input("Press enter to continue") elif choice == "2": - menu_otp(api) + await menu_otp(api) elif choice == "3": choices = ["audit", "enforcement"] print(colorText("Move devices to which state?:", "yellow")) - direction = Selector.select_string(choices, False, False) - devices = selectAgents(api) + direction = await Selector.select_string(choices, False, False) + devices = await selectAgents(api) print(colorText("Would you like to continue with these devices?","white")) for device in devices: print(device.hostname) - confirm = Selector.confirm() + confirm =await Selector.confirm() if direction and devices and confirm: for device in devices: - moveAgentToRelatedPolicy(api,device, direction[0]) + await moveAgentToRelatedPolicy(api,device, direction) elif choice == "4": - findAgents(api,False) + await findAgents(api,False) elif choice == "5": - findQuietAgents(api) + await findQuietAgents(api) elif choice == "6": - if extras == "POLICYPREP": menu_policymanagment(api) + if extras == "POLICYPREP": await menu_policymanagment(api, queue) elif choice.upper() == "F": - open_directory(working_dir) + await open_directory(working_dir) elif choice.upper() == "S": - menu_settings() + await menu_settings() elif choice.upper() == "Q": break else: @@ -98,7 +99,7 @@ def menu_main(api: AirlockAPIWrapper): -def menu_policy_enforce(api: AirlockAPIWrapper): +async def menu_policy_enforce(api: AirlockAPIWrapper, queue: AsyncTaskQueue): selected_policies = [] destination_policy = [] destination_allowlist = [] @@ -106,34 +107,35 @@ def menu_policy_enforce(api: AirlockAPIWrapper): processed_hashes = [] processed_publishers = [] tested = False - working_dir = load_env("WORKING_DIR") + working_dir = await load_env("WORKING_DIR") while True: printEnforceChecklist(selected_policies, destination_policy, destination_allowlist) - choice = get_sanitized_input("\nEnter your choice: ") + choice = await get_sanitized_input("\nEnter your choice: ") if choice == "1": - selected_policies = selectPolicies(api,True) + selected_policies = await selectPolicies(api,True) elif choice == "2": print(colorText("Please choose destination_name Policy for Path Exclusions", "white")) - destination_policy = selectPolicies(api, False) + destination_policy = await selectPolicies(api, False) print(colorText("Please choose Allowlist for Hashes", "white")) - destination_allowlist = selectAllowlists(api, destination_policy, False) + destination_allowlist = await selectAllowlists(api, destination_policy, False) # pyright: ignore[reportArgumentType] elif choice == "3": - sortHashes( + await sortHashes( api, + queue, selected_policies, type=[1, 2, 6, 7], ) elif choice == "4": if os.path.exists(f"{working_dir}\\Needs_Review\\Review_First\\approved_executions.csv"): - buildPathsandPublishers(False) + await buildPathsandPublishers(False) else: print("File not found. Please make sure it's saved correctly and try again.") @@ -141,7 +143,7 @@ def menu_policy_enforce(api: AirlockAPIWrapper): if os.path.exists(f"{working_dir}\\Approved\\hashes_to_add.csv") and os.path.exists( f"{working_dir}\\Approved\\primary_Paths.csv" ): - buildPreflights() + await buildPreflights() else: print("File not found. Please make sure it's saved correctly and try again.") @@ -209,7 +211,7 @@ def menu_policy_enforce(api: AirlockAPIWrapper): elif choice == "7": areYouSure() - confirmation = get_sanitized_input("Type 'I AGREE' to continue: ") + confirmation = await get_sanitized_input("Type 'I AGREE' to continue: ") if ( tested and destination_policy @@ -217,10 +219,10 @@ def menu_policy_enforce(api: AirlockAPIWrapper): and confirmation.strip() == "I AGREE" ): print(colorText("Proceeding with the code...", "yellow")) - api.hash_add_to_allowlist(destination_allowlist[0].applicationid, processed_hashes) - api.policy_add_path_exclusions(destination_policy[0].groupid, processed_paths) + await api.hash_add_to_allowlist(destination_allowlist[0].applicationid, processed_hashes) + await api.policy_add_path_exclusions(destination_policy[0].groupid, processed_paths) if processed_publishers: - api.policy_add_publishers(destination_policy[0].groupid, processed_publishers) + await api.policy_add_publishers(destination_policy[0].groupid, processed_publishers) else: logger.error("Confirmation block failed. Reasons:") if not tested: @@ -233,9 +235,9 @@ def menu_policy_enforce(api: AirlockAPIWrapper): logger.error(" - User did not confirm with 'I AGREE'. Received: '%s'", confirmation.strip()) elif choice.upper() == "F": - open_directory(working_dir) + await open_directory(working_dir) elif choice.upper() == "S": - menu_settings() + await menu_settings() elif choice.upper == "B": break @@ -244,7 +246,7 @@ def menu_policy_enforce(api: AirlockAPIWrapper): print(colorText("Invalid choice. Please try again.", "red")) -def menu_otp(api: AirlockAPIWrapper): +async def menu_otp(api: AirlockAPIWrapper): working_dir = load_env("WORKING_DIR") while True: @@ -256,23 +258,23 @@ def menu_otp(api: AirlockAPIWrapper): print(colorText("S. šŸ› ļø - Settings", "yellow")) print(colorText("B. šŸ”™ - Back", "yellow")) - choice = get_sanitized_input("Enter your choice: ") + choice = await get_sanitized_input("Enter your choice: ") if choice == "1": - otp_list = generate(api) + otp_list = await generate(api) print(colorText(otp_list,"green")) elif choice == "2": - otp_activities_by_agent(api) + await otp_activities_by_agent(api) elif choice == "3": - revoke(api) + await revoke(api) elif choice.upper() == "F": - open_directory(working_dir) + await open_directory(working_dir) elif choice.upper() == "S": - menu_settings() + await menu_settings() elif choice.upper() == "B": break -def menu_policymanagment(api: AirlockAPIWrapper): +async def menu_policymanagment(api: AirlockAPIWrapper, queue: AsyncTaskQueue): working_dir = load_env("WORKING_DIR") while True: print(colorText("1. šŸ”’ - Prepare Policy For Enforcement", "yellow")) @@ -280,31 +282,31 @@ def menu_policymanagment(api: AirlockAPIWrapper): print(colorText("F. šŸ“‚ - Open Working Directory", "yellow")) print(colorText("S. šŸ› ļø - Settings", "yellow")) print(colorText("B. šŸ”™ - Back", "yellow")) - choice = get_sanitized_input("\n Enter Menu Item: ") + choice = await get_sanitized_input("\n Enter Menu Item: ") if choice == "1": - menu_policy_enforce(api) + await menu_policy_enforce(api, queue) elif choice == "2": areYouSure() - confirmation = get_sanitized_input("Type 'I AGREE' to continue: ") + confirmation = await get_sanitized_input("Type 'I AGREE' to continue: ") if confirmation.strip() == "I AGREE": - policyh.updateAuditPoliciesFromEnforcementPolices(api) + await policyh.updateAuditPoliciesFromEnforcementPolices(api) elif choice.upper() == "F": - open_directory(working_dir) + await open_directory(working_dir) elif choice.upper() == "S": - menu_settings() + await menu_settings() elif choice.upper() == "B": break else: print(colorText("Invalid choice. Please try again.", "red")) -def menu_settings(): +async def menu_settings(): while True: print(colorText("\n--- šŸ› ļø Settings Submenu šŸ› ļø ---", "cyan")) print(colorText("This Feature is still in development", "cyan")) # print(colorText("2. Sub-option B","cyan")) print(colorText("B. šŸ”™ - Back", "yellow")) - choice = get_sanitized_input("Enter your choice: ") + choice = await get_sanitized_input("Enter your choice: ") if choice == "1": pass #TODO ADD CHANGE WORKDIR CODE diff --git a/utils/selector.py b/utils/selector.py index 87f1bfb..dd49e44 100644 --- a/utils/selector.py +++ b/utils/selector.py @@ -1,33 +1,16 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . - - import logging from typing import Any, Callable, List, Optional, Union -from utils.utils import get_sanitized_input +from utils.utils import get_sanitized_input # new async version logger = logging.getLogger(__name__) - class Selector: @staticmethod def _display_choices( items: List[Any], label_func: Callable[[Any], str], - num_columns: int = 4, + num_columns: int = 3, header: str = "Available Choices:" ) -> None: sorted_items = sorted(items, key=lambda item: label_func(item).lower()) @@ -43,7 +26,7 @@ class Selector: print(line) @staticmethod - def _select_from_list( + async def _select_from_list( items: List[Any], label_func: Callable[[Any], str], allow_multiple: bool = False, @@ -60,7 +43,7 @@ class Selector: if allow_multiple: while True: - choice = get_sanitized_input("Select an item by number (or Q to finish): ").strip().lower() + choice = (await get_sanitized_input("Select an item by number (or Q to finish): ")).strip().lower() if choice == "q": break try: @@ -80,7 +63,7 @@ class Selector: return selected if selected else None else: try: - choice = int(get_sanitized_input("Select one item by number: ")) + choice = int(await get_sanitized_input("Select one item by number: ")) if 1 <= choice <= len(sorted_items): selected_item = sorted_items[choice - 1] logger.info(f"Selected: {label_func(selected_item)}") @@ -92,12 +75,8 @@ class Selector: return None @staticmethod - def select_objects( - objects: List[Any], - allow_multiple: bool = False, - prompt_each: bool = False - ) -> Union[Optional[Any], List[Any]]: - return Selector._select_from_list( + async def select_objects(objects: List[Any], allow_multiple: bool = False, prompt_each: bool = False) -> Union[Optional[Any], List[Any]]: + return await Selector._select_from_list( objects, label_func=lambda obj: getattr(obj, "name", str(obj)), allow_multiple=allow_multiple, @@ -106,12 +85,8 @@ class Selector: ) @staticmethod - def select_string( - options: List[str], - allow_multiple: bool = False, - prompt_each: bool = False - ) -> Union[Optional[str], List[str]]: - return Selector._select_from_list( + async def select_string(options: List[str], allow_multiple: bool = False, prompt_each: bool = False) -> Union[Optional[str], List[str]]: + return await Selector._select_from_list( options, label_func=str, allow_multiple=allow_multiple, @@ -120,12 +95,8 @@ class Selector: ) @staticmethod - def select_int( - options: List[int], - allow_multiple: bool = False, - prompt_each: bool = False - ) -> Union[Optional[int], List[int]]: - return Selector._select_from_list( + async def select_int(options: List[int], allow_multiple: bool = False, prompt_each: bool = False) -> Union[Optional[int], List[int]]: + return await Selector._select_from_list( options, label_func=lambda x: str(x), allow_multiple=allow_multiple, @@ -134,16 +105,11 @@ class Selector: ) @staticmethod - def select_value( - prompt: str, - value_type: type = int, - valid_range: Optional[tuple] = None, - allow_quit: bool = False - ) -> Optional[Any]: + async def select_value(prompt: str, value_type: type = int, valid_range: Optional[tuple] = None, allow_quit: bool = False) -> Optional[Any]: while True: - user_input = get_sanitized_input(prompt).strip().lower() + user_input = (await get_sanitized_input(prompt)).strip().lower() if allow_quit and user_input == "q": - logger.info("User opted to quit value selection.") + logger.debug("User opted to quit value selection.") return None try: value = value_type(user_input) @@ -152,20 +118,20 @@ class Selector: if not (min_val <= value <= max_val): logger.warning(f"Value out of range ({min_val}–{max_val}).") continue - logger.info(f"User selected value: {value}") + logger.debug(f"User selected value: {value}") return value except ValueError: logger.warning(f"Invalid input. Expected a {value_type.__name__}.") @staticmethod - def confirm(prompt: str = "Are you sure? (Y/N): ") -> bool: + async def confirm(prompt: str = "Are you sure? (Y/N): ") -> bool: while True: - response = get_sanitized_input(prompt).strip().lower() + response = (await get_sanitized_input(prompt)).strip().lower() if response in ["y", "yes"]: - logger.info("User confirmed action.") + logger.debug("User confirmed action.") return True elif response in ["n", "no"]: - logger.info("User declined action.") + logger.debug("User declined action.") return False else: - logger.warning("Invalid confirmation input. Expected 'Y' or 'N'.") \ No newline at end of file + logger.warning("Invalid confirmation input. Expected 'Y' or 'N'.") diff --git a/utils/setup.py b/utils/setup.py index d2d70ef..5d3e153 100644 --- a/utils/setup.py +++ b/utils/setup.py @@ -1,17 +1,3 @@ -# Copyright (C) 2025 James Brotosky, Brandon Wickline -# -# This program is free software: you can redistribute it and/or modify -# it under the terms of the GNU Affero 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 Affero General Public License for more details. -# -# You should have received a copy of the GNU Affero General Public License -# along with this program. If not, see . import json import logging @@ -41,7 +27,11 @@ def configure_logging(log_dir: Path, log_level: str = "DEBUG"): logger = logging.getLogger() logger.setLevel(getattr(logging, log_level.upper(), logging.DEBUG)) - # šŸ”§ Clear existing handlers + + httpx_logger = logging.getLogger("httpx") + httpx_logger.setLevel(getattr(logging, log_level.upper(), logging.DEBUG)) + + for handler in logger.handlers[:]: logger.removeHandler(handler) @@ -104,14 +94,14 @@ def load_user_config(config_dir: Path) -> dict: def write_config_to_env(config: dict, env_path: Path): for key, value in config.items(): if key in PROTECTED_KEYS: - continue # Skip protected keys + continue try: serialized = json.dumps(value) if isinstance(value, (list, dict)) else str(value) set_key(env_path, key, serialized) except Exception as e: logging.warning(f"Failed to write {key} to .env: {e}") -def setup() -> Path: +async def setup() -> Path: base_dir = get_base_directory() dirs = { 'config': base_dir / 'config', @@ -129,6 +119,7 @@ def setup() -> Path: env_path = base_dir / ".env" if not env_path.exists(): env_path.touch() + load_dotenv(dotenv_path=env_path, override=True) working_dir = Path(os.getenv("WORKING_DIR") or (base_dir / "data")) @@ -155,14 +146,10 @@ def setup() -> Path: user_config = load_user_config(dirs['config']) merged_config = {**system_config, **user_config} - - protected_config = load_protected_config() + protected_config = await load_protected_config() merged_config.update(protected_config) - # āœ… URL resolution order: system_config → .env → user prompt - url = system_config.get("URL") - if not url: - url = os.getenv("URL") + url = system_config.get("URL") or os.getenv("URL") if not url: url = input("🌐 Enter the service URL (e.g., https://example.com/api): ").strip() merged_config["URL"] = url @@ -171,5 +158,4 @@ def setup() -> Path: logging.debug(f"Service URL set to: {url}") write_config_to_env(merged_config, env_path) - return working_dir \ No newline at end of file diff --git a/utils/utils.py b/utils/utils.py index 6803411..4e87604 100644 --- a/utils/utils.py +++ b/utils/utils.py @@ -14,6 +14,7 @@ # along with this program. If not, see . +import asyncio import logging import os import platform @@ -99,15 +100,15 @@ def choose_file(initial_directory=None, required_substring=None): - -def get_sanitized_input(prompt: str) -> str: +async def get_sanitized_input(prompt: str) -> str: while True: - user_input = input(prompt) + user_input = await asyncio.to_thread(input, prompt) if user_input.strip() == "": - return user_input # Allow blank lines - if re.match(r'^[a-zA-Z0-9_\- .]+$', user_input.strip()): + return user_input + if re.match(r'^[a-zA-Z0-9_ .-]+$', user_input.strip()): return user_input else: + logger.debug("User entered invalid input") print("Invalid input. Only letters, numbers, underscores, spaces, hyphens, and periods are allowed.") @@ -695,14 +696,11 @@ def formatHTML(df, output_html_path=None, overwrite=True): -def open_directory(path): +async def open_directory(path): system = platform.system() - if system == "Windows": - os.startfile(path) + await asyncio.to_thread(os.startfile, path) elif system == "Linux": - subprocess.run(["xdg-open", path]) + await asyncio.to_thread(subprocess.run, ["xdg-open", path]) else: - raise OSError(f"Unsupported operating system: {system}") - - + raise OSError(f"Unsupported operating system: {system}") \ No newline at end of file