main.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302
  1. import os
  2. import threading
  3. import gc
  4. from typing import Any, Dict
  5. import toml
  6. from textual import work
  7. from textual.app import App, ComposeResult
  8. from textual.containers import Horizontal, Vertical
  9. from textual.widgets import (
  10. Button,
  11. Footer,
  12. Header,
  13. Input,
  14. Label,
  15. ProgressBar,
  16. RichLog,
  17. Rule,
  18. )
  19. from tasks import PIPELINE_TASKS, SCENARIOS
  20. from util.config_manager import handle_directory_config
  21. from util.progress import ProgressTracker
  22. from util.screens import OverwriteConfirmScreen
  23. from util.ui_logger import PipelineLogger
  24. class CNNHarnessApp(App):
  25. CSS_PATH = "style.tcss"
  26. def __init__(self):
  27. super().__init__()
  28. self.stop_event = threading.Event()
  29. self.pause_event = threading.Event()
  30. self.current_config: Dict[str, Any] = {}
  31. self.ui_logger = PipelineLogger(self, widget_id="#live_log")
  32. # ==========================================
  33. # UI Layout
  34. # ==========================================
  35. def compose(self) -> ComposeResult:
  36. yield Header()
  37. with Vertical(id="main_container"):
  38. # --- TOP SECTION: Two-Column Configuration ---
  39. with Horizontal(id="config_section"):
  40. # Left Column: Path entry and buttons
  41. with Vertical(id="config_left_col"):
  42. yield Label("Working Directory:")
  43. yield Input(placeholder="./outputs/experiment_1", id="work_dir")
  44. yield Button("Load / Init Dir", variant="primary", id="load_btn")
  45. yield Button("Reload Config", variant="default", id="reload_btn")
  46. # Right Column: Read-only parameters display
  47. with Vertical(id="config_right_col"):
  48. yield Label("Current Parameters:")
  49. yield RichLog(id="config_display", highlight=True, markup=True)
  50. yield Rule()
  51. # --- BOTTOM SECTION: Controls, Progress, and Logs ---
  52. with Horizontal(id="bottom_area"):
  53. # --- Pipeline Controls ---
  54. with Vertical(id="control_panel"):
  55. yield Button("Start Pipeline", variant="success", id="btn_start")
  56. with Horizontal(id="active_controls"):
  57. yield Button("Pause", variant="warning", id="btn_pause")
  58. yield Button("Stop", variant="error", id="btn_stop")
  59. # --- Progress Trackers ---
  60. with Vertical(id="progress_area"):
  61. for i in range(5):
  62. with Horizontal(
  63. id=f"progress_row_{i}", classes="progress_row"
  64. ):
  65. yield Label(
  66. id=f"progress_title_{i}", classes="progress_title"
  67. )
  68. yield ProgressBar(
  69. id=f"progress_bar_{i}",
  70. classes="progress_bar",
  71. show_eta=True,
  72. )
  73. yield Label(
  74. id=f"progress_stats_{i}", classes="progress_stats"
  75. )
  76. # --- Live Log Output ---
  77. with Vertical(id="log_panel"):
  78. with Horizontal(id="log_status_row"):
  79. yield Label("STATUS: READY", id="pipeline_status_label")
  80. with Horizontal(id="log_header_row"):
  81. yield Label("Live Logs:", classes="section_label")
  82. yield RichLog(id="live_log", highlight=True, markup=True)
  83. yield Footer()
  84. def on_mount(self) -> None:
  85. self.query_one("#active_controls").display = False
  86. self.ui_logger.info("Application initialized and ready.")
  87. # ==========================================
  88. # Helper Methods
  89. # ==========================================
  90. def _set_pipeline_status(self, text: str) -> None:
  91. try:
  92. self.query_one("#pipeline_status_label", Label).update(text)
  93. except Exception:
  94. pass
  95. def _update_config_display(self) -> None:
  96. config_log = self.query_one("#config_display", RichLog)
  97. config_log.clear()
  98. if not self.current_config:
  99. config_log.write("[italic]No configuration loaded.[/italic]")
  100. return
  101. toml_string = toml.dumps(self.current_config)
  102. config_log.write(toml_string)
  103. def _toggle_config_inputs(self, disabled: bool) -> None:
  104. self.query_one("#work_dir", Input).disabled = disabled
  105. self.query_one("#load_btn", Button).disabled = disabled
  106. self.query_one("#reload_btn", Button).disabled = disabled
  107. # ==========================================
  108. # Event Handlers
  109. # ==========================================
  110. def on_button_pressed(self, event: Button.Pressed) -> None:
  111. if event.button.id in ("load_btn", "reload_btn"):
  112. work_dir = self.query_one("#work_dir", Input).value
  113. success, message, config_data = handle_directory_config(work_dir)
  114. if success:
  115. self.current_config = config_data
  116. self.current_config["work_dir"] = work_dir
  117. self._update_config_display()
  118. self.ui_logger.info(message)
  119. else:
  120. self.ui_logger.error(message)
  121. elif event.button.id == "btn_start":
  122. if not self.current_config:
  123. self.ui_logger.error(
  124. "No configuration loaded! Please load a directory first."
  125. )
  126. return
  127. scenario_id = self.current_config.get("scenario")
  128. if not scenario_id or scenario_id not in SCENARIOS:
  129. self.ui_logger.error(
  130. f"Invalid or missing scenario '{scenario_id}' in config.toml!"
  131. )
  132. return
  133. work_dir = self.current_config["work_dir"]
  134. # Check if the working directory exists and is non-empty (except for the config.toml file)
  135. if os.path.exists(work_dir) and any(
  136. f for f in os.listdir(work_dir) if f != "config.toml"
  137. ):
  138. def check_overwrite_callback(proceed: bool) -> None:
  139. if proceed:
  140. self._execute_pipeline()
  141. else:
  142. self.ui_logger.error(
  143. "Pipeline start cancelled by user (folder non-empty)."
  144. )
  145. self.app.push_screen(
  146. OverwriteConfirmScreen(work_dir),
  147. check_overwrite_callback, # pyright: ignore
  148. )
  149. else:
  150. self._execute_pipeline()
  151. elif event.button.id == "btn_pause":
  152. if self.pause_event.is_set():
  153. self.pause_event.clear()
  154. event.button.label = "Pause"
  155. event.button.variant = "warning"
  156. self._set_pipeline_status("[bold green]PIPELINE RUNNING[/bold green]")
  157. self.ui_logger.info("Pipeline Resumed.")
  158. else:
  159. self.pause_event.set()
  160. event.button.label = "Resume"
  161. event.button.variant = "success"
  162. self._set_pipeline_status("[bold yellow]PIPELINE PAUSED[/bold yellow]")
  163. self.ui_logger.info(
  164. "Pipeline Paused. Waiting for current operation to yield..."
  165. )
  166. elif event.button.id == "btn_stop":
  167. self._set_pipeline_status("[bold red]STOPPING PIPELINE...[/bold red]")
  168. self.ui_logger.error("Stop requested! Terminating gracefully...")
  169. event.button.disabled = True
  170. self.stop_event.set()
  171. # ==========================================
  172. # Pipeline Execution
  173. # ==========================================
  174. def _execute_pipeline(self) -> None:
  175. self._toggle_config_inputs(disabled=True)
  176. self.query_one("#btn_start").display = False
  177. self.query_one("#active_controls").display = True
  178. self._set_pipeline_status("[bold green]PIPELINE RUNNING[/bold green]")
  179. self.stop_event.clear()
  180. self.pause_event.clear()
  181. btn_pause = self.query_one("#btn_pause", Button)
  182. btn_pause.label = "Pause"
  183. btn_pause.variant = "warning"
  184. self.run_background_pipeline(self.current_config.copy())
  185. @work(exclusive=True, thread=True)
  186. def run_background_pipeline(self, config: Dict[str, Any]) -> None:
  187. scenario_id = config["scenario"]
  188. steps = SCENARIOS[scenario_id]["tasks"]
  189. scenario_label = SCENARIOS[scenario_id]["label"]
  190. root_tracker = ProgressTracker(self, level=0)
  191. pipeline_state: Dict[str, Any] = {
  192. "events": {"stop": self.stop_event, "pause": self.pause_event}
  193. }
  194. try:
  195. self.ui_logger.info(f"Initiating Scenario: {scenario_label}")
  196. total_tasks = len(steps)
  197. root_tracker.reset()
  198. root_tracker.set_title("Pipeline Status")
  199. root_tracker.update(total=total_tasks, advance=0)
  200. for i, task_id in enumerate(steps):
  201. if self.stop_event.is_set():
  202. break
  203. task_name = PIPELINE_TASKS[task_id]["task_name"]
  204. task_func = PIPELINE_TASKS[task_id]["task_func"]
  205. task_logger = self.ui_logger.get_task_logger(task_name)
  206. self.ui_logger.info(f"Starting Task {i + 1}/{total_tasks}: {task_name}")
  207. task_tracker = root_tracker.get_sub_tracker(task_name)
  208. task_func(task_tracker, task_logger, config, pipeline_state)
  209. if not self.stop_event.is_set():
  210. task_logger.info("Task completed.")
  211. root_tracker.update(advance=1)
  212. if self.stop_event.is_set():
  213. self.call_from_thread(
  214. self._set_pipeline_status, "[bold red]PIPELINE STOPPED[/bold red]"
  215. )
  216. self.ui_logger.error("Pipeline stopped by user.")
  217. root_tracker.reset()
  218. else:
  219. self.call_from_thread(
  220. self._set_pipeline_status,
  221. "[bold blue]PIPELINE COMPLETE[/bold blue]",
  222. )
  223. self.ui_logger.info("All tasks finished successfully!")
  224. except InterruptedError as e:
  225. self.ui_logger.error(f"STOPPED: {str(e)}")
  226. self.call_from_thread(
  227. self._set_pipeline_status, "[bold red]PIPELINE STOPPED[/bold red]"
  228. )
  229. root_tracker.reset()
  230. except Exception as e:
  231. self.ui_logger.file_logger.error(
  232. f"Pipeline failed: {str(e)}", exc_info=True
  233. )
  234. self.ui_logger.error(f"ERROR: {str(e)}")
  235. self.call_from_thread(
  236. self._set_pipeline_status, "[bold red]PIPELINE FAILED[/bold red]"
  237. )
  238. finally:
  239. pipeline_state.clear()
  240. gc.collect()
  241. self.call_from_thread(self._reset_ui)
  242. def _reset_ui(self) -> None:
  243. self.query_one("#btn_stop", Button).disabled = False
  244. self.query_one("#btn_start").display = True
  245. self.query_one("#active_controls").display = False
  246. self._toggle_config_inputs(disabled=False)
  247. if __name__ == "__main__":
  248. app = CNNHarnessApp()
  249. app.run()