| from __future__ import annotations |
|
|
| import re |
| from dataclasses import dataclass |
| from pathlib import Path |
|
|
| from adam.config import ConfigManager |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class ToolFolderDefinition: |
| tool_id: str |
| name: str |
| expected_files: tuple[str, ...] = () |
|
|
|
|
| @dataclass(frozen=True, slots=True) |
| class ToolFolderStatus: |
| tool_id: str |
| name: str |
| path: str |
| exists: bool |
| valid: bool |
| entry_points: tuple[str, ...] |
| message: str |
|
|
|
|
| TOOL_FOLDER_DEFINITIONS = ( |
| ToolFolderDefinition( |
| "dataset_collector", |
| "Dataset Collector", |
| ("collector.py", "dataset_collector.py", "main.py"), |
| ), |
| ToolFolderDefinition( |
| "caption_generator", |
| "Caption Generator", |
| ("caption.py", "caption_generator.py", "main.py"), |
| ), |
| ToolFolderDefinition( |
| "lora_trainer", |
| "LoRA Trainer", |
| ("train.py", "lora_train.py", "main.py", "src/loratrainer/main.py"), |
| ), |
| ToolFolderDefinition( |
| "ddpm_trainer", |
| "DDPM Trainer", |
| ("train.py", "appStableDiffusion.py"), |
| ), |
| ToolFolderDefinition( |
| "flow_trainer", |
| "Flow Matching Trainer", |
| ("flow_matching_app.py", "roblox_action_flow_app.py"), |
| ), |
| ToolFolderDefinition( |
| "preview_generator", |
| "Preview Generator", |
| ("generate.py", "preview.py", "main.py"), |
| ), |
| ) |
|
|
|
|
| class ToolFolderManager: |
| def __init__(self, config: ConfigManager) -> None: |
| self.config = config |
| self.definitions = { |
| definition.tool_id: definition |
| for definition in TOOL_FOLDER_DEFINITIONS |
| } |
|
|
| def paths(self) -> dict[str, str]: |
| stored = self.config.get("tool_folders", {}) |
| return dict(stored) if isinstance(stored, dict) else {} |
|
|
| def get(self, tool_id: str) -> str: |
| return str(self.paths().get(tool_id, "")) |
|
|
| def set(self, tool_id: str, path: str) -> ToolFolderStatus: |
| if tool_id not in self.definitions: |
| raise KeyError(f"Unknown tool folder: {tool_id}") |
| stored = self.paths() |
| stored[tool_id] = path.strip().strip('"') |
| self.config.update({"tool_folders": stored}) |
| return self.scan(tool_id) |
|
|
| def update(self, values: dict[str, str]) -> dict[str, ToolFolderStatus]: |
| stored = self.paths() |
| for tool_id, value in values.items(): |
| if tool_id in self.definitions: |
| stored[tool_id] = value.strip().strip('"') |
| self.config.update({"tool_folders": stored}) |
| return {tool_id: self.scan(tool_id) for tool_id in values} |
|
|
| def scan(self, tool_id: str) -> ToolFolderStatus: |
| definition = self.definitions[tool_id] |
| raw_path = self.get(tool_id) |
| if not raw_path: |
| return ToolFolderStatus( |
| tool_id, |
| definition.name, |
| "", |
| False, |
| False, |
| (), |
| "Not configured", |
| ) |
| folder = Path(raw_path).expanduser() |
| if not folder.is_dir(): |
| return ToolFolderStatus( |
| tool_id, |
| definition.name, |
| str(folder), |
| False, |
| False, |
| (), |
| "Folder not found", |
| ) |
| entry_points = tuple( |
| filename |
| for filename in definition.expected_files |
| if (folder / filename).is_file() |
| ) |
| if entry_points: |
| message = f"Detected 路 {', '.join(entry_points)}" |
| valid = True |
| else: |
| top_level_python = sorted(path.name for path in folder.glob("*.py")) |
| entry_points = tuple(top_level_python[:5]) |
| valid = bool(entry_points) |
| message = ( |
| f"Python project detected 路 review {entry_points[0]}" |
| if entry_points |
| else "Folder found 路 no Python entry point detected" |
| ) |
| return ToolFolderStatus( |
| tool_id, |
| definition.name, |
| str(folder.resolve()), |
| True, |
| valid, |
| entry_points, |
| message, |
| ) |
|
|
| def scan_all(self) -> dict[str, ToolFolderStatus]: |
| return { |
| tool_id: self.scan(tool_id) |
| for tool_id in self.definitions |
| } |
|
|
| def parse_assignments(self, text: str) -> dict[str, str]: |
| """Recognize folder assignments pasted into chat without executing them.""" |
| patterns = { |
| "ddpm_trainer": r"(?im)^\s*DDPM(?:\s+Trainer)?\s*:\s*(.+?)\s*$", |
| "flow_trainer": ( |
| r"(?im)^\s*Flow(?:\s+Matching)?(?:\s+Trainer)?\s*:\s*(.+?)\s*$" |
| ), |
| "lora_trainer": r"(?im)^\s*LoRA(?:\s+Trainer)?\s*:\s*(.+?)\s*$", |
| "dataset_collector": ( |
| r"(?im)^\s*Dataset(?:\s+Collector)?\s*:\s*(.+?)\s*$" |
| ), |
| "caption_generator": ( |
| r"(?im)^\s*Caption(?:\s+Generator)?\s*:\s*(.+?)\s*$" |
| ), |
| "preview_generator": ( |
| r"(?im)^\s*Preview(?:\s+Generator)?\s*:\s*(.+?)\s*$" |
| ), |
| } |
| assignments: dict[str, str] = {} |
| for tool_id, pattern in patterns.items(): |
| match = re.search(pattern, text) |
| if match: |
| assignments[tool_id] = match.group(1).strip().strip('"') |
| return assignments |
|
|