"""Dataset management commands.""" import click from tabulate import tabulate @click.group("datasets") def datasets_group() -> None: """Manage evaluation datasets.""" @datasets_group.command("sync") @click.argument("repo") @click.option("--force", "-f", is_flag=True, help="Force re-download") @click.pass_context def sync_dataset(ctx: click.Context, repo: str, force: bool) -> None: """Sync a dataset from HuggingFace.""" cfg = ctx.obj["config"] if cfg.is_remote: from solar_eval.cli.client import EvalClient client = EvalClient(cfg.remote_url, cfg.timeout) click.secho(f"Syncing via server: {repo}...", fg="blue") result = client.post("/api/datasets/sync", json={"repo": repo, "force": force}) if result.get("error"): click.secho(f"Sync failed: {result['error']}", fg="red") raise SystemExit(1) click.secho(f"Synced {len(result.get('synced_files', []))} file(s)", fg="green") else: from solar_eval.core.dataset_loader import DatasetLoader click.secho(f"Syncing locally: {repo}...", fg="blue") loader = DatasetLoader() files = loader.list_files(repo) for f in files: click.echo(f" Downloading: {f}") loader.download_file(repo, f, force=force) click.secho(f"Synced {len(files)} file(s)", fg="green") @datasets_group.command("list") @click.pass_context def list_datasets(ctx: click.Context) -> None: """List available datasets.""" cfg = ctx.obj["config"] if cfg.is_remote: from solar_eval.cli.client import EvalClient client = EvalClient(cfg.remote_url, cfg.timeout) datasets = client.get("/api/datasets") else: # List datasets from project configs from solar_eval.core.project_loader import load_all_project_configs configs = load_all_project_configs(cfg.projects_dir, cfg.config_dirs) seen = set() datasets = [] for c in configs: repo = c.get("dataset", {}).get("repo", "") for t in c.get("tasks", []): path = t.get("dataset_path", "") key = f"{repo}:{path}" if key not in seen: seen.add(key) datasets.append({"repo": repo, "path": path, "cached": False}) if not datasets: click.secho("No datasets found.", fg="yellow") return rows = [[d.get("repo", ""), d.get("path", ""), d.get("num_samples", "-"), "Yes" if d.get("cached") else "No"] for d in datasets] click.secho(tabulate(rows, headers=["Repository", "Path", "Samples", "Cached"], tablefmt="simple"), fg="blue")