File size: 4,303 Bytes
9c84f9d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
"""Prompt version registry manager for local mode."""

import json
import logging
from pathlib import Path
from typing import Any

import yaml

from solar_eval.models.prompt_version import PromptRegistry, PromptVersionEntry

logger = logging.getLogger(__name__)


def resolve_prompt_path(
    entry: PromptVersionEntry,
    prompts_dir: Path,
    task_dir: Path | None = None,
) -> Path:
    """Resolve prompt file path from a version entry.

    Resolution order:
    1. entry.prompt → prompts_dir/_library/{prompt}.yaml
    2. entry.file  → task_dir/{file} (legacy per-task path)
    """
    if entry.prompt:
        return prompts_dir / "_library" / f"{entry.prompt}.yaml"
    if task_dir and entry.file:
        return task_dir / entry.file
    raise ValueError(f"Cannot resolve prompt path: no prompt= or file= set on {entry.version}")


def resolve_pipeline_path(
    entry: PromptVersionEntry,
    pipelines_dir: Path,
) -> Path | None:
    """Resolve pipeline template path, or None if entry has no pipeline= field."""
    if not entry.pipeline:
        return None
    return pipelines_dir / f"{entry.pipeline}.yaml"


class PromptRegistryManager:
    """Manages the prompts.yaml registry file for a single task."""

    REGISTRY_FILE = "prompts.yaml"

    def __init__(self, prompts_task_dir: Path, experiments_dir: Path | None = None) -> None:
        """Initialize with the task-level prompts directory.

        e.g., projects/my-project/prompts/translate/

        Args:
            prompts_task_dir: Legacy per-task prompts directory.
            experiments_dir: New experiments/ directory (fallback when prompts.yaml not found).
        """
        self.dir = Path(prompts_task_dir)
        self.dir.mkdir(parents=True, exist_ok=True)
        self.experiments_dir = Path(experiments_dir) if experiments_dir else None
        self._registry_path = self.dir / self.REGISTRY_FILE

    def load(self) -> PromptRegistry:
        """Load or create the registry.

        Resolution order:
        1. Legacy: prompts/{task}/prompts.yaml
        2. New: experiments/{task}.yaml (if experiments_dir is set)
        3. Empty registry with task name inferred from directory
        """
        # Try legacy location first
        if self._registry_path.exists():
            data = yaml.safe_load(self._registry_path.read_text())
            if isinstance(data, dict):
                return PromptRegistry(**data)
        # Try experiments/ directory
        if self.experiments_dir:
            exp_path = self.experiments_dir / f"{self.dir.name}.yaml"
            if exp_path.exists():
                data = yaml.safe_load(exp_path.read_text())
                if isinstance(data, dict):
                    return PromptRegistry(task=self.dir.name, versions=data.get("versions", []))
        # Infer task name from directory name
        return PromptRegistry(task=self.dir.name)

    def save(self, registry: PromptRegistry) -> None:
        """Save registry to YAML."""
        self._registry_path.write_text(
            yaml.dump(registry.model_dump(), allow_unicode=True, sort_keys=False, default_flow_style=False)
        )

    def update_run_ref(self, version: int | str, run_ref: str) -> None:
        """Link a version to its run artifacts (appends to list)."""
        registry = self.load()
        entry = registry.get_version(version)
        if entry:
            # Migrate legacy str to list
            if isinstance(entry.run_ref, str):
                entry.run_ref = [entry.run_ref]
            elif entry.run_ref is None:
                entry.run_ref = []
            entry.run_ref.append(run_ref)
            self.save(registry)

    def get_evaluation(self, version: int | str, artifacts_dir: Path) -> dict[str, Any] | None:
        """Get evaluation results for a version by reading its latest run's evaluation.json."""
        registry = self.load()
        entry = registry.get_version(version)
        if not entry or not entry.run_ref:
            return None
        refs = entry.run_ref if isinstance(entry.run_ref, list) else [entry.run_ref]
        # Return latest run's evaluation
        eval_file = artifacts_dir / refs[-1] / "evaluation.json"
        if not eval_file.exists():
            return None
        return json.loads(eval_file.read_text())