File size: 6,776 Bytes
c34ff1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
"""Smoke tests for the speculators CLI."""

import click
from typer.testing import CliRunner

from speculators.cli import app

runner = CliRunner()


def unstyled_output(result):
    """Return CLI output without terminal styling for stable assertions."""
    return click.unstyle(result.output)


class TestRootApp:
    def test_no_args_shows_help(self):
        result = runner.invoke(app, [])
        assert "Usage" in unstyled_output(result)

    def test_help(self):
        result = runner.invoke(app, ["--help"])
        assert result.exit_code == 0
        assert "Pipeline" in unstyled_output(result)
        assert "Tools" in unstyled_output(result)

    def test_version(self):
        result = runner.invoke(app, ["--version"])
        assert result.exit_code == 0
        assert "speculators version:" in unstyled_output(result)

    def test_pipeline_commands_in_help(self):
        result = runner.invoke(app, ["--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "prepare-data" in output
        assert "stitch-mtp" in output
        assert "generate-offline-data" in output
        assert "regenerate-responses" in output
        assert "train" in output

    def test_tools_commands_in_help(self):
        result = runner.invoke(app, ["--help"])
        assert result.exit_code == 0
        assert "convert" in unstyled_output(result)


class TestConvertCommand:
    def test_help(self):
        result = runner.invoke(app, ["convert", "--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "--verifier" in output
        assert "--algorithm" in output

    def test_algorithm_choices_in_help(self):
        result = runner.invoke(app, ["convert", "--help"])
        assert result.exit_code == 0
        for algo in ("eagle3", "mtp", "dflash"):
            assert algo in unstyled_output(result)

    def test_missing_required_args(self):
        result = runner.invoke(app, ["convert"])
        assert result.exit_code != 0


class TestPrepareDataCommand:
    def test_help(self):
        result = runner.invoke(app, ["prepare-data", "--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "--model" in output
        assert "--data" in output
        assert "--output" in output
        assert "--seq-length" in output

    def test_missing_required_args(self):
        result = runner.invoke(app, ["prepare-data"])
        assert result.exit_code != 0

    def test_allow_empty_output_in_help(self):
        result = runner.invoke(app, ["prepare-data", "--help"])
        assert result.exit_code == 0
        assert "--allow-empty-output" in unstyled_output(result)

    def test_overwrite_in_help(self):
        result = runner.invoke(app, ["prepare-data", "--help"])
        assert result.exit_code == 0
        assert "--overwrite" in unstyled_output(result)

    def test_render_endpoint_in_help(self):
        result = runner.invoke(app, ["prepare-data", "--help"])
        assert result.exit_code == 0
        assert "--render-endpoint" in unstyled_output(result)


class TestStitchCommand:
    def test_help(self):
        result = runner.invoke(app, ["stitch-mtp", "--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "finetuned_checkpoint" in output
        assert "verifier_path" in output

    def test_missing_required_args(self):
        result = runner.invoke(app, ["stitch-mtp"])
        assert result.exit_code != 0


class TestGenerateOfflineDataCommand:
    def test_help(self):
        result = runner.invoke(app, ["generate-offline-data", "--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "--endpoint" in output
        assert "--preprocessed-data" in output
        assert "--concurrency" in output
        assert "--world-size" in output
        assert "--rank" in output

    def test_fail_on_error_in_help(self):
        result = runner.invoke(app, ["generate-offline-data", "--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "--fail-on-error" in output
        assert "--max-retries" in output
        assert "--validate-outputs" in output

    def test_invalid_rank(self):
        result = runner.invoke(
            app, ["generate-offline-data", "--rank", "5", "--world-size", "2"]
        )
        assert result.exit_code != 0

    def test_invalid_concurrency(self):
        result = runner.invoke(app, ["generate-offline-data", "--concurrency", "0"])
        assert result.exit_code != 0


class TestRegenerateResponsesCommand:
    def test_help(self):
        result = runner.invoke(app, ["regenerate-responses", "--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "--endpoint" in output
        assert "--dataset" in output
        assert "--concurrency" in output
        assert "--max-tokens" in output

    def test_invalid_max_retries(self):
        result = runner.invoke(app, ["regenerate-responses", "--max-retries", "-1"])
        assert result.exit_code != 0

    def test_invalid_sampling_params(self):
        result = runner.invoke(
            app, ["regenerate-responses", "--sampling-params", "not-json"]
        )
        assert result.exit_code != 0

    def test_sampling_params_must_be_object(self):
        result = runner.invoke(
            app, ["regenerate-responses", "--sampling-params", "[1,2,3]"]
        )
        assert result.exit_code != 0

    def test_split_only_applies_to_presets(self, tmp_path):
        dataset = tmp_path / "prompts.jsonl"
        dataset.touch()
        result = runner.invoke(
            app,
            [
                "regenerate-responses",
                "--dataset",
                str(dataset),
                "--split",
                "custom",
            ],
        )
        assert result.exit_code != 0
        assert "only apply to dataset presets" in unstyled_output(result)

    def test_invalid_temperature_cycle(self):
        result = runner.invoke(
            app, ["regenerate-responses", "--temperature-cycle", "0.6,notnum"]
        )
        assert result.exit_code != 0


class TestTrainCommand:
    def test_help(self):
        result = runner.invoke(app, ["train", "--help"])
        assert result.exit_code == 0
        output = unstyled_output(result)
        assert "--verifier-name-or-path" in output
        assert "--config" in output
        assert "--speculator-type" in output

    def test_train_appears_in_pipeline_panel(self):
        result = runner.invoke(app, ["--help"])
        assert result.exit_code == 0
        assert "train" in unstyled_output(result)