136 lines
3.6 KiB
Python
136 lines
3.6 KiB
Python
"""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()
|