File size: 5,203 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
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
"""Versioned prompt models for local prompt registry."""

from pathlib import Path

import yaml
from pydantic import Field

from solar_eval.models.base import CustomBaseModel


class PromptVersionEntry(CustomBaseModel):
    """A single versioned prompt entry in the registry."""

    version: int | str
    file: str = ""  # backward compat (legacy per-task path)
    prompt: str = ""  # prompt library reference (e.g. "R9ci_json_ctx")
    pipeline: str = ""  # pipeline template reference (e.g. "4step_default")
    model: str = "solar-pro2"
    temperature: float = 0.0
    max_tokens: int = 8000
    reasoning_effort: str | None = None
    memo: str = ""
    created_at: str = ""
    run_ref: str | list[str] | None = None


class PromptRegistry(CustomBaseModel):
    """Registry of prompt versions for a single task."""

    task: str
    next_version: int | None = None  # deprecated, kept for backward compat
    versions: list[PromptVersionEntry] = Field(default_factory=list)

    def get_version(self, version: int | str) -> PromptVersionEntry | None:
        """Get a specific version entry by int or string name."""
        # Try exact match first
        for v in self.versions:
            if v.version == version:
                return v
        # Try string/int coercion (e.g. "1" matches 1)
        for v in self.versions:
            if str(v.version) == str(version):
                return v
        return None

    def get_best_version(self) -> PromptVersionEntry | None:
        """Get the version with the highest score (requires run_ref to be set)."""
        scored = [v for v in self.versions if v.run_ref]
        return scored[-1] if scored else None


def load_prompt_messages(path: Path) -> list[dict[str, str]]:
    """Load prompt messages from a YAML file, TXT file, or TXT directory.

    Supported formats:
      - Directory with system.txt (+ optional user.txt)
      - YAML with messages key
      - Plain .txt file (treated as system prompt)

    Returns list of {role, content} dicts.
    """
    path = Path(path)
    if path.is_dir():
        return _load_prompt_messages_from_dir(path)

    if path.suffix in (".txt",):
        content = path.read_text().strip()
        return [{"role": "system", "content": content}]

    data = yaml.safe_load(path.read_text())
    if not isinstance(data, dict) or "messages" not in data:
        raise ValueError(f"Prompt file must contain 'messages' key: {path}")

    return data["messages"]


def load_step_prompts(path: Path) -> dict[str, str]:
    """Load step prompts from YAML file or TXT directory.

    Supported formats:
      - Directory with step*.txt files
      - YAML with step_prompts key

    Returns dict of {step_key: prompt_text}.
    """
    path = Path(path)
    if path.is_dir():
        return _load_step_prompts_from_dir(path)

    data = yaml.safe_load(path.read_text())
    if not isinstance(data, dict) or "step_prompts" not in data:
        raise ValueError(f"Prompt file must contain 'step_prompts' key: {path}")
    return data["step_prompts"]


def detect_prompt_format(path: Path) -> str:
    """Return 'multi_step' or 'single_step' for a prompt file or directory."""
    path = Path(path)
    if path.is_dir():
        return _detect_prompt_format_dir(path)

    data = yaml.safe_load(path.read_text())
    if isinstance(data, dict) and "step_prompts" in data:
        return "multi_step"
    return "single_step"


# -- Directory-based TXT loading helpers --


def _load_step_prompts_from_dir(dir_path: Path) -> dict[str, str]:
    """Load step prompts from a directory of .txt files.

    Reads all .txt files except system.txt/user.txt (reserved for single_step).
    File stems become step keys: step1.txt -> "step1", proofread.txt -> "proofread".
    """
    result = {}
    txt_files = sorted(dir_path.glob("*.txt"))
    if not txt_files:
        raise ValueError(f"No .txt files found in prompt directory: {dir_path}")
    for f in txt_files:
        key = f.stem
        if key in ("system", "user"):
            continue
        result[key] = f.read_text().strip()
    if not result:
        raise ValueError(f"No step prompt .txt files found in {dir_path}")
    return result


def _load_prompt_messages_from_dir(dir_path: Path) -> list[dict[str, str]]:
    """Load prompt messages from a directory with system.txt and optional user.txt."""
    system_file = dir_path / "system.txt"
    if not system_file.exists():
        raise ValueError(f"system.txt not found in prompt directory: {dir_path}")
    messages = [{"role": "system", "content": system_file.read_text().strip()}]
    user_file = dir_path / "user.txt"
    if user_file.exists():
        messages.append({"role": "user", "content": user_file.read_text().strip()})
    return messages


def _detect_prompt_format_dir(dir_path: Path) -> str:
    """Detect format from a prompt directory."""
    step_files = list(dir_path.glob("step*.txt"))
    if step_files:
        return "multi_step"
    if (dir_path / "system.txt").exists():
        return "single_step"
    raise ValueError(
        f"Cannot detect prompt format in {dir_path}: "
        "expected step*.txt files (multi_step) or system.txt (single_step)"
    )