Updated OTP Generate with new workflow
This commit is contained in:
@@ -8,6 +8,7 @@ Requests==2.32.5
|
||||
textual==6.5.0
|
||||
tqdm==4.67.1
|
||||
urllib3==2.5.0
|
||||
pyperclip==1.11.0
|
||||
|
||||
--extra-index-url https://git.racooncity.org/api/packages/brotoskyj/pypi/simple/
|
||||
airlock_libs==2.0.0
|
||||
@@ -0,0 +1,41 @@
|
||||
from typing import List
|
||||
|
||||
from textual.app import ComposeResult
|
||||
from textual.screen import Screen
|
||||
|
||||
from models.agent import Agent
|
||||
from widgets.multiagentselector import MultiAgentSelector
|
||||
from widgets.OTP_generate import OTPGenerator
|
||||
|
||||
|
||||
class OTPWorkflowScreen(Screen):
|
||||
"""Screen that handles the OTP generation workflow."""
|
||||
|
||||
def __init__(self, all_agents: List[Agent]):
|
||||
super().__init__()
|
||||
self.all_agents = all_agents
|
||||
self.selected_devices = None
|
||||
|
||||
def compose(self) -> ComposeResult:
|
||||
"""Start with the multi-agent selector."""
|
||||
yield MultiAgentSelector(self.all_agents)
|
||||
|
||||
def on_multi_agent_selector_agents_selected(
|
||||
self, message: MultiAgentSelector.AgentsSelected
|
||||
) -> None:
|
||||
"""Handle selected agents - switch to OTP generator."""
|
||||
self.selected_devices = message.selected_agents
|
||||
|
||||
# Remove the MultiAgentSelector
|
||||
selector = self.query_one(MultiAgentSelector)
|
||||
selector.remove()
|
||||
|
||||
# Mount the OTPGenerator with the selected Agent objects
|
||||
# No need to pass API - it will access self.app.api directly
|
||||
self.mount(OTPGenerator(self.selected_devices))
|
||||
|
||||
def on_otp_generator_otp_info(self, message: OTPGenerator.OTPInfo) -> None:
|
||||
"""Handle OTP generation request - call the actual OTP generation function."""
|
||||
# This will be handled by the main app, but we can also do it here
|
||||
# For now, just pass it up to the app level
|
||||
pass
|
||||
+79
-12
@@ -18,9 +18,12 @@ from textual.widgets import (
|
||||
Tabs,
|
||||
)
|
||||
|
||||
from flows.otp import otp_activities_by_agent, otp_generate, otp_revoke
|
||||
from flows.otp import otp_activities_by_agent, otp_revoke
|
||||
from flows.prepPolicy import menu_policy_enforce
|
||||
from flows.quietAgent import findQuietAgents
|
||||
from models.agent import Agent
|
||||
from models.policy import Policy
|
||||
from screens.otpworkflowscreen import OTPWorkflowScreen
|
||||
from services.agenthandler import findAgents, moveAgents, toggleEnforcement
|
||||
from services.API import AirlockAPIWrapper
|
||||
from services.policyhandler import confirmUpdateAfromE
|
||||
@@ -28,6 +31,7 @@ from utils.configmanager import load_env
|
||||
from utils.setup import get_base_directory, load_user_config
|
||||
from utils.utils import open_directory
|
||||
from widgets.multiagentselector import MultiAgentSelector
|
||||
from widgets.OTP_generate import OTPGenerator
|
||||
from widgets.policytreewidget import PolicyTreeWidget
|
||||
from widgets.themeselector import ThemeSelector
|
||||
|
||||
@@ -107,7 +111,7 @@ class MainMenuScreen(Screen):
|
||||
("🔀 - Move - Other", "move_other_button"),
|
||||
],
|
||||
"otp": [
|
||||
("🔐 - Generate OTPs", "otp_generate_button"),
|
||||
("🎫 - Generate OTPs", "otp_generate_button"),
|
||||
("📊 - OTP Activities By Agent", "otp_activities_button"),
|
||||
("❌ - Revoke OTPs", "otp_revoke_button"),
|
||||
],
|
||||
@@ -117,8 +121,9 @@ class MainMenuScreen(Screen):
|
||||
],
|
||||
}
|
||||
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, api: AirlockAPIWrapper) -> None:
|
||||
super().__init__()
|
||||
self.api = api
|
||||
self.extras = load_env("EXTRAS")
|
||||
wd = load_env("WORKING_DIR") or os.getcwd()
|
||||
if not os.path.isdir(wd):
|
||||
@@ -144,7 +149,6 @@ class MainMenuScreen(Screen):
|
||||
Tab("OTP", id="otp"),
|
||||
Tab("Directory", id="dir"),
|
||||
Tab("Settings", id="settings"),
|
||||
Tab("Multi Select", id="multi_select"),
|
||||
]
|
||||
|
||||
if self.extras == "POLICYPREP":
|
||||
@@ -207,8 +211,6 @@ class MainMenuScreen(Screen):
|
||||
content.mount(PolicyTreeWidget(self.app.policies, self.app.devices))
|
||||
elif tab_id == "settings":
|
||||
content.mount(ThemeSelector())
|
||||
elif tab_id == "multi_select":
|
||||
content.mount(MultiAgentSelector(self.app.devices.to_dict("records")))
|
||||
else:
|
||||
content.mount(Static(f"Unknown tab: {tab_id}"))
|
||||
|
||||
@@ -235,6 +237,30 @@ class MainMenuScreen(Screen):
|
||||
_PENDING_JOB = ("restart",)
|
||||
self.app.exit()
|
||||
|
||||
def on_otp_generator_otp_info(self, message: OTPGenerator.OTPInfo) -> None:
|
||||
"""Handle OTP generation request from the workflow."""
|
||||
global _PENDING_JOB
|
||||
|
||||
# Log what we received
|
||||
logger.info(
|
||||
"OTP Generation requested: %d devices, requestor=%s, reason=%s, duration=%d",
|
||||
len(message.devices),
|
||||
message.requestor,
|
||||
message.reasoning,
|
||||
message.duration,
|
||||
)
|
||||
|
||||
# Set up the job to run the OTP generation
|
||||
_PENDING_JOB = (
|
||||
"otp_workflow",
|
||||
message.devices,
|
||||
message.requestor,
|
||||
message.reasoning,
|
||||
message.duration,
|
||||
)
|
||||
|
||||
self.app.exit()
|
||||
|
||||
def on_directory_tree_file_selected(
|
||||
self, event: DirectoryTree.FileSelected
|
||||
) -> None:
|
||||
@@ -268,7 +294,10 @@ class MainMenuScreen(Screen):
|
||||
case "move_other_button":
|
||||
_PENDING_JOB = ("legacy", moveAgents, (self.app.api,), {})
|
||||
case "otp_generate_button":
|
||||
_PENDING_JOB = ("legacy", otp_generate, (self.app.api,), {})
|
||||
# NEW: Push OTP workflow screen instead of legacy function
|
||||
self.app.push_screen(OTPWorkflowScreen(self.app.devices))
|
||||
event.stop()
|
||||
return # Don't exit the app
|
||||
case "otp_activities_button":
|
||||
_PENDING_JOB = ("legacy", otp_activities_by_agent, (self.app.api,), {})
|
||||
case "otp_revoke_button":
|
||||
@@ -314,16 +343,20 @@ class Loxide(App):
|
||||
|
||||
# Add error handling for API calls
|
||||
try:
|
||||
self.policies = api.policy_find_all()
|
||||
self.devices = api.agent_find_all()
|
||||
self.policies = [
|
||||
Policy(**row.to_dict()) for _, row in api.policy_find_all().iterrows()
|
||||
]
|
||||
self.devices = [
|
||||
Agent(**row.to_dict()) for _, row in api.agent_find_all().iterrows()
|
||||
]
|
||||
except Exception as exc:
|
||||
logger.error("Failed to load policies/devices: %s", exc)
|
||||
self.policies = None
|
||||
self.devices = None
|
||||
|
||||
def on_mount(self) -> None:
|
||||
def on_mount(self, api: AirlockAPIWrapper) -> None:
|
||||
self.theme = self._textual_theme
|
||||
self.push_screen(MainMenuScreen())
|
||||
self.push_screen(MainMenuScreen(api))
|
||||
|
||||
def action_quit(self) -> None:
|
||||
global _PENDING_JOB
|
||||
@@ -410,10 +443,44 @@ def run_Loxide(api: AirlockAPIWrapper) -> None:
|
||||
|
||||
if job[0] == "multi_agent_action":
|
||||
# Handle multi-agent selection
|
||||
# TODO: Implement actual multi-agent action handling
|
||||
logger.info("Multi-agent action with selected agents: %s", job[1])
|
||||
continue
|
||||
|
||||
# NEW: Handle OTP workflow
|
||||
if job[0] == "otp_workflow":
|
||||
_, devices, requestor, reasoning, duration = job
|
||||
|
||||
# Call your OTP generation with the parameters
|
||||
def otp_generate_with_params():
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print("OTP GENERATION")
|
||||
print(f"{'='*60}")
|
||||
print(f"Requestor: {requestor}")
|
||||
print(f"Reasoning: {reasoning}")
|
||||
print(f"Duration: {duration} minutes")
|
||||
print(f"\nGenerating OTPs for {len(devices)} devices:")
|
||||
print(f"{'='*60}\n")
|
||||
|
||||
# Call your actual OTP generation function
|
||||
# You'll need to adapt otp_generate to accept these parameters
|
||||
# For now, this is a placeholder showing the structure
|
||||
for device in devices:
|
||||
print(f"Device: {device}")
|
||||
print(f" Requestor: {requestor}")
|
||||
print(f" Reason: {reasoning}")
|
||||
print(f" Duration: {duration} minutes")
|
||||
# TODO: Actually call your API to generate OTP
|
||||
# result = api.generate_otp(device, requestor, reasoning, duration)
|
||||
print()
|
||||
|
||||
print(f"{'='*60}")
|
||||
print("OTP Generation Complete!")
|
||||
print(f"{'='*60}")
|
||||
|
||||
_run_legacy_job(otp_generate_with_params, (), {})
|
||||
continue
|
||||
|
||||
break
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,346 @@
|
||||
import logging
|
||||
from typing import List
|
||||
|
||||
from textual.containers import Horizontal, Vertical
|
||||
from textual.css.query import NoMatches
|
||||
from textual.message import Message
|
||||
from textual.reactive import reactive
|
||||
from textual.widget import Widget
|
||||
from textual.widgets import Button, Input, RadioButton, RadioSet, Static, TextArea
|
||||
|
||||
from models.agent import Agent
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OTPGenerator(Widget):
|
||||
# Reactive properties to track form completion
|
||||
requestor_filled = reactive(False)
|
||||
reasoning_filled = reactive(False)
|
||||
duration_selected = reactive(True) # Default is selected
|
||||
otp_generated = reactive(False)
|
||||
|
||||
class OTPInfo(Message):
|
||||
def __init__(
|
||||
self, devices: List[Agent], requestor: str, reasoning: str, duration: int
|
||||
):
|
||||
super().__init__()
|
||||
self.devices = devices
|
||||
self.requestor = requestor
|
||||
self.reasoning = reasoning
|
||||
self.duration = duration
|
||||
|
||||
# Duration options in minutes
|
||||
DURATION_OPTIONS = [
|
||||
(15, "15 minutes"),
|
||||
(60, "1 hour"),
|
||||
(360, "6 hours"),
|
||||
(1440, "1 day"),
|
||||
(10080, "7 days"),
|
||||
]
|
||||
|
||||
def __init__(self, devices: List[Agent]):
|
||||
"""Initialize with a list of Agent objects."""
|
||||
super().__init__()
|
||||
self.devices = devices
|
||||
|
||||
def watch_requestor_filled(self, old_value: bool, new_value: bool) -> None:
|
||||
"""Update button state when requestor changes."""
|
||||
self._update_button_state()
|
||||
|
||||
def watch_reasoning_filled(self, old_value: bool, new_value: bool) -> None:
|
||||
"""Update button state when reasoning changes."""
|
||||
self._update_button_state()
|
||||
|
||||
def watch_otp_generated(self, old_value: bool, new_value: bool) -> None:
|
||||
"""Update button state when OTP is generated."""
|
||||
self._update_button_state()
|
||||
|
||||
def _update_button_state(self) -> None:
|
||||
"""Enable/disable the generate button based on form state."""
|
||||
try:
|
||||
button = self.query_one("#generate_button", Button)
|
||||
# Enable only if all fields filled and OTP not yet generated
|
||||
button.disabled = not (
|
||||
self.requestor_filled
|
||||
and self.reasoning_filled
|
||||
and not self.otp_generated
|
||||
)
|
||||
except NoMatches:
|
||||
pass
|
||||
|
||||
def compose(self):
|
||||
title_text = Static(
|
||||
f"🎫 Generate One Time Passes for {len(self.devices)} device(s)",
|
||||
id="otpgen_title",
|
||||
)
|
||||
title_text.styles.margin = (0, 0, 1, 0)
|
||||
yield title_text
|
||||
|
||||
with Horizontal() as main_layout:
|
||||
main_layout.styles.height = "auto"
|
||||
|
||||
# Left side - Inputs and controls
|
||||
with Vertical() as left_side:
|
||||
left_side.styles.width = "1fr"
|
||||
left_side.styles.height = "auto"
|
||||
|
||||
# Requestor input
|
||||
requestor_label = Static("Who is requesting OTP?")
|
||||
requestor_label.styles.margin = (0, 0, 0, 0)
|
||||
yield requestor_label
|
||||
|
||||
requestor_box = Input(
|
||||
placeholder="Enter requestor name", id="requestor_input"
|
||||
)
|
||||
requestor_box.styles.margin = (0, 0, 1, 0)
|
||||
yield requestor_box
|
||||
|
||||
# Reasoning input
|
||||
reasoning_label = Static("What work are they doing?")
|
||||
reasoning_label.styles.margin = (0, 0, 0, 0)
|
||||
yield reasoning_label
|
||||
|
||||
reasoning_box = Input(
|
||||
placeholder="Enter reason for OTP", id="reasoning_input"
|
||||
)
|
||||
reasoning_box.styles.margin = (0, 0, 1, 0)
|
||||
yield reasoning_box
|
||||
|
||||
# Duration selection
|
||||
duration_label = Static("Duration:")
|
||||
duration_label.styles.margin = (0, 0, 0, 0)
|
||||
yield duration_label
|
||||
|
||||
with RadioSet(id="duration_radio") as radio_set:
|
||||
radio_set.styles.margin = (0, 0, 1, 0)
|
||||
for minutes, label in self.DURATION_OPTIONS:
|
||||
radio = RadioButton(label, id=f"duration_{minutes}")
|
||||
if minutes == 360: # Default to 6 hours
|
||||
radio.value = True
|
||||
yield radio
|
||||
|
||||
# Buttons in a horizontal layout
|
||||
with Horizontal() as button_row:
|
||||
button_row.styles.height = "auto"
|
||||
button_row.styles.margin = (1, 0, 0, 0)
|
||||
|
||||
back_button = Button("← Back", id="back_button")
|
||||
back_button.styles.width = "1fr"
|
||||
yield back_button
|
||||
|
||||
generate_button = Button(
|
||||
"Generate OTP", id="generate_button", variant="primary"
|
||||
)
|
||||
generate_button.styles.width = "2fr"
|
||||
yield generate_button
|
||||
|
||||
# Right side - Show device list initially, then output after generation
|
||||
with Vertical() as right_side:
|
||||
right_side.styles.width = "2fr"
|
||||
right_side.styles.height = "100%"
|
||||
|
||||
output_label = Static(
|
||||
f"Selected Devices ({len(self.devices)}):", id="output_label"
|
||||
)
|
||||
output_label.styles.margin = (0, 0, 0, 0)
|
||||
yield output_label
|
||||
|
||||
# Container for either device list or output
|
||||
with Vertical(id="output_container") as output_container:
|
||||
output_container.styles.height = "1fr"
|
||||
output_container.styles.margin = (1, 0, 0, 0)
|
||||
output_container.styles.overflow_y = "auto"
|
||||
output_container.styles.border = ("round", "green")
|
||||
|
||||
# Show device list initially
|
||||
device_list_text = "\n".join(
|
||||
f"• {device.hostname}" for device in self.devices
|
||||
)
|
||||
device_display = Static(device_list_text, id="device_display")
|
||||
yield device_display
|
||||
|
||||
# Copy to clipboard button (hidden initially)
|
||||
copy_button = Button("📋 Copy to Clipboard", id="copy_clipboard_button")
|
||||
copy_button.styles.margin = (1, 0, 0, 0)
|
||||
copy_button.styles.display = "none"
|
||||
yield copy_button
|
||||
|
||||
def on_mount(self) -> None:
|
||||
"""Set initial button state."""
|
||||
self._update_button_state()
|
||||
|
||||
def on_input_changed(self, event: Input.Changed) -> None:
|
||||
"""Handle input field changes."""
|
||||
input_id = event.input.id
|
||||
|
||||
if input_id == "requestor_input":
|
||||
self.requestor_filled = bool(event.value.strip())
|
||||
elif input_id == "reasoning_input":
|
||||
self.reasoning_filled = bool(event.value.strip())
|
||||
|
||||
def on_button_pressed(self, event: Button.Pressed):
|
||||
btn_id = event.button.id
|
||||
|
||||
if btn_id == "back_button":
|
||||
self.app.pop_screen()
|
||||
event.stop()
|
||||
|
||||
elif btn_id == "copy_clipboard_button":
|
||||
try:
|
||||
output_area = self.query_one("#otp_output", TextArea)
|
||||
text_to_copy = output_area.text
|
||||
|
||||
import pyperclip
|
||||
|
||||
pyperclip.copy(text_to_copy)
|
||||
self.app.notify(
|
||||
"✅ Copied to clipboard!", severity="information", timeout=2
|
||||
)
|
||||
except ImportError:
|
||||
self.app.notify(
|
||||
"⚠️ pyperclip not installed. Run: pip install pyperclip",
|
||||
severity="warning",
|
||||
)
|
||||
except Exception as e:
|
||||
self.app.notify(f"❌ Failed to copy: {str(e)}", severity="error")
|
||||
event.stop()
|
||||
|
||||
elif btn_id == "generate_button":
|
||||
try:
|
||||
requestor = self.query_one("#requestor_input", Input).value.strip()
|
||||
reasoning = self.query_one("#reasoning_input", Input).value.strip()
|
||||
|
||||
radio_set = self.query_one("#duration_radio", RadioSet)
|
||||
selected_button_id = (
|
||||
radio_set.pressed_button.id if radio_set.pressed_button else None
|
||||
)
|
||||
|
||||
if not selected_button_id:
|
||||
self._show_error("Please select a duration")
|
||||
return
|
||||
|
||||
duration = int(selected_button_id.replace("duration_", ""))
|
||||
|
||||
if not requestor or not reasoning:
|
||||
self._show_error("Please fill in all fields")
|
||||
return
|
||||
|
||||
self.otp_generated = True
|
||||
|
||||
# Access API from the app - this is the key change!
|
||||
api = self.app.api
|
||||
|
||||
output_lines = [
|
||||
"=" * 60,
|
||||
"OTP GENERATION RESULTS",
|
||||
"=" * 60,
|
||||
]
|
||||
|
||||
otp_dict = {}
|
||||
for device in self.devices:
|
||||
try:
|
||||
otp_code = api.otp_generate(device.agentid, duration, reasoning)
|
||||
otp_dict[device.hostname] = otp_code
|
||||
logger.debug(f"Generated OTP for {device.hostname}: {otp_code}")
|
||||
except Exception as e:
|
||||
otp_dict[device.hostname] = f"ERROR: {str(e)}"
|
||||
logger.error(
|
||||
f"Failed to generate OTP for {device.hostname}: {e}"
|
||||
)
|
||||
|
||||
for hostname, otp_code in otp_dict.items():
|
||||
output_lines.append(f"{hostname:30} | {otp_code}")
|
||||
|
||||
output_lines.append("=" * 60)
|
||||
result_text = "\n".join(output_lines)
|
||||
self._show_result(result_text)
|
||||
|
||||
# Post message with the OTP info
|
||||
self.post_message(
|
||||
self.OTPInfo(self.devices, requestor, reasoning, duration)
|
||||
)
|
||||
event.stop()
|
||||
|
||||
except NoMatches:
|
||||
self._show_error("UI elements not found")
|
||||
except Exception as e:
|
||||
self._show_error(f"Error: {str(e)}")
|
||||
logger.exception("Error generating OTP")
|
||||
|
||||
def _show_error(self, message: str):
|
||||
"""Display error message in output area."""
|
||||
try:
|
||||
container = self.query_one("#output_container", Vertical)
|
||||
try:
|
||||
device_display = self.query_one("#device_display", Static)
|
||||
device_display.remove()
|
||||
except NoMatches:
|
||||
pass
|
||||
|
||||
try:
|
||||
output_area = self.query_one("#otp_output", TextArea)
|
||||
except NoMatches:
|
||||
output_area = TextArea(id="otp_output", read_only=True)
|
||||
container.mount(output_area)
|
||||
|
||||
output_area.text = f"❌ ERROR: {message}"
|
||||
except Exception as e:
|
||||
logger.debug(f"Error showing error message: {e}")
|
||||
|
||||
def _show_result(self, message: str):
|
||||
"""Display result message in output area."""
|
||||
try:
|
||||
container = self.query_one("#output_container", Vertical)
|
||||
try:
|
||||
device_display = self.query_one("#device_display", Static)
|
||||
device_display.remove()
|
||||
except NoMatches:
|
||||
pass
|
||||
|
||||
try:
|
||||
output_area = self.query_one("#otp_output", TextArea)
|
||||
except NoMatches:
|
||||
output_area = TextArea(id="otp_output", read_only=True)
|
||||
container.mount(output_area)
|
||||
|
||||
output_area.text = message
|
||||
output_label = self.query_one("#output_label", Static)
|
||||
output_label.update("Generated OTP Details:")
|
||||
copy_button = self.query_one("#copy_clipboard_button", Button)
|
||||
copy_button.styles.display = "block"
|
||||
except Exception as e:
|
||||
logger.debug(f"Error showing result: {e}")
|
||||
|
||||
def display_otp_result(self, result_text: str):
|
||||
"""Display OTP generation result in the output area."""
|
||||
try:
|
||||
container = self.query_one("#output_container", Vertical)
|
||||
try:
|
||||
device_display = self.query_one("#device_display", Static)
|
||||
device_display.remove()
|
||||
except NoMatches:
|
||||
pass
|
||||
|
||||
try:
|
||||
output_area = self.query_one("#otp_output", TextArea)
|
||||
except NoMatches:
|
||||
output_area = TextArea(id="otp_output", read_only=True)
|
||||
container.mount(output_area)
|
||||
|
||||
output_area.text = result_text
|
||||
except Exception as e:
|
||||
logger.debug(f"Error displaying OTP result: {e}")
|
||||
|
||||
def clear_form(self):
|
||||
"""Clear all input fields and reset state."""
|
||||
try:
|
||||
self.query_one("#requestor_input", Input).value = ""
|
||||
self.query_one("#reasoning_input", Input).value = ""
|
||||
self.query_one("#otp_output", TextArea).text = ""
|
||||
self.otp_generated = False
|
||||
self.requestor_filled = False
|
||||
self.reasoning_filled = False
|
||||
self._update_button_state()
|
||||
except NoMatches:
|
||||
pass
|
||||
@@ -1,4 +1,6 @@
|
||||
import difflib
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
from textual.containers import Horizontal, Vertical
|
||||
from textual.css.query import NoMatches
|
||||
@@ -6,14 +8,16 @@ from textual.message import Message
|
||||
from textual.widget import Widget
|
||||
from textual.widgets import Button, SelectionList, Static, Switch, TextArea
|
||||
|
||||
from models.agent import Agent
|
||||
|
||||
|
||||
class MultiAgentSelector(Widget):
|
||||
class AgentsSelected(Message):
|
||||
def __init__(self, selected_agents):
|
||||
def __init__(self, selected_agents: List[Agent]):
|
||||
super().__init__()
|
||||
self.selected_agents = selected_agents
|
||||
|
||||
def __init__(self, all_agents: list[dict]):
|
||||
def __init__(self, all_agents: List[Agent]):
|
||||
super().__init__()
|
||||
self.all_agents = all_agents
|
||||
self._match_type = "exact"
|
||||
@@ -48,7 +52,6 @@ class MultiAgentSelector(Widget):
|
||||
yield text_area
|
||||
|
||||
with Horizontal(id="switch_search_container") as switch_search:
|
||||
|
||||
switch = Switch(value=False, id="match_switch")
|
||||
switch.styles.width = "auto"
|
||||
switch.styles.margin = (1, 0, 0, 0)
|
||||
@@ -62,7 +65,6 @@ class MultiAgentSelector(Widget):
|
||||
|
||||
search = Button("🔍 Search", id="search_button")
|
||||
search.styles.margin = (1, 0, 0, 0)
|
||||
|
||||
yield search
|
||||
|
||||
with Horizontal() as select_buttons:
|
||||
@@ -76,8 +78,16 @@ class MultiAgentSelector(Widget):
|
||||
select_none_button.styles.margin = (1, 0, 0, 1)
|
||||
yield select_none_button
|
||||
|
||||
with Horizontal() as button_row:
|
||||
button_row.styles.height = "auto"
|
||||
button_row.styles.margin = (1, 0, 0, 0)
|
||||
|
||||
back_button = Button("← Back", id="back_button")
|
||||
back_button.styles.width = "1fr"
|
||||
yield back_button
|
||||
|
||||
submit_button = Button(
|
||||
"► Continue with Selected", id="submit_selection", variant="primary"
|
||||
"▶ Select & Continue", id="submit_selection", variant="primary"
|
||||
)
|
||||
submit_button.styles.margin = (0, 5, 2, 1)
|
||||
submit_button.styles.padding = (0, 6, 0, 0)
|
||||
@@ -86,7 +96,6 @@ class MultiAgentSelector(Widget):
|
||||
# Right side - Results
|
||||
with Vertical() as right_pane:
|
||||
right_pane.styles.width = "2fr"
|
||||
|
||||
yield SelectionList(id="match_results")
|
||||
yield Static(id="unmatched_label")
|
||||
|
||||
@@ -97,24 +106,31 @@ class MultiAgentSelector(Widget):
|
||||
)
|
||||
|
||||
def on_button_pressed(self, event: Button.Pressed):
|
||||
btn_id = event.button.id # can be None for internal buttons
|
||||
btn_id = event.button.id
|
||||
|
||||
# Only query when needed
|
||||
try:
|
||||
match_list = self.query_one("#match_results", SelectionList)
|
||||
except NoMatches:
|
||||
# UI not mounted yet or id changed—just ignore gracefully
|
||||
return
|
||||
|
||||
if btn_id == "select_all":
|
||||
if btn_id == "back_button":
|
||||
self.app.pop_screen()
|
||||
event.stop()
|
||||
elif btn_id == "select_all":
|
||||
match_list.select_all()
|
||||
event.stop()
|
||||
elif btn_id == "select_none":
|
||||
match_list.deselect_all()
|
||||
event.stop()
|
||||
elif btn_id == "submit_selection":
|
||||
selected = list(match_list.selected)
|
||||
self.post_message(self.AgentsSelected(selected))
|
||||
# Get selected hostnames
|
||||
selected_hostnames = list(match_list.selected)
|
||||
# Convert back to Agent objects
|
||||
selected_agents = [
|
||||
agent
|
||||
for agent in self.all_agents
|
||||
if agent.hostname in selected_hostnames
|
||||
]
|
||||
self.post_message(self.AgentsSelected(selected_agents))
|
||||
event.stop()
|
||||
elif btn_id == "search_button":
|
||||
self.update_matches()
|
||||
@@ -137,22 +153,18 @@ class MultiAgentSelector(Widget):
|
||||
def match_devices(self, device_names: list[str]) -> tuple[list[str], list[str]]:
|
||||
if not self.all_agents or not device_names:
|
||||
return [], device_names
|
||||
agent_names = [agent["hostname"] for agent in self.all_agents]
|
||||
agent_names = [agent.hostname for agent in self.all_agents]
|
||||
matched = set()
|
||||
unmatched = []
|
||||
|
||||
for name in device_names:
|
||||
# Check if the name contains wildcards
|
||||
has_wildcards = "*" in name or "?" in name
|
||||
|
||||
if has_wildcards:
|
||||
# Use regex for wildcard matching
|
||||
import re
|
||||
|
||||
# Escape special regex characters except * and ?
|
||||
pattern = re.escape(name)
|
||||
# Convert wildcards to regex
|
||||
pattern = pattern.replace(r"\*", ".*").replace(r"\?", ".")
|
||||
# Make it case-insensitive and match full string
|
||||
regex = re.compile(f"^{pattern}$", re.IGNORECASE)
|
||||
|
||||
wildcard_matches = [
|
||||
|
||||
+21
-14
@@ -49,26 +49,26 @@ class PolicyTreeWidget(Widget):
|
||||
node_map = {}
|
||||
|
||||
# Top-level policies
|
||||
for _, policy in self.policies.iterrows():
|
||||
if policy["parent"] == "global-policy-settings":
|
||||
node = policy_tree.root.add(label=policy["name"], data=policy.to_dict())
|
||||
node_map[policy["groupid"]] = node
|
||||
for policy in self.policies:
|
||||
if policy.parent == "global-policy-settings":
|
||||
node = policy_tree.root.add(label=policy.name, data=policy)
|
||||
node_map[policy.groupid] = node
|
||||
|
||||
# Child policies
|
||||
for _, policy in self.policies.iterrows():
|
||||
parent_id = policy["parent"]
|
||||
for policy in self.policies:
|
||||
parent_id = policy.parent
|
||||
if parent_id in node_map:
|
||||
parent_node = node_map[parent_id]
|
||||
node = parent_node.add(label=policy["name"], data=policy.to_dict())
|
||||
node_map[policy["groupid"]] = node
|
||||
node = parent_node.add(label=policy.name, data=policy)
|
||||
node_map[policy.groupid] = node
|
||||
|
||||
# Devices under policies
|
||||
for _, device in self.devices.iterrows():
|
||||
group_id = device["groupid"]
|
||||
for device in self.devices:
|
||||
group_id = device.groupid
|
||||
if group_id in node_map:
|
||||
parent_node = node_map[group_id]
|
||||
label = device["hostname"]
|
||||
parent_node.add(label=label, data=device.to_dict())
|
||||
label = device.hostname
|
||||
parent_node.add(label=label, data=device)
|
||||
|
||||
def _collect_tree_nodes(self, node, all_nodes):
|
||||
"""Helper to recursively collect all nodes from a tree."""
|
||||
@@ -110,7 +110,10 @@ class PolicyTreeWidget(Widget):
|
||||
|
||||
# Update details pane
|
||||
if data:
|
||||
details = "\n".join(f"{key}: {value}" for key, value in data.items())
|
||||
# Work with dataclass objects using __dict__
|
||||
details = "\n".join(
|
||||
f"{key}: {value}" for key, value in data.__dict__.items()
|
||||
)
|
||||
else:
|
||||
details = f"Selected: {node.label}"
|
||||
details_pane.update(details)
|
||||
@@ -135,7 +138,11 @@ class PolicyTreeWidget(Widget):
|
||||
label_text = str(node.label).lower()
|
||||
label_to_node[label_text] = node
|
||||
if node.data:
|
||||
for key, value in node.data.items():
|
||||
# Use __dict__ for dataclass objects
|
||||
data_dict = (
|
||||
node.data.__dict__ if hasattr(node.data, "__dict__") else node.data
|
||||
)
|
||||
for key, value in data_dict.items():
|
||||
if isinstance(value, str):
|
||||
label_to_node[value.lower()] = node
|
||||
|
||||
|
||||
Reference in New Issue
Block a user