feat: add version updater and statistics enhancements (fixes #29)
- Implemented version checking system with update notifications - Integrated Git for fetching and downloading the latest version - Added statistics updates - Removed unused code across the project - Condensed project structure - Updated README - Cleaned up UI
This commit is contained in:
@@ -1,388 +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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime, timedelta
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from flows.prepPolicy import selectPolicies
|
||||
from models.agent import Agent
|
||||
from models.policy import Policy
|
||||
from services.API import AirlockAPIWrapper
|
||||
from utils.configmanager import get_system_json, load_env
|
||||
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(
|
||||
prompt="Enter how many days of history to pull (1–365): ",
|
||||
value_type=int,
|
||||
valid_range=(1, 365),
|
||||
)
|
||||
|
||||
if not agents or not history_days:
|
||||
print(colorText("No agents selected or invalid history range.", "red"))
|
||||
return
|
||||
|
||||
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:
|
||||
try:
|
||||
exechistory = api.history_execution(today, historical_date, agent.hostname)
|
||||
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")
|
||||
)
|
||||
|
||||
if outputjson:
|
||||
print(json.dumps(all_history, indent=2))
|
||||
|
||||
|
||||
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()]
|
||||
|
||||
for agent in agents:
|
||||
agent.enrich_with_policies(policies)
|
||||
|
||||
return agents
|
||||
|
||||
|
||||
def findAgents(api, return_dataframe):
|
||||
agents = selectAgents(api)
|
||||
working_dir = 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)
|
||||
|
||||
if 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()
|
||||
)
|
||||
if user_input == "y":
|
||||
timestamp = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
filename = f"agentsearch_{timestamp}.csv"
|
||||
file_path = os.path.join(str(working_dir), filename)
|
||||
|
||||
agent_df.to_csv(file_path, index=False)
|
||||
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 collect_device_names() -> List[str]:
|
||||
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 (Three times if you have a single device).\n",
|
||||
"cyan",
|
||||
)
|
||||
)
|
||||
print(colorText("Example:", "cyan"))
|
||||
print(colorText("H00000\nUTN00000\ni-hSuperSecretServer\nu-hVenderBroke\n", "cyan"))
|
||||
print(colorText("Paste or type your device names below:", "white"))
|
||||
|
||||
device_input_lines = []
|
||||
empty_line_count = 0
|
||||
valid_line_pattern = re.compile(r"^[a-zA-Z0-9_\- ]+$")
|
||||
|
||||
while True:
|
||||
line = get_sanitized_input("")
|
||||
stripped_line = line.strip()
|
||||
|
||||
if stripped_line == "":
|
||||
empty_line_count += 1
|
||||
if empty_line_count == 2:
|
||||
break
|
||||
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",
|
||||
)
|
||||
)
|
||||
|
||||
return [name for name in device_input_lines if name]
|
||||
|
||||
|
||||
def choose_match_type() -> bool:
|
||||
print(colorText("Use exact match? (Y for exact, N for fuzzy):", "white"))
|
||||
return get_sanitized_input("").strip().lower() in ["y", "yes"]
|
||||
|
||||
|
||||
def match_agents(
|
||||
device_names: List[str], agents: List["Agent"], use_exact: bool
|
||||
) -> List["Agent"]:
|
||||
if use_exact:
|
||||
return [
|
||||
agent
|
||||
for agent in agents
|
||||
if agent.hostname.lower() in [name.lower() for name in device_names]
|
||||
]
|
||||
else:
|
||||
pattern = "|".join(map(re.escape, device_names))
|
||||
regex = re.compile(pattern, re.IGNORECASE)
|
||||
return [agent for agent in agents if regex.search(agent.hostname)]
|
||||
|
||||
|
||||
def show_unmatched(
|
||||
device_names: List[str], matched_agents: List["Agent"], use_exact: bool
|
||||
):
|
||||
if use_exact:
|
||||
unmatched = [
|
||||
name
|
||||
for name in device_names
|
||||
if not any(
|
||||
agent.hostname.lower() == name.lower() for agent in matched_agents
|
||||
)
|
||||
]
|
||||
else:
|
||||
unmatched = [
|
||||
name
|
||||
for name in device_names
|
||||
if not any(
|
||||
re.search(re.escape(name), agent.hostname, re.IGNORECASE)
|
||||
for agent in matched_agents
|
||||
)
|
||||
]
|
||||
|
||||
if unmatched:
|
||||
logger.debug(f"âš ï¸ No matches for: {', '.join(unmatched)}")
|
||||
print(colorText(f"âš ï¸ No matches for: {', '.join(unmatched)}", "yellow"))
|
||||
|
||||
|
||||
def enrich_agents(agents: List["Agent"], policies: List["Policy"]):
|
||||
for agent in agents:
|
||||
agent.enrich_with_policies(policies)
|
||||
|
||||
|
||||
def selectAgents(api: "AirlockAPIWrapper") -> List["Agent"]:
|
||||
device_names = collect_device_names()
|
||||
if not device_names:
|
||||
logger.debug("No device names entered")
|
||||
print(colorText("âš ï¸ No device names entered.", "red"))
|
||||
return []
|
||||
|
||||
use_exact = choose_match_type()
|
||||
|
||||
policies = [Policy(**row.to_dict()) for _, row in api.policy_find_all().iterrows()]
|
||||
agents = [Agent(**row.to_dict()) for _, row in api.agent_find_all().iterrows()]
|
||||
matched_agents = match_agents(device_names, agents, use_exact)
|
||||
matched_agents.sort(key=lambda agent: agent.hostname.lower())
|
||||
|
||||
show_unmatched(device_names, matched_agents, use_exact)
|
||||
|
||||
if not matched_agents:
|
||||
logger.debug("⌠No matching devices found.")
|
||||
print(colorText("⌠No matching devices found.", "red"))
|
||||
return []
|
||||
|
||||
print(colorText(f"✓ Found {len(matched_agents)} matching device(s).", "green"))
|
||||
logger.info("Matched agent hostnames:")
|
||||
rows = (len(matched_agents) + 2) // 3 # 3 columns
|
||||
for row in range(rows):
|
||||
line = ""
|
||||
for col in range(3):
|
||||
idx = row + col * rows
|
||||
if idx < len(matched_agents):
|
||||
line += f"{matched_agents[idx].hostname:<30}"
|
||||
logger.info(line)
|
||||
|
||||
matched_agents = Selector.select_with_mode(
|
||||
matched_agents,
|
||||
label_func=lambda agent: agent.hostname,
|
||||
header="Matched Devices:",
|
||||
)
|
||||
|
||||
if not matched_agents:
|
||||
logger.debug("⌠No matching devices remain after refinement.")
|
||||
print(colorText("⌠No matching devices remain after refinement.", "red"))
|
||||
return []
|
||||
|
||||
enrich_agents(matched_agents, 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_system_json("POLICY_MAP_ENF_AUD", "{}")
|
||||
|
||||
if mode == "audit":
|
||||
if agent.groupid in policy_relationship_map:
|
||||
target_policy = policy_relationship_map[agent.groupid]
|
||||
elif agent.groupid in policy_relationship_map.values():
|
||||
logger.debug(
|
||||
f"Agent {agent.hostname} is already in an audit group. No action needed."
|
||||
)
|
||||
print(
|
||||
f"Agent {agent.hostname} is already in an audit group. No action needed."
|
||||
)
|
||||
return
|
||||
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:
|
||||
target_policy = inverse_map[agent.groupid]
|
||||
elif agent.groupid in inverse_map.values():
|
||||
logger.info(
|
||||
f"Agent {agent.hostname} is already in an enforcement group. No action needed."
|
||||
)
|
||||
return
|
||||
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
|
||||
|
||||
result = api.agent_move(agent.agentid, target_policy)
|
||||
return result
|
||||
|
||||
|
||||
def toggleEnforcement(api: AirlockAPIWrapper):
|
||||
choices = ["Audit", "Enforcement", "Exit"]
|
||||
print(colorText("Move devices to which state?:", "yellow"))
|
||||
direction = Selector.select_string(choices, False, False)
|
||||
if direction == "Exit":
|
||||
pass
|
||||
else:
|
||||
devices = selectAgents(api)
|
||||
for device in devices:
|
||||
print(device.hostname)
|
||||
confirm = Selector.confirm(
|
||||
"Would you like to continue with these devices? Y/N: "
|
||||
)
|
||||
if direction and devices and confirm:
|
||||
for device in devices:
|
||||
result = moveAgentToRelatedPolicy(api, device, str(direction).lower())
|
||||
logger.info(f"{device.hostname}: result: {result}")
|
||||
get_sanitized_input("Press enter to continue")
|
||||
|
||||
|
||||
def moveAgents(api: AirlockAPIWrapper):
|
||||
devices = selectAgents(api)
|
||||
for device in devices:
|
||||
print(device.hostname)
|
||||
confirm_devices = Selector.confirm(
|
||||
"Would you like to continue with these devices? Y/N: "
|
||||
)
|
||||
if devices and confirm_devices:
|
||||
policies = selectPolicies(api, False)
|
||||
confirm_move = Selector.confirm(
|
||||
f"Would you like to move these devices to {policies[0].name}?"
|
||||
)
|
||||
if confirm_move:
|
||||
for device in devices:
|
||||
result = api.agent_move(device.agentid, policies[0].groupid)
|
||||
logger.info(f"{device.hostname}: result: {result}")
|
||||
else:
|
||||
logger.info("Exiting without change")
|
||||
get_sanitized_input("Press enter to continue")
|
||||
@@ -1,234 +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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
|
||||
import datetime
|
||||
import gc
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
|
||||
from bson import ObjectId
|
||||
import pandas as pd
|
||||
import tqdm
|
||||
|
||||
from models.policy import Policy
|
||||
from services.API import AirlockAPIWrapper
|
||||
from utils.setup import get_base_directory
|
||||
from utils.utils import colorText
|
||||
|
||||
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.")
|
||||
|
||||
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]
|
||||
)
|
||||
|
||||
# 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.total = len(histories)
|
||||
|
||||
if not histories:
|
||||
break
|
||||
|
||||
for index, history_item in enumerate(histories):
|
||||
if (
|
||||
"checkpoint" not in history_item
|
||||
or "datetime" not in history_item
|
||||
):
|
||||
continue # Skip malformed entries
|
||||
|
||||
# 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", ""
|
||||
), # 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)
|
||||
|
||||
filebar.update(1)
|
||||
filebar.refresh()
|
||||
|
||||
# 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
|
||||
|
||||
deduplicated = list(seen.values())
|
||||
with open(file_path, "w") as file:
|
||||
json.dump(
|
||||
{
|
||||
"error": "Success",
|
||||
"response": {"exechistories": deduplicated},
|
||||
},
|
||||
file,
|
||||
)
|
||||
|
||||
json_output["response"]["exechistories"].clear()
|
||||
|
||||
# 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
|
||||
|
||||
# Final output
|
||||
with open(file_path, "r") as file:
|
||||
final_output = json.load(file)
|
||||
os.remove(file_path)
|
||||
|
||||
return json.dumps(final_output) if outputjson else None
|
||||
|
||||
|
||||
def getPolicyInfo(api: AirlockAPIWrapper, policy, type, days):
|
||||
import airlock_libs
|
||||
|
||||
executionhist_policy = pd.DataFrame()
|
||||
exehist = airlock_libs.pull_policy_exec_histories(api, policy.name, str(type), days)
|
||||
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",
|
||||
]
|
||||
]
|
||||
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",
|
||||
)
|
||||
)
|
||||
del data
|
||||
del exehist
|
||||
gc.collect()
|
||||
return executionhist_policy
|
||||
|
||||
|
||||
def skipback(days):
|
||||
"""
|
||||
Generate a MongoDB ObjectId for a given number of days ago from today.
|
||||
"""
|
||||
adjusted_days = days
|
||||
date_days_ago = datetime.datetime.now(datetime.UTC) - datetime.timedelta(
|
||||
days=adjusted_days
|
||||
)
|
||||
timestamp = int(date_days_ago.timestamp())
|
||||
hex_timestamp = format(timestamp, "08x")
|
||||
objectid_hex = hex_timestamp + "0000000000000000"
|
||||
return ObjectId(objectid_hex)
|
||||
@@ -1,185 +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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
import base64
|
||||
from getpass import getpass
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import re
|
||||
import sys
|
||||
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
|
||||
import keyring
|
||||
|
||||
# Constants
|
||||
KDF_ITERATIONS = 200_000
|
||||
SALT_SIZE = 16 # 128-bit Salt
|
||||
NONCE_SIZE = 12 # AES-GCM
|
||||
KEY_SIZE = 32 # AES-256
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _derive_key(password: bytes, salt: bytes) -> bytes:
|
||||
kdf = PBKDF2HMAC(
|
||||
algorithm=hashes.SHA256(),
|
||||
length=KEY_SIZE,
|
||||
salt=salt,
|
||||
iterations=KDF_ITERATIONS,
|
||||
)
|
||||
return kdf.derive(password)
|
||||
|
||||
|
||||
def configure_keyring_backend():
|
||||
system = platform.system()
|
||||
if system == "Windows":
|
||||
import keyring.backends.Windows
|
||||
|
||||
keyring.set_keyring(keyring.backends.Windows.WinVaultKeyring())
|
||||
elif system == "Linux":
|
||||
import keyring.backends.kwallet
|
||||
|
||||
keyring.set_keyring(keyring.backends.kwallet.DBusKeyring())
|
||||
else:
|
||||
raise EnvironmentError(f"Unsupported OS: {system}")
|
||||
|
||||
|
||||
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)
|
||||
aesgcm = AESGCM(key)
|
||||
nonce = os.urandom(NONCE_SIZE)
|
||||
ct = aesgcm.encrypt(nonce, api_key.encode(), associated_data=None)
|
||||
blob = salt + nonce + ct
|
||||
b64 = base64.b64encode(blob).decode()
|
||||
keyring.set_password(service, username, b64)
|
||||
|
||||
logger.debug(
|
||||
f"API key for service '{service}' and user '{username}' stored successfully."
|
||||
)
|
||||
|
||||
print("\n✅ API key stored securely.")
|
||||
print("The program will now exit. Press Enter to continue...")
|
||||
|
||||
try:
|
||||
_ = input()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
_ = None
|
||||
sys.exit(0)
|
||||
|
||||
|
||||
def retrieve_api_key(service: str, username: str, password: str) -> str:
|
||||
configure_keyring_backend()
|
||||
b64 = keyring.get_password(service, username)
|
||||
if b64 is None:
|
||||
raise ValueError("No stored secret for this service/username.")
|
||||
blob = base64.b64decode(b64)
|
||||
salt = blob[:SALT_SIZE]
|
||||
nonce = blob[SALT_SIZE : SALT_SIZE + NONCE_SIZE]
|
||||
ct = blob[SALT_SIZE + NONCE_SIZE :]
|
||||
key = _derive_key(password.encode(), salt)
|
||||
aesgcm = AESGCM(key)
|
||||
pt = aesgcm.decrypt(nonce, ct, associated_data=None)
|
||||
return pt.decode()
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
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):
|
||||
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)
|
||||
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.")
|
||||
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()
|
||||
print(
|
||||
"Please exit and relaunch program after saving your credential to avoid errors"
|
||||
)
|
||||
|
||||
while True:
|
||||
password = getpass("Create a password to encrypt your API key: ")
|
||||
confirm_password = getpass("Confirm your password: ")
|
||||
|
||||
if password != confirm_password:
|
||||
logging.warning("Passwords do not match. Try again.")
|
||||
continue
|
||||
|
||||
if check_password_complexity(password):
|
||||
try:
|
||||
store_api_key(SERVICE_NAME, USERNAME, api_key, password)
|
||||
logging.info("API key stored securely.")
|
||||
break
|
||||
except Exception as e:
|
||||
logging.error(f"Failed to store API key: {e}")
|
||||
break
|
||||
else:
|
||||
logging.warning(
|
||||
"Password does not meet complexity requirements. Try again."
|
||||
)
|
||||
|
||||
|
||||
class APIKeyManager:
|
||||
_api_key = None
|
||||
|
||||
@classmethod
|
||||
def load(cls, service: str, username: str, password: str):
|
||||
cls._api_key = retrieve_api_key(service, username, password)
|
||||
|
||||
@classmethod
|
||||
def get(cls) -> str:
|
||||
if cls._api_key is None:
|
||||
raise ValueError("API key not loaded. Call APIKeyManager.load() first.")
|
||||
return cls._api_key
|
||||
Reference in New Issue
Block a user