"""Full-factorial experiment matrix generation.""" from itertools import product from typing import Any, Dict, List from uuid import uuid4 from sqlalchemy.orm import Session from app.models import Experiment, ExperimentRun, PromptTemplateVersion, Sample from app.prompts.service import PromptService def generate_full_factorial_matrix( db: Session, name: str, models: List[str], levels: List[str], sample_ids: List[str], repeats: int = 3, sampling_params: Dict[str, Any] = None, strategy_id: str = "code_review", ) -> Experiment: """Generate experiment runs for a full-factorial matrix. Models × Levels × Samples × Repeats. """ sampling_params = sampling_params or {"temperature": 0.7, "max_tokens": 2048} experiment = Experiment( name=name, models=models, levels=levels, sample_ids=sample_ids, repeats=repeats, sampling_params=sampling_params, status="pending", ) db.add(experiment) db.flush() prompt_service = PromptService(db) runs = [] for model_id, level in product(models, levels): # Resolve latest template version for this level versions = prompt_service.list_versions(strategy_id, level) if not versions: raise ValueError(f"No prompt template found for {strategy_id}/{level}") template_version = versions[-1] for sample_id in sample_ids: from uuid import UUID sample_uuid = UUID(sample_id) if isinstance(sample_id, str) else sample_id sample = db.query(Sample).filter_by(id=sample_uuid).first() if not sample: raise ValueError(f"Sample not found: {sample_id}") for repeat_index in range(1, repeats + 1): run = ExperimentRun( run_id=str(uuid4()), experiment_id=experiment.id, sample_id=sample.id, model_id=model_id, template_version_id=template_version.id, repeat_index=repeat_index, status="pending", sampling_params=sampling_params, ) runs.append(run) db.add_all(runs) db.commit() db.refresh(experiment) return experiment def count_pending_runs(db: Session, experiment_id: str) -> int: return ( db.query(ExperimentRun) .filter_by(experiment_id=experiment_id, status="pending") .count() )