658 lines
24 KiB
Python
658 lines
24 KiB
Python
from __future__ import annotations
|
||
|
||
import logging
|
||
from typing import Optional
|
||
|
||
import pandas as pd
|
||
from textual.app import ComposeResult
|
||
from textual.binding import Binding
|
||
from textual.containers import Horizontal, Vertical
|
||
from textual.screen import Screen
|
||
from textual.widgets import (
|
||
Button,
|
||
DataTable,
|
||
Footer,
|
||
Header,
|
||
Static,
|
||
TextArea,
|
||
)
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
class AllowlistSelectionWidget(Static):
|
||
"""
|
||
Widget for selecting an allowlist and adding hashes to it.
|
||
Can be reused in different workflows.
|
||
"""
|
||
|
||
DEFAULT_CSS = """
|
||
AllowlistSelectionWidget {
|
||
height: 1fr;
|
||
layout: vertical;
|
||
}
|
||
#allowlist_main {
|
||
height: 1fr;
|
||
width: 100%;
|
||
}
|
||
#left_panel {
|
||
width: 50%;
|
||
padding: 1;
|
||
border: solid $primary;
|
||
}
|
||
#right_panel {
|
||
width: 50%;
|
||
padding: 1;
|
||
border: solid $primary;
|
||
}
|
||
#allowlist_table {
|
||
height: 70%;
|
||
margin: 1 0;
|
||
}
|
||
#allowlist_table > .datatable--header {
|
||
text-style: bold;
|
||
background: $boost;
|
||
}
|
||
#allowlist_table Row {
|
||
height: 1;
|
||
}
|
||
#preview_area {
|
||
height: 60%;
|
||
margin: 1 0;
|
||
}
|
||
#action_buttons {
|
||
height: auto;
|
||
min-height: 3;
|
||
padding: 1;
|
||
content-align: center middle;
|
||
}
|
||
.panel-title {
|
||
text-style: bold;
|
||
margin: 0 0 1 0;
|
||
}
|
||
.info-text {
|
||
margin: 1 0;
|
||
}
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
selected_data: pd.DataFrame,
|
||
api=None,
|
||
hostname: Optional[str] = None,
|
||
otpid: Optional[str] = None,
|
||
hash_column: str = "sha256", # Default hash column name
|
||
):
|
||
"""
|
||
Initialize the allowlist selection widget.
|
||
|
||
Args:
|
||
selected_data: DataFrame containing the selected activities
|
||
api: API instance for making allowlist calls
|
||
hostname: Optional hostname for context
|
||
otpid: Optional OTP ID for context
|
||
hash_column: Name of the column containing hashes (default: "sha256")
|
||
"""
|
||
super().__init__()
|
||
self.selected_data = selected_data
|
||
self.api = api
|
||
self.hostname = hostname
|
||
self.otpid = otpid
|
||
self.hash_column = hash_column
|
||
self.allowlists = []
|
||
self.selected_allowlist = None
|
||
self.hashes_to_add = []
|
||
|
||
def compose(self) -> ComposeResult:
|
||
with Horizontal(id="allowlist_main"):
|
||
# Left panel - Allowlist selection
|
||
with Vertical(id="left_panel"):
|
||
yield Static("Select Allowlist", classes="panel-title")
|
||
yield Static(
|
||
f"Choose an allowlist to add {len(self.selected_data)} selected items",
|
||
classes="info-text",
|
||
)
|
||
|
||
# Allowlist table
|
||
self.allowlist_table = DataTable(id="allowlist_table")
|
||
self.allowlist_table.cursor_type = "row"
|
||
yield self.allowlist_table
|
||
|
||
# Refresh button
|
||
self.refresh_btn = Button(
|
||
"🔄 Refresh Allowlists", id="refresh_allowlists_btn"
|
||
)
|
||
yield self.refresh_btn
|
||
|
||
# Right panel - Preview and actions
|
||
with Vertical(id="right_panel"):
|
||
yield Static("Preview", classes="panel-title")
|
||
|
||
# Context information
|
||
context_text = []
|
||
if self.hostname:
|
||
context_text.append(f"Host: {self.hostname}")
|
||
if self.otpid:
|
||
context_text.append(f"OTP: {self.otpid}")
|
||
context_text.append(f"Selected Activities: {len(self.selected_data)}")
|
||
|
||
yield Static(" | ".join(context_text), classes="info-text")
|
||
|
||
# Preview text area
|
||
self.preview_area = TextArea(
|
||
id="preview_area", read_only=True, language="markdown"
|
||
)
|
||
yield self.preview_area
|
||
|
||
# Hash statistics
|
||
self.stats_label = Static("", id="stats_label", classes="info-text")
|
||
yield self.stats_label
|
||
|
||
# Action buttons at bottom
|
||
with Horizontal(id="action_buttons"):
|
||
self.back_btn = Button("⬅ Back", id="back_btn")
|
||
self.add_btn = Button("➕ Add to Allowlist", id="add_to_allowlist_btn")
|
||
|
||
self.back_btn.styles.width = "50%"
|
||
self.add_btn.styles.width = "50%"
|
||
self.add_btn.disabled = True # Disabled until allowlist selected
|
||
|
||
yield self.back_btn
|
||
yield self.add_btn
|
||
|
||
async def on_mount(self) -> None:
|
||
"""Load allowlists when widget mounts."""
|
||
await self.load_allowlists()
|
||
await self.extract_and_preview_hashes()
|
||
|
||
async def load_allowlists(self) -> None:
|
||
"""Load available allowlists from API, grouped by policy association."""
|
||
if not self.api:
|
||
logger.error("No API available")
|
||
self.allowlist_table.add_column("Error")
|
||
self.allowlist_table.add_row("No API available")
|
||
return
|
||
|
||
try:
|
||
# First, try to get the host's policy if hostname is provided
|
||
host_policy_allowlists = []
|
||
host_policy_ids = set()
|
||
policy_name = "Unknown Policy" # Default value
|
||
group_id = None
|
||
|
||
if self.hostname:
|
||
try:
|
||
# Get agent info to find its policy
|
||
agents_df = self.api.agent_find_by_hostname(self.hostname)
|
||
if not agents_df.empty:
|
||
# Get the policy group ID for this host
|
||
group_id = agents_df.iloc[0].get("groupid")
|
||
|
||
# Look up the policy name from app's cached policies
|
||
if (
|
||
group_id
|
||
and hasattr(self.app, "policies")
|
||
and self.app.policies
|
||
):
|
||
for policy in self.app.policies:
|
||
if policy.groupid == group_id:
|
||
policy_name = policy.name
|
||
logger.info(
|
||
f"Found policy name: '{policy_name}' for group_id: {group_id}"
|
||
)
|
||
break
|
||
|
||
logger.info(
|
||
f"Found host '{self.hostname}' in policy '{policy_name}' (group_id: {group_id})"
|
||
)
|
||
|
||
if group_id:
|
||
# Get allowlists for this policy
|
||
policy_allowlists_df = self.api.policy_list_allowlists(
|
||
group_id
|
||
)
|
||
if not policy_allowlists_df.empty:
|
||
host_policy_allowlists = policy_allowlists_df.to_dict(
|
||
orient="records"
|
||
)
|
||
host_policy_ids = {
|
||
al.get("applicationid")
|
||
for al in host_policy_allowlists
|
||
}
|
||
logger.info(
|
||
f"Found {len(host_policy_allowlists)} allowlists for host's policy"
|
||
)
|
||
except Exception as e:
|
||
logger.warning(f"Could not get host's policy allowlists: {e}")
|
||
|
||
# If we still don't have a policy name, try to get it from the first allowlist or use a default
|
||
if not policy_name:
|
||
# Get all policies and try to find which one has allowlists
|
||
try:
|
||
all_policies_df = self.api.policy_find_all()
|
||
if not all_policies_df.empty:
|
||
# If we have a group_id from somewhere, use it
|
||
if group_id:
|
||
policy_row = all_policies_df[
|
||
all_policies_df["groupid"] == group_id
|
||
]
|
||
if not policy_row.empty:
|
||
policy_name = policy_row.iloc[0].get(
|
||
"groupname", "Unknown Policy"
|
||
)
|
||
else:
|
||
# Use the first policy as fallback
|
||
policy_name = all_policies_df.iloc[0].get(
|
||
"groupname", "Default Policy"
|
||
)
|
||
logger.info(f"Using first available policy: {policy_name}")
|
||
else:
|
||
policy_name = "Unknown Policy"
|
||
except Exception as e:
|
||
logger.warning(f"Could not fetch policies: {e}")
|
||
policy_name = "Unknown Policy"
|
||
|
||
# Get all allowlists
|
||
all_allowlists_df = self.api.allowlist_find_all()
|
||
|
||
if all_allowlists_df.empty:
|
||
self.allowlist_table.add_column("No Allowlists")
|
||
self.allowlist_table.add_row("No allowlists found")
|
||
return
|
||
|
||
all_allowlists = all_allowlists_df.to_dict(orient="records")
|
||
|
||
# Separate into two groups: policy-associated and others
|
||
other_allowlists = [
|
||
al
|
||
for al in all_allowlists
|
||
if al.get("applicationid") not in host_policy_ids
|
||
]
|
||
|
||
# Sort each group alphabetically by name
|
||
host_policy_allowlists.sort(key=lambda x: x.get("name", "").lower())
|
||
other_allowlists.sort(key=lambda x: x.get("name", "").lower())
|
||
|
||
# Combine lists with policy-associated first
|
||
self.allowlists = host_policy_allowlists + other_allowlists
|
||
|
||
# Setup table columns
|
||
self.allowlist_table.clear()
|
||
self.allowlist_table.add_columns("Name", "Application ID", "Type")
|
||
|
||
# Track which rows are headers vs actual allowlists
|
||
self._row_to_allowlist_map = {}
|
||
current_row = 0
|
||
|
||
# Add policy-associated allowlists if any
|
||
if host_policy_allowlists:
|
||
# Add section header
|
||
header_text = f"=== Policy: {policy_name or 'Host Policy'} ==="
|
||
self.allowlist_table.add_row(header_text, "", "", key="header_policy")
|
||
current_row += 1
|
||
|
||
# Add policy allowlists
|
||
for idx, allowlist in enumerate(host_policy_allowlists):
|
||
name = allowlist.get("name", "Unknown")
|
||
app_id = allowlist.get("applicationid", "Unknown")
|
||
|
||
self.allowlist_table.add_row(
|
||
f" {name}", # Indent to show grouping
|
||
app_id,
|
||
"Policy",
|
||
key=f"policy_{idx}",
|
||
)
|
||
self._row_to_allowlist_map[current_row] = idx
|
||
current_row += 1
|
||
|
||
# Add other allowlists
|
||
if other_allowlists:
|
||
# Add section header
|
||
if host_policy_allowlists:
|
||
# Add spacer if we have policy allowlists above
|
||
self.allowlist_table.add_row("", "", "", key="spacer")
|
||
current_row += 1
|
||
|
||
self.allowlist_table.add_row(
|
||
"=== Other Available Allowlists ===", "", "", key="header_other"
|
||
)
|
||
current_row += 1
|
||
|
||
# Add other allowlists
|
||
for idx, allowlist in enumerate(other_allowlists):
|
||
name = allowlist.get("name", "Unknown")
|
||
app_id = allowlist.get("applicationid", "Unknown")
|
||
|
||
self.allowlist_table.add_row(
|
||
f" {name}", # Indent to show grouping
|
||
app_id,
|
||
"General",
|
||
key=f"other_{idx}",
|
||
)
|
||
# Map to the correct index in the combined list
|
||
actual_idx = len(host_policy_allowlists) + idx
|
||
self._row_to_allowlist_map[current_row] = actual_idx
|
||
current_row += 1
|
||
|
||
# Log summary
|
||
logger.info(
|
||
f"Loaded {len(self.allowlists)} total allowlists: "
|
||
f"{len(host_policy_allowlists)} policy-associated, "
|
||
f"{len(other_allowlists)} others"
|
||
)
|
||
|
||
# Update stats label if no allowlists in policy
|
||
if self.hostname and not host_policy_allowlists:
|
||
self.stats_label.update(
|
||
f"Note: No allowlists found for {self.hostname}'s policy | "
|
||
+ self.stats_label.content.plain
|
||
)
|
||
|
||
except Exception as exc:
|
||
logger.exception(f"Failed to load allowlists: {exc}")
|
||
self.allowlist_table.add_column("Error")
|
||
self.allowlist_table.add_row(f"Failed to load: {str(exc)}")
|
||
|
||
async def extract_and_preview_hashes(self) -> None:
|
||
"""Extract hashes from selected data and show preview."""
|
||
preview_lines = ["## Hash Extraction Summary\n"]
|
||
|
||
# Check for hash column
|
||
if self.hash_column not in self.selected_data.columns:
|
||
# Try to find a hash column
|
||
possible_hash_cols = [
|
||
"sha256",
|
||
"SHA256",
|
||
"hash",
|
||
"Hash",
|
||
"sha1",
|
||
"SHA1",
|
||
"md5",
|
||
"MD5",
|
||
"filehash",
|
||
"file_hash",
|
||
]
|
||
found_col = None
|
||
for col in possible_hash_cols:
|
||
if col in self.selected_data.columns:
|
||
found_col = col
|
||
break
|
||
|
||
if found_col:
|
||
self.hash_column = found_col
|
||
preview_lines.append(f"âÅâ Found hash column: **{found_col}**\n")
|
||
else:
|
||
preview_lines.append("⚠︠**No hash column found**\n")
|
||
preview_lines.append("Available columns:\n")
|
||
for col in self.selected_data.columns:
|
||
if col != "_row_id":
|
||
preview_lines.append(f" - {col}\n")
|
||
|
||
self.preview_area.text = "".join(preview_lines)
|
||
self.stats_label.update("No hashes to add")
|
||
return
|
||
|
||
# Extract unique hashes
|
||
hashes = self.selected_data[self.hash_column].dropna().unique()
|
||
self.hashes_to_add = [h for h in hashes if h and str(h).strip()]
|
||
|
||
# Build preview
|
||
preview_lines.append(f"### Found {len(self.hashes_to_add)} unique hashes\n\n")
|
||
|
||
# Show sample of hashes (first 10)
|
||
preview_lines.append("**Sample hashes to be added:**\n```\n")
|
||
for i, hash_val in enumerate(self.hashes_to_add[:10]):
|
||
preview_lines.append(f"{i+1}. {hash_val}\n")
|
||
if len(self.hashes_to_add) > 10:
|
||
preview_lines.append(f"... and {len(self.hashes_to_add) - 10} more\n")
|
||
preview_lines.append("```\n\n")
|
||
|
||
# Show sample of source data
|
||
preview_lines.append("**Sample source activities:**\n")
|
||
sample_cols = [
|
||
col
|
||
for col in self.selected_data.columns
|
||
if col not in ["_row_id"] and col in ["filename", "path", "action", "user"]
|
||
]
|
||
if not sample_cols:
|
||
sample_cols = [
|
||
col for col in self.selected_data.columns if col != "_row_id"
|
||
][:3]
|
||
|
||
if sample_cols:
|
||
preview_lines.append("```\n")
|
||
for i, row in self.selected_data[sample_cols].head(5).iterrows():
|
||
row_text = " | ".join([f"{col}: {row[col]}" for col in sample_cols])
|
||
preview_lines.append(f"{row_text}\n")
|
||
preview_lines.append("```\n")
|
||
|
||
self.preview_area.text = "".join(preview_lines)
|
||
|
||
# Update statistics
|
||
self.stats_label.update(
|
||
f"Ready to add {len(self.hashes_to_add)} unique hashes | "
|
||
f"From {len(self.selected_data)} selected activities"
|
||
)
|
||
|
||
async def on_data_table_row_selected(self, event) -> None:
|
||
"""Handle allowlist selection."""
|
||
try:
|
||
# Extract row index from event - handle different event structures
|
||
row_index = None
|
||
|
||
# Try to get row index from coordinate
|
||
if hasattr(event, "coordinate") and hasattr(event.coordinate, "row"):
|
||
row_index = event.coordinate.row
|
||
# Try cursor_row as fallback
|
||
elif hasattr(event, "cursor_row"):
|
||
row_index = event.cursor_row
|
||
# Try getting from the table itself
|
||
else:
|
||
table = self.allowlist_table
|
||
if hasattr(table, "cursor_row"):
|
||
row_index = table.cursor_row
|
||
|
||
# Validate row index
|
||
if row_index is not None and isinstance(row_index, int):
|
||
# Account for group headers in the row count
|
||
actual_allowlist_index = self._get_allowlist_index_from_row(row_index)
|
||
|
||
if (
|
||
actual_allowlist_index is not None
|
||
and 0 <= actual_allowlist_index < len(self.allowlists)
|
||
):
|
||
self.selected_allowlist = self.allowlists[actual_allowlist_index]
|
||
self.add_btn.disabled = False
|
||
self.add_btn.label = (
|
||
f"â Add to '{self.selected_allowlist.get('name', 'Unknown')}'"
|
||
)
|
||
|
||
# Update preview with selection
|
||
await self._update_preview_with_selection()
|
||
|
||
logger.info(
|
||
f"Selected allowlist: {self.selected_allowlist.get('name')}"
|
||
)
|
||
else:
|
||
logger.debug(f"Row {row_index} is a header or invalid")
|
||
else:
|
||
logger.warning(f"Could not extract valid row index from event: {event}")
|
||
|
||
except Exception as exc:
|
||
logger.exception(f"Failed to select allowlist: {exc}")
|
||
|
||
def _get_allowlist_index_from_row(self, row_index: int) -> Optional[int]:
|
||
"""Convert table row index to allowlist list index, accounting for group headers."""
|
||
# This will be updated when we have group headers
|
||
if hasattr(self, "_row_to_allowlist_map"):
|
||
return self._row_to_allowlist_map.get(row_index)
|
||
return row_index
|
||
|
||
async def _update_preview_with_selection(self) -> None:
|
||
"""Update preview when an allowlist is selected."""
|
||
if not self.selected_allowlist:
|
||
return
|
||
|
||
current_text = self.preview_area.text
|
||
# Remove any existing selection header
|
||
if "### Selected Allowlist:" in current_text:
|
||
lines = current_text.split("\n")
|
||
# Find and remove the selection lines
|
||
new_lines = []
|
||
skip_next = False
|
||
for line in lines:
|
||
if line.startswith("### Selected Allowlist:"):
|
||
skip_next = True
|
||
continue
|
||
if skip_next and line.startswith("Application ID:"):
|
||
skip_next = False
|
||
continue
|
||
if not skip_next:
|
||
new_lines.append(line)
|
||
current_text = "\n".join(new_lines)
|
||
|
||
# Add new selection at the top
|
||
selection_text = (
|
||
f"### Selected Allowlist: **{self.selected_allowlist.get('name')}**\n"
|
||
f"Application ID: {self.selected_allowlist.get('applicationid')}\n\n"
|
||
)
|
||
self.preview_area.text = selection_text + current_text
|
||
|
||
async def on_button_pressed(self, event) -> None:
|
||
"""Handle button presses."""
|
||
btn = getattr(event, "button", None) or getattr(event, "sender", None)
|
||
btn_id = getattr(btn, "id", None) or getattr(event, "button_id", None)
|
||
|
||
if btn is self.back_btn or btn_id == "back_btn":
|
||
await self.app.pop_screen()
|
||
event.stop()
|
||
return
|
||
|
||
if btn is self.refresh_btn or btn_id == "refresh_allowlists_btn":
|
||
await self.load_allowlists()
|
||
event.stop()
|
||
return
|
||
|
||
if btn is self.add_btn or btn_id == "add_to_allowlist_btn":
|
||
await self.add_hashes_to_allowlist()
|
||
event.stop()
|
||
return
|
||
|
||
async def add_hashes_to_allowlist(self) -> None:
|
||
"""Add the extracted hashes to the selected allowlist."""
|
||
if not self.selected_allowlist or not self.hashes_to_add:
|
||
self.app.notify(
|
||
"No allowlist selected or no hashes to add", severity="warning"
|
||
)
|
||
return
|
||
|
||
if not self.api:
|
||
self.app.notify("API not available", severity="error")
|
||
return
|
||
|
||
try:
|
||
# Disable button during operation
|
||
self.add_btn.disabled = True
|
||
self.add_btn.label = "⏳ Adding hashes..."
|
||
|
||
# Call API to add hashes
|
||
app_id = self.selected_allowlist.get("applicationid")
|
||
allowlist_name = self.selected_allowlist.get("name", "Unknown")
|
||
|
||
logger.info(
|
||
f"Adding {len(self.hashes_to_add)} hashes to allowlist {allowlist_name} (ID: {app_id})"
|
||
)
|
||
|
||
result = self.api.hash_add_to_allowlist(app_id, self.hashes_to_add)
|
||
|
||
# Success notification
|
||
self.app.notify(
|
||
f"✅ Successfully added {len(self.hashes_to_add)} hashes to '{allowlist_name}'",
|
||
title="Success",
|
||
severity="information",
|
||
timeout=5,
|
||
)
|
||
|
||
# Update preview to show success
|
||
self.preview_area.text = (
|
||
f"## ✅ SUCCESS\n\n"
|
||
f"Added **{len(self.hashes_to_add)} hashes** to allowlist:\n"
|
||
f"**{allowlist_name}** (ID: {app_id})\n\n"
|
||
f"### Operation Details:\n"
|
||
f"- Source: {self.hostname or 'Multiple hosts'}\n"
|
||
f"- OTP ID: {self.otpid or 'N/A'}\n"
|
||
f"- Activities processed: {len(self.selected_data)}\n"
|
||
f"- Unique hashes added: {len(self.hashes_to_add)}\n"
|
||
)
|
||
|
||
# Change button to "Done"
|
||
self.add_btn.label = "✅ Done"
|
||
self.add_btn.disabled = True
|
||
|
||
except Exception as exc:
|
||
logger.exception(f"Failed to add hashes to allowlist: {exc}")
|
||
self.app.notify(
|
||
f"❌ Failed to add hashes: {str(exc)}",
|
||
title="Error",
|
||
severity="error",
|
||
timeout=10,
|
||
)
|
||
|
||
# Re-enable button
|
||
self.add_btn.disabled = False
|
||
self.add_btn.label = "⟳ Retry Add to Allowlist"
|
||
|
||
|
||
class AllowlistSelectionScreen(Screen):
|
||
"""
|
||
Screen wrapper for the AllowlistSelectionWidget.
|
||
"""
|
||
|
||
BINDINGS = [
|
||
Binding("b", "back", "Back"),
|
||
Binding("r", "refresh", "Refresh Allowlists"),
|
||
Binding("enter", "confirm", "Add to Allowlist"),
|
||
]
|
||
|
||
def __init__(
|
||
self,
|
||
selected_data: pd.DataFrame,
|
||
api=None,
|
||
hostname: Optional[str] = None,
|
||
otpid: Optional[str] = None,
|
||
hash_column: str = "sha256",
|
||
):
|
||
super().__init__()
|
||
self.selected_data = selected_data
|
||
self.api = api
|
||
self.hostname = hostname
|
||
self.otpid = otpid
|
||
self.hash_column = hash_column
|
||
|
||
def compose(self) -> ComposeResult:
|
||
yield Header(show_clock=True)
|
||
self.widget = AllowlistSelectionWidget(
|
||
self.selected_data,
|
||
api=self.api,
|
||
hostname=self.hostname,
|
||
otpid=self.otpid,
|
||
hash_column=self.hash_column,
|
||
)
|
||
yield self.widget
|
||
yield Footer()
|
||
|
||
async def action_back(self) -> None:
|
||
"""Go back to previous screen."""
|
||
await self.app.pop_screen()
|
||
|
||
async def action_refresh(self) -> None:
|
||
"""Refresh the allowlists."""
|
||
if hasattr(self, "widget") and self.widget:
|
||
await self.widget.load_allowlists()
|
||
|
||
async def action_confirm(self) -> None:
|
||
"""Confirm and add to allowlist."""
|
||
if hasattr(self, "widget") and self.widget:
|
||
if self.widget.selected_allowlist and self.widget.hashes_to_add:
|
||
await self.widget.add_hashes_to_allowlist()
|