first commit
This commit is contained in:
@@ -0,0 +1,135 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user