cli-modelarium 0.1.3__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (34) hide show
  1. cli_modelarium/__init__.py +6 -0
  2. cli_modelarium/__main__.py +8 -0
  3. cli_modelarium/assertions.py +596 -0
  4. cli_modelarium/banner.py +96 -0
  5. cli_modelarium/batch.py +425 -0
  6. cli_modelarium/cli.py +2577 -0
  7. cli_modelarium/exceptions.py +88 -0
  8. cli_modelarium/hallucination.py +384 -0
  9. cli_modelarium/io_safety.py +112 -0
  10. cli_modelarium/judging.py +469 -0
  11. cli_modelarium/models_registry.py +138 -0
  12. cli_modelarium/output_formatters.py +1108 -0
  13. cli_modelarium/pricing.py +199 -0
  14. cli_modelarium/providers/__init__.py +7 -0
  15. cli_modelarium/providers/_utils.py +26 -0
  16. cli_modelarium/providers/anthropic_provider.py +148 -0
  17. cli_modelarium/providers/base.py +87 -0
  18. cli_modelarium/providers/deepseek_provider.py +15 -0
  19. cli_modelarium/providers/google_provider.py +135 -0
  20. cli_modelarium/providers/groq_provider.py +15 -0
  21. cli_modelarium/providers/local_provider.py +94 -0
  22. cli_modelarium/providers/mistral_provider.py +172 -0
  23. cli_modelarium/providers/openai_provider.py +163 -0
  24. cli_modelarium/providers/openrouter_provider.py +33 -0
  25. cli_modelarium/providers/xai_provider.py +15 -0
  26. cli_modelarium/run_statistics.py +1202 -0
  27. cli_modelarium/security.py +202 -0
  28. cli_modelarium/streaming.py +416 -0
  29. cli_modelarium-0.1.3.dist-info/METADATA +764 -0
  30. cli_modelarium-0.1.3.dist-info/RECORD +34 -0
  31. cli_modelarium-0.1.3.dist-info/WHEEL +4 -0
  32. cli_modelarium-0.1.3.dist-info/entry_points.txt +2 -0
  33. cli_modelarium-0.1.3.dist-info/licenses/LICENSE +201 -0
  34. cli_modelarium-0.1.3.dist-info/licenses/NOTICE +102 -0
