Files
2026-09-19 12:54:45 +08:00

136 lines
3.6 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Typer CLI equivalent to the REST API."""
import uuid
from typing import List, Optional
import typer
from sqlalchemy.orm import Session
from app.api.router import (
AnalysisService,
DatasetBuildRequest,
ExperimentCreate,
PromptService,
create_adapter,
generate_full_factorial_matrix,
seed_default_templates,
)
from app.db import SessionLocal, init_db
from app.experiments.runner import ExperimentRunner
from app.model_adapters.factory import list_models
from app.models import Sample
app = typer.Typer(help="PromptCR-Lab command-line interface")
def get_db() -> Session:
init_db()
return SessionLocal()
@app.command()
def build_dataset(
repo_path: str,
count: int = typer.Option(12, "--count", "-c"),
languages: Optional[List[str]] = typer.Option(None, "--language", "-l"),
):
"""Build dataset from a Git repository."""
from app.dataset.builder import build_dataset_from_repo
db = get_db()
languages = languages or ["python", "java", "javascript"]
samples = build_dataset_from_repo(db, repo_path, count, languages)
typer.echo(f"Created {len(samples)} samples")
@app.command()
def list_samples():
"""List all samples."""
db = get_db()
samples = db.query(Sample).all()
for s in samples:
typer.echo(f"{s.id} {s.repo} {s.commit_sha} {s.language}")
@app.command()
def create_experiment(
name: str,
models: List[str] = typer.Option(..., "--model", "-m"),
levels: List[str] = typer.Option(..., "--level", "-l"),
sample_ids: List[str] = typer.Option(..., "--sample", "-s"),
repeats: int = typer.Option(3, "--repeats", "-r"),
):
"""Create a full-factorial experiment."""
db = get_db()
seed_default_templates(db)
experiment = generate_full_factorial_matrix(
db,
name=name,
models=models,
levels=levels,
sample_ids=sample_ids,
repeats=repeats,
)
typer.echo(f"Created experiment {experiment.id} with {len(experiment.runs)} runs")
@app.command()
def run_experiment(
experiment_id: str,
concurrency: int = typer.Option(5, "--concurrency", "-c"),
):
"""Run pending experiment units."""
import asyncio
db = get_db()
runner = ExperimentRunner(db, concurrency=concurrency)
summary = asyncio.run(runner.run_experiment(experiment_id=uuid.UUID(experiment_id)))
typer.echo(f"Total: {summary['total']}, Completed: {summary['completed']}")
@app.command()
def smoke(
model: str = typer.Option("deepseek", "--model", "-m"),
level: str = typer.Option("L1", "--level", "-l"),
sample_id: str = typer.Option(..., "--sample", "-s"),
):
"""Run a 1×1×1×1 smoke test against a real model API."""
import asyncio
db = get_db()
seed_default_templates(db)
sample = db.query(Sample).filter_by(id=uuid.UUID(sample_id)).first()
if not sample:
typer.echo("Sample not found", err=True)
raise typer.Exit(1)
prompt_service = PromptService(db)
prompt = prompt_service.render("code_review", None, {"language": sample.language, "diff": sample.diff})
adapter = create_adapter(model)
async def call():
response = await adapter.chat(prompt)
typer.echo(response.text)
asyncio.run(call())
@app.command()
def list_models_cmd():
"""List supported model IDs."""
for model in list_models():
typer.echo(model)
@app.command()
def aggregate(experiment_id: str):
"""Aggregate metrics for an experiment."""
db = get_db()
service = AnalysisService(db)
# AnalysisService converts the id internally and expects a string.
typer.echo(service.aggregate_metrics(experiment_id))
if __name__ == "__main__":
app()