|
| 1 | +from __future__ import annotations |
| 2 | + |
| 3 | +import os |
| 4 | +import sys |
| 5 | +import threading |
| 6 | +import time |
| 7 | +from dataclasses import dataclass, field |
| 8 | +from typing import Iterable |
| 9 | + |
| 10 | +import click |
| 11 | + |
| 12 | +from ifixai.types import InspectionCategory |
| 13 | + |
| 14 | + |
| 15 | +_SPINNER_FRAMES = "⠋⠙⠹⠸⠼⠴⠦⠧⠇⠏" |
| 16 | + |
| 17 | + |
| 18 | +class _Spinner: |
| 19 | + def __init__(self, message: str) -> None: |
| 20 | + self._message = message |
| 21 | + self._stop_event = threading.Event() |
| 22 | + self._thread: threading.Thread | None = None |
| 23 | + self._frame = 0 |
| 24 | + self._start_time = 0.0 |
| 25 | + |
| 26 | + def start(self) -> None: |
| 27 | + if self._thread is not None: |
| 28 | + return |
| 29 | + self._start_time = time.monotonic() |
| 30 | + sys.stdout.write(self._render() + "\n") |
| 31 | + sys.stdout.flush() |
| 32 | + self._thread = threading.Thread(target=self._run, daemon=True) |
| 33 | + self._thread.start() |
| 34 | + |
| 35 | + def stop(self) -> None: |
| 36 | + if self._thread is None: |
| 37 | + return |
| 38 | + self._stop_event.set() |
| 39 | + self._thread.join(timeout=0.5) |
| 40 | + self._thread = None |
| 41 | + |
| 42 | + def _render(self) -> str: |
| 43 | + glyph = _SPINNER_FRAMES[self._frame % len(_SPINNER_FRAMES)] |
| 44 | + elapsed = int(time.monotonic() - self._start_time) if self._start_time else 0 |
| 45 | + suffix = f" ({elapsed}s)" if elapsed >= 3 else "" |
| 46 | + return _truecolor(f" {glyph} {self._message}{suffix}", _DIM_RGB) |
| 47 | + |
| 48 | + def _run(self) -> None: |
| 49 | + while not self._stop_event.wait(0.12): |
| 50 | + self._frame += 1 |
| 51 | + sys.stdout.write("\033[1F\033[2K" + self._render() + "\n") |
| 52 | + sys.stdout.flush() |
| 53 | + |
| 54 | + |
| 55 | +_LOGO_LINES: tuple[str, ...] = ( |
| 56 | + "██ ███████ ██ ██ ██ █████ ██", |
| 57 | + "██ ██ ██ ██ ██ ██ ██ ██", |
| 58 | + "██ █████ ██ ███ ███████ ██", |
| 59 | + "██ ██ ██ ██ ██ ██ ██ ██", |
| 60 | + "██ ██ ██ ██ ██ ██ ██ ██", |
| 61 | +) |
| 62 | + |
| 63 | +_ACCENT_RGB = (232, 99, 42) |
| 64 | +_DIM_RGB = (110, 110, 117) |
| 65 | + |
| 66 | +_CATEGORY_COLORS: dict[InspectionCategory, tuple[int, int, int]] = { |
| 67 | + InspectionCategory.FABRICATION: (255, 139, 92), |
| 68 | + InspectionCategory.MANIPULATION: (255, 99, 99), |
| 69 | + InspectionCategory.DECEPTION: (167, 139, 250), |
| 70 | + InspectionCategory.UNPREDICTABILITY: (251, 191, 36), |
| 71 | + InspectionCategory.OPACITY: (96, 165, 250), |
| 72 | +} |
| 73 | + |
| 74 | +_CATEGORY_ORDER: tuple[InspectionCategory, ...] = ( |
| 75 | + InspectionCategory.FABRICATION, |
| 76 | + InspectionCategory.MANIPULATION, |
| 77 | + InspectionCategory.DECEPTION, |
| 78 | + InspectionCategory.UNPREDICTABILITY, |
| 79 | + InspectionCategory.OPACITY, |
| 80 | +) |
| 81 | + |
| 82 | +_BAR_WIDTH = 26 |
| 83 | +_BLOCK_FULL = "█" |
| 84 | +_BLOCK_EMPTY = "·" |
| 85 | + |
| 86 | + |
| 87 | +def supports_color(stream=None) -> bool: |
| 88 | + s = stream or sys.stdout |
| 89 | + if os.environ.get("NO_COLOR"): |
| 90 | + return False |
| 91 | + if not hasattr(s, "isatty"): |
| 92 | + return False |
| 93 | + return bool(s.isatty()) |
| 94 | + |
| 95 | + |
| 96 | +def _truecolor(text: str, rgb: tuple[int, int, int], bold: bool = False) -> str: |
| 97 | + if not supports_color(): |
| 98 | + return text |
| 99 | + r, g, b = rgb |
| 100 | + prefix = f"\033[38;2;{r};{g};{b}m" |
| 101 | + if bold: |
| 102 | + prefix = "\033[1m" + prefix |
| 103 | + return f"{prefix}{text}\033[0m" |
| 104 | + |
| 105 | + |
| 106 | +def print_startup_banner(version: str, *, quiet: bool = False) -> None: |
| 107 | + if quiet or not supports_color(): |
| 108 | + return |
| 109 | + click.echo() |
| 110 | + for line in _LOGO_LINES: |
| 111 | + click.echo(" " + _truecolor(line, _ACCENT_RGB, bold=True)) |
| 112 | + click.echo() |
| 113 | + click.echo(_truecolor(f" ™ · v{version} · powered by iMe", _DIM_RGB)) |
| 114 | + click.echo() |
| 115 | + |
| 116 | + |
| 117 | +@dataclass |
| 118 | +class _CategoryRow: |
| 119 | + category: InspectionCategory |
| 120 | + total: int = 0 |
| 121 | + done: int = 0 |
| 122 | + failed: int = 0 |
| 123 | + |
| 124 | + |
| 125 | +@dataclass |
| 126 | +class CategoryProgress: |
| 127 | + rows: dict[InspectionCategory, _CategoryRow] = field(default_factory=dict) |
| 128 | + _printed_lines: int = 0 |
| 129 | + _started: bool = False |
| 130 | + _interactive: bool = False |
| 131 | + _spinner: _Spinner | None = None |
| 132 | + |
| 133 | + @classmethod |
| 134 | + def from_totals(cls, totals: dict[InspectionCategory, int]) -> "CategoryProgress": |
| 135 | + rows = {cat: _CategoryRow(category=cat, total=totals.get(cat, 0)) for cat in _CATEGORY_ORDER} |
| 136 | + return cls(rows=rows) |
| 137 | + |
| 138 | + def start(self) -> None: |
| 139 | + if self._started: |
| 140 | + return |
| 141 | + self._started = True |
| 142 | + self._interactive = supports_color() |
| 143 | + if not self._interactive: |
| 144 | + return |
| 145 | + total = sum(r.total for r in self.rows.values()) |
| 146 | + if total == 0: |
| 147 | + return |
| 148 | + self._spinner = _Spinner(f"Running {total} tests") |
| 149 | + self._spinner.start() |
| 150 | + self._printed_lines = 1 |
| 151 | + |
| 152 | + def record(self, category: InspectionCategory, passing: bool) -> None: |
| 153 | + row = self.rows.get(category) |
| 154 | + if row is None: |
| 155 | + return |
| 156 | + if row.total == 0: |
| 157 | + row.total = max(1, row.total) |
| 158 | + row.done += 1 |
| 159 | + if not passing: |
| 160 | + row.failed += 1 |
| 161 | + self._redraw() |
| 162 | + |
| 163 | + def finalize(self) -> None: |
| 164 | + if self._spinner is not None: |
| 165 | + self._spinner.stop() |
| 166 | + self._spinner = None |
| 167 | + self._redraw(final=True) |
| 168 | + if self._interactive and self._printed_lines: |
| 169 | + click.echo() |
| 170 | + |
| 171 | + def _redraw(self, *, final: bool = False) -> None: |
| 172 | + if not self._started: |
| 173 | + return |
| 174 | + rendered = self._render() |
| 175 | + if not rendered: |
| 176 | + return |
| 177 | + if self._spinner is not None: |
| 178 | + self._spinner.stop() |
| 179 | + self._spinner = None |
| 180 | + if self._interactive and self._printed_lines: |
| 181 | + sys.stdout.write(f"\033[{self._printed_lines}F") |
| 182 | + for line in rendered: |
| 183 | + sys.stdout.write("\033[2K") |
| 184 | + sys.stdout.write(line + "\n") |
| 185 | + sys.stdout.flush() |
| 186 | + self._printed_lines = len(rendered) |
| 187 | + elif self._interactive and not self._printed_lines: |
| 188 | + for line in rendered: |
| 189 | + sys.stdout.write(line + "\n") |
| 190 | + sys.stdout.flush() |
| 191 | + self._printed_lines = len(rendered) |
| 192 | + elif not self._interactive and final: |
| 193 | + for line in rendered: |
| 194 | + click.echo(line) |
| 195 | + |
| 196 | + def _render(self) -> list[str]: |
| 197 | + out: list[str] = [] |
| 198 | + for cat in _CATEGORY_ORDER: |
| 199 | + row = self.rows[cat] |
| 200 | + if row.total == 0: |
| 201 | + continue |
| 202 | + label = cat.value.upper().ljust(16) |
| 203 | + ratio = row.done / row.total if row.total else 0.0 |
| 204 | + filled = int(round(ratio * _BAR_WIDTH)) |
| 205 | + bar = _BLOCK_FULL * filled + _BLOCK_EMPTY * (_BAR_WIDTH - filled) |
| 206 | + colored_bar = _truecolor(bar, _CATEGORY_COLORS[cat], bold=True) |
| 207 | + count = f"{row.done}/{row.total}".rjust(7) |
| 208 | + tail = "" |
| 209 | + if row.done >= row.total: |
| 210 | + if row.failed == 0: |
| 211 | + tail = " " + _truecolor("✓", (74, 222, 128)) |
| 212 | + else: |
| 213 | + tail = " " + _truecolor(f"✗ {row.failed} failed", (239, 68, 68)) |
| 214 | + out.append(f" {label} {colored_bar} {count}{tail}") |
| 215 | + return out |
| 216 | + |
| 217 | + |
| 218 | +def category_totals_from_specs(specs: Iterable[object]) -> dict[InspectionCategory, int]: |
| 219 | + counts: dict[InspectionCategory, int] = {cat: 0 for cat in _CATEGORY_ORDER} |
| 220 | + for spec in specs: |
| 221 | + cat = getattr(spec, "category", None) |
| 222 | + if isinstance(cat, InspectionCategory): |
| 223 | + counts[cat] = counts.get(cat, 0) + 1 |
| 224 | + return counts |
0 commit comments