@@ -0,0 +1,425 @@
1
+ """Multi-prompt batch mode.
2
+
3
+ Three concerns:
4
+
5
+ 1. Parsing & validation - load_batch_file() reads .txt or .json into
6
+ BatchPrompt dataclasses with size-limit and structure checks.
7
+ 2. Cost-and-volume guards - check_batch_size_limits() rejects oversized
8
+ batches before any API call is made.
9
+ 3. Orchestration - run_batch() builds StreamStates for every
10
+ (prompt x system x model x temperature) tuple and runs them via the
11
+ same `_call_with_retry` helper as the compare command, but with a
12
+ Rich Progress bar instead of per-token Live panels.
13
+
14
+ Per-prompt system-prompt override (when a BatchPrompt has a `system` field)
15
+ takes precedence over command-line system prompts for that specific prompt
16
+ only - other prompts still use the command-line list.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import asyncio
22
+ import json
23
+ from collections.abc import Callable
24
+ from dataclasses import dataclass, field
25
+ from pathlib import Path
26
+ from typing import Any
27
+
28
+ from rich.console import Console
29
+ from rich.progress import (
30
+ BarColumn,
31
+ MofNCompleteColumn,
32
+ Progress,
33
+ TextColumn,
34
+ TimeRemainingColumn,
35
+ )
36
+
37
+ from cli_modelarium.exceptions import (
38
+ BatchSizeError,
39
+ BatchValidationError,
40
+ ModelariumError,
41
+ )
42
+ from cli_modelarium.io_safety import BATCH_INPUT_MAX_BYTES, safe_input_path
43
+ from cli_modelarium.models_registry import get_provider_for_model
44
+ from cli_modelarium.pricing import calculate_cost, is_local_model
45
+ from cli_modelarium.providers.base import BaseProvider
46
+ from cli_modelarium.streaming import (
47
+ DEFAULT_MAX_RETRIES,
48
+ StreamState,
49
+ _call_with_retry,
50
+ )
51
+
52
+ # Safety caps. The build prompt sets these; --force-large bypasses both.
53
+ MAX_PROMPTS_PER_BATCH = 1000
54
+ MAX_TOTAL_CALLS = 10_000
55
+
56
+ # Cost estimation defaults. Assumed input/output tokens per call when we
57
+ # don't know the real shape yet. Deliberately on the high side - we'd rather
58
+ # refuse a borderline run than burn through somebody's quota.
59
+ ESTIMATE_INPUT_TOKENS = 500
60
+ ESTIMATE_OUTPUT_TOKENS = 500
61
+
62
+
63
+ @dataclass
64
+ class BatchPrompt:
65
+ """One row of a batch input file."""
66
+
67
+ id: str
68
+ prompt: str
69
+ system: str | None = None
70
+ # Raw assertion dicts as-loaded; executed downstream by assertions.run_assertions.
71
+ assertions: list[dict[str, Any]] = field(default_factory=list)
72
+
73
+
74
+ # ===== File parsing =====
75
+
76
+
77
+ def load_batch_file(file_path: str) -> list[BatchPrompt]:
78
+ """Load a batch input file. Format is auto-detected from extension.
79
+
80
+ Supported:
81
+ .txt - one prompt per non-blank, non-comment line
82
+ .json - top-level array of {"prompt": "...", "id"?: "...", ...}
83
+
84
+ Files larger than `BATCH_INPUT_MAX_BYTES` are rejected by
85
+ `safe_input_path` before any parsing happens.
86
+
87
+ Raises:
88
+ FileNotFoundError, ValueError: from safe_input_path.
89
+ BatchValidationError: malformed content or unknown extension.
90
+ json.JSONDecodeError: malformed JSON.
91
+ """
92
+ path = safe_input_path(file_path, max_size_bytes=BATCH_INPUT_MAX_BYTES)
93
+ suffix = path.suffix.lower()
94
+ if suffix == ".txt":
95
+ return _parse_txt(path)
96
+ if suffix == ".json":
97
+ return _parse_json(path)
98
+ raise BatchValidationError(
99
+ f"Cannot detect batch file format from extension {suffix!r}.\n"
100
+ f" Supported: .txt (one prompt per line), .json (array of objects).\n"
101
+ f" At: {path}"
102
+ )
103
+
104
+
105
+ def _parse_txt(path: Path) -> list[BatchPrompt]:
106
+ """One prompt per line. Lines starting with `#` are comments."""
107
+ text = path.read_text(encoding="utf-8-sig")
108
+ prompts: list[BatchPrompt] = []
109
+ for line in text.splitlines():
110
+ stripped = line.strip()
111
+ if not stripped:
112
+ continue
113
+ if stripped.startswith("#"):
114
+ # Comment line - inline `#` mid-line is NOT a comment marker
115
+ # (would conflict with model output that legitimately contains `#`).
116
+ continue
117
+ prompts.append(BatchPrompt(id=f"p{len(prompts) + 1}", prompt=stripped))
118
+ return prompts
119
+
120
+
121
+ def _parse_json(path: Path) -> list[BatchPrompt]:
122
+ """Top-level JSON array of prompt objects."""
123
+ raw = path.read_text(encoding="utf-8-sig")
124
+ data = json.loads(raw) # propagates JSONDecodeError verbatim
125
+ if not isinstance(data, list):
126
+ raise BatchValidationError(
127
+ f"Batch JSON file must be an array at top level, "
128
+ f"got {type(data).__name__!r}. At: {path}"
129
+ )
130
+
131
+ seen_ids: set[str] = set()
132
+ prompts: list[BatchPrompt] = []
133
+ for i, item in enumerate(data):
134
+ if not isinstance(item, dict):
135
+ raise BatchValidationError(
136
+ f"Batch element #{i} is not an object (got {type(item).__name__!r}). At: {path}"
137
+ )
138
+ if "prompt" not in item:
139
+ raise BatchValidationError(
140
+ f"Batch element #{i} is missing required 'prompt' field. At: {path}"
141
+ )
142
+ prompt = item["prompt"]
143
+ if not isinstance(prompt, str):
144
+ raise BatchValidationError(f"Batch element #{i} 'prompt' must be a string. At: {path}")
145
+ prompt_id = item.get("id") or f"p{i + 1}"
146
+ if not isinstance(prompt_id, str):
147
+ raise BatchValidationError(f"Batch element #{i} 'id' must be a string. At: {path}")
148
+ if prompt_id in seen_ids:
149
+ raise BatchValidationError(
150
+ f"Batch contains duplicate prompt id {prompt_id!r}. At: {path}"
151
+ )
152
+ seen_ids.add(prompt_id)
153
+
154
+ system = item.get("system")
155
+ if system is not None and not isinstance(system, str):
156
+ raise BatchValidationError(
157
+ f"Batch element {prompt_id!r} 'system' must be a string. At: {path}"
158
+ )
159
+
160
+ assertions = item.get("assertions", [])
161
+ if not isinstance(assertions, list):
162
+ raise BatchValidationError(
163
+ f"Batch element {prompt_id!r} 'assertions' must be a list. At: {path}"
164
+ )
165
+
166
+ prompts.append(
167
+ BatchPrompt(
168
+ id=prompt_id,
169
+ prompt=prompt,
170
+ system=system,
171
+ assertions=list(assertions),
172
+ )
173
+ )
174
+ return prompts
175
+
176
+
177
+ # ===== Size validation =====
178
+
179
+
180
+ def check_batch_size_limits(
181
+ prompts: list[BatchPrompt],
182
+ models: list[str],
183
+ temperatures: list[float],
184
+ command_system_prompts: list[str | None],
185
+ *,
186
+ force_large: bool = False,
187
+ ) -> int:
188
+ """Confirm the batch fits within the safety caps.
189
+
190
+ Returns the total task count (useful for callers to display).
191
+
192
+ Raises BatchSizeError when limits are exceeded and `force_large` is False.
193
+ """
194
+ if len(prompts) > MAX_PROMPTS_PER_BATCH and not force_large:
195
+ raise BatchSizeError(
196
+ f"Too many prompts: {len(prompts)} (max {MAX_PROMPTS_PER_BATCH}).\n"
197
+ f" Pass --force-large to bypass this safety cap."
198
+ )
199
+
200
+ total = _count_total_calls(prompts, models, temperatures, command_system_prompts)
201
+ if total > MAX_TOTAL_CALLS and not force_large:
202
+ raise BatchSizeError(
203
+ f"Too many total API calls: {total} (max {MAX_TOTAL_CALLS}).\n"
204
+ f" Composition: {len(prompts)} prompts x {len(models)} models "
205
+ f"x {len(temperatures)} temperatures (some prompts add their own system "
206
+ f"prompts on top).\n"
207
+ f" Pass --force-large to bypass this safety cap, or reduce one "
208
+ f"of the dimensions."
209
+ )
210
+ return total
211
+
212
+
213
+ def estimate_batch_cost(
214
+ prompts: list[BatchPrompt],
215
+ models: list[str],
216
+ temperatures: list[float],
217
+ command_system_prompts: list[str | None],
218
+ ) -> float:
219
+ """Return an upper-bound USD cost estimate for the batch.
220
+
221
+ Assumes `ESTIMATE_INPUT_TOKENS` + `ESTIMATE_OUTPUT_TOKENS` per call.
222
+ Unknown models are treated as $0 (we silently skip them in the estimate
223
+ rather than failing - the actual run will fail more loudly when it
224
+ reaches that model).
225
+ """
226
+ total = 0.0
227
+ for bp in prompts:
228
+ effective_sps = [bp.system] if bp.system else command_system_prompts
229
+ for _sp in effective_sps:
230
+ for model in models:
231
+ for _temp in temperatures:
232
+ if is_local_model(model):
233
+ continue
234
+ try:
235
+ total += calculate_cost(
236
+ model,
237
+ input_tokens=ESTIMATE_INPUT_TOKENS,
238
+ output_tokens=ESTIMATE_OUTPUT_TOKENS,
239
+ )
240
+ except ModelariumError:
241
+ # Unknown model: skip in the estimate; real call surfaces it.
242
+ pass
243
+ return total
244
+
245
+
246
+ def estimate_compare_cost(
247
+ models: list[str],
248
+ temperatures: list[float],
249
+ system_prompts: list[str | None],
250
+ ) -> float:
251
+ """Upper-bound USD cost estimate for a compare run.
252
+
253
+ Compare runs 1 prompt x M models x T temperatures x S system_prompts.
254
+ Uses ESTIMATE_INPUT_TOKENS and ESTIMATE_OUTPUT_TOKENS as the per-call
255
+ upper bound. Local models contribute $0. Unknown models are silently
256
+ skipped (real call will fail at runtime if model truly invalid).
257
+
258
+ Does NOT include judge cost (judge output length unknown until run).
259
+ """
260
+ total = 0.0
261
+ for _sp in system_prompts:
262
+ for model in models:
263
+ if is_local_model(model):
264
+ continue
265
+ for _temp in temperatures:
266
+ try:
267
+ total += calculate_cost(
268
+ model,
269
+ input_tokens=ESTIMATE_INPUT_TOKENS,
270
+ output_tokens=ESTIMATE_OUTPUT_TOKENS,
271
+ )
272
+ except ModelariumError:
273
+ pass
274
+ return total
275
+
276
+
277
+ def _count_total_calls(
278
+ prompts: list[BatchPrompt],
279
+ models: list[str],
280
+ temperatures: list[float],
281
+ command_system_prompts: list[str | None],
282
+ ) -> int:
283
+ """Count the total task count, accounting for per-prompt system overrides."""
284
+ n_models = max(1, len(models))
285
+ n_temps = max(1, len(temperatures))
286
+ n_command_sps = max(1, len(command_system_prompts))
287
+ total = 0
288
+ for bp in prompts:
289
+ n_sps = 1 if bp.system else n_command_sps
290
+ total += n_sps * n_models * n_temps
291
+ return total
292
+
293
+
294
+ # ===== Orchestration =====
295
+
296
+
297
+ def build_batch_states(
298
+ prompts: list[BatchPrompt],
299
+ models: list[str],
300
+ temperatures: list[float],
301
+ command_system_prompts: list[str | None],
302
+ ) -> list[tuple[StreamState, BatchPrompt]]:
303
+ """Build StreamStates for every (prompt x system x model x temperature) tuple.
304
+
305
+ Each returned tuple pairs the state with its source BatchPrompt so the
306
+ caller can look up the prompt text and assertions later.
307
+
308
+ Per-prompt `system` overrides win for that prompt only; other prompts
309
+ use the command-line `command_system_prompts` list.
310
+ """
311
+ pairs: list[tuple[StreamState, BatchPrompt]] = []
312
+ for bp in prompts:
313
+ effective_sps = [bp.system] if bp.system else command_system_prompts
314
+ for sp in effective_sps:
315
+ for model in models:
316
+ provider_name = get_provider_for_model(model)
317
+ for temperature in temperatures:
318
+ pairs.append(
319
+ (
320
+ StreamState(
321
+ model=model,
322
+ provider_name=provider_name,
323
+ temperature=temperature,
324
+ system_prompt=sp,
325
+ ),
326
+ bp,
327
+ )
328
+ )
329
+ return pairs
330
+
331
+
332
+ async def run_batch(
333
+ *,
334
+ pairs: list[tuple[StreamState, BatchPrompt]],
335
+ provider_factory: Callable[[str], BaseProvider],
336
+ console: Console,
337
+ concurrency: int,
338
+ max_retries: int = DEFAULT_MAX_RETRIES,
339
+ show_progress: bool = True,
340
+ sleep: Callable[[float], asyncio.Future[None]] = asyncio.sleep,
341
+ ) -> list[tuple[StreamState, BatchPrompt]]:
342
+ """Run every (state, prompt) pair in parallel under per-provider semaphores.
343
+
344
+ Reuses `streaming._call_with_retry` so the 429/529 retry behavior matches
345
+ the compare command exactly. The only difference is the display: a Rich
346
+ Progress bar showing "Completed X/Y" instead of per-task Live panels.
347
+
348
+ Returns the same `pairs` list (states mutated in place) for convenience.
349
+ """
350
+ provider_names = sorted({s.provider_name for s, _ in pairs})
351
+ instances: dict[str, BaseProvider] = {name: provider_factory(name) for name in provider_names}
352
+ semaphores: dict[str, asyncio.Semaphore] = {
353
+ name: asyncio.Semaphore(concurrency) for name in provider_names
354
+ }
355
+
356
+ progress: Progress | None = None
357
+ task_id = None
358
+ if show_progress and pairs:
359
+ progress = Progress(
360
+ TextColumn("[progress.description]{task.description}"),
361
+ BarColumn(),
362
+ MofNCompleteColumn(),
363
+ TextColumn("[progress.percentage]{task.percentage:>3.0f}%"),
364
+ TimeRemainingColumn(),
365
+ console=console,
366
+ )
367
+ task_id = progress.add_task(
368
+ f"Running {len(pairs)} call{'s' if len(pairs) != 1 else ''}",
369
+ total=len(pairs),
370
+ )
371
+
372
+ async def _run_one(state: StreamState, bp: BatchPrompt) -> None:
373
+ provider = instances[state.provider_name]
374
+ async with semaphores[state.provider_name]:
375
+ state.mark_started()
376
+ try:
377
+ result = await _call_with_retry(
378
+ provider=provider,
379
+ state=state,
380
+ prompt=bp.prompt,
381
+ max_retries=max_retries,
382
+ sleep=sleep,
383
+ )
384
+ state.mark_complete(result)
385
+ except ModelariumError as e:
386
+ state.mark_error(str(e))
387
+ except Exception as e: # noqa: BLE001 - become an error row, not a crash
388
+ state.mark_error(f"unexpected: {e}")
389
+ finally:
390
+ if progress is not None and task_id is not None:
391
+ progress.advance(task_id)
392
+
393
+ coro = asyncio.gather(*[_run_one(s, bp) for s, bp in pairs])
394
+
395
+ if progress is not None:
396
+ with progress:
397
+ await coro
398
+ else:
399
+ await coro
400
+
401
+ return pairs
402
+
403
+
404
+ # ===== Filename safety =====
405
+
406
+ # Detect output formats from common extensions.
407
+ _EXTENSION_TO_FORMAT = {
408
+ ".csv": "csv",
409
+ ".json": "json",
410
+ ".md": "markdown",
411
+ ".markdown": "markdown",
412
+ }
413
+
414
+
415
+ def detect_output_format(path: Path) -> str | None:
416
+ """Infer 'csv' / 'json' / 'markdown' from a file extension. None if unknown."""
417
+ return _EXTENSION_TO_FORMAT.get(path.suffix.lower())
418
+
419
+
420
+ def output_overlaps_input(input_path: Path, output_path: Path) -> bool:
421
+ """True if `output_path` would overwrite `input_path` (same resolved file)."""
422
+ try:
423
+ return input_path.resolve() == output_path.resolve()
424
+ except OSError:
425
+ return False