quantdiff 0.1.0rc1__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 (45) hide show
  1. quantdiff/__init__.py +53 -0
  2. quantdiff/__main__.py +5 -0
  3. quantdiff/_http.py +151 -0
  4. quantdiff/_text.py +13 -0
  5. quantdiff/_version.py +1 -0
  6. quantdiff/api.py +340 -0
  7. quantdiff/backends/__init__.py +28 -0
  8. quantdiff/backends/_common.py +342 -0
  9. quantdiff/backends/base.py +91 -0
  10. quantdiff/backends/llamacpp.py +428 -0
  11. quantdiff/backends/ollama.py +359 -0
  12. quantdiff/backends/openai_compat.py +338 -0
  13. quantdiff/cache.py +240 -0
  14. quantdiff/card.py +1664 -0
  15. quantdiff/cli.py +377 -0
  16. quantdiff/discover.py +488 -0
  17. quantdiff/errors.py +45 -0
  18. quantdiff/metrics/__init__.py +36 -0
  19. quantdiff/metrics/codeexec.py +428 -0
  20. quantdiff/metrics/jsonschema.py +610 -0
  21. quantdiff/metrics/logit.py +214 -0
  22. quantdiff/metrics/tasks.py +114 -0
  23. quantdiff/metrics/textsim.py +66 -0
  24. quantdiff/metrics/toolcheck.py +99 -0
  25. quantdiff/png.py +360 -0
  26. quantdiff/preflight.py +365 -0
  27. quantdiff/progress.py +283 -0
  28. quantdiff/py.typed +0 -0
  29. quantdiff/report.py +780 -0
  30. quantdiff/runner.py +492 -0
  31. quantdiff/spec.py +154 -0
  32. quantdiff/stats.py +226 -0
  33. quantdiff/suites/__init__.py +462 -0
  34. quantdiff/suites/data/chat.jsonl +22 -0
  35. quantdiff/suites/data/code.jsonl +32 -0
  36. quantdiff/suites/data/json.jsonl +34 -0
  37. quantdiff/suites/data/scoring.jsonl +41 -0
  38. quantdiff/suites/data/tools.jsonl +32 -0
  39. quantdiff/types.py +322 -0
  40. quantdiff/verdict.py +1513 -0
  41. quantdiff-0.1.0rc1.dist-info/METADATA +514 -0
  42. quantdiff-0.1.0rc1.dist-info/RECORD +45 -0
  43. quantdiff-0.1.0rc1.dist-info/WHEEL +4 -0
  44. quantdiff-0.1.0rc1.dist-info/entry_points.txt +2 -0
  45. quantdiff-0.1.0rc1.dist-info/licenses/LICENSE +202 -0
quantdiff/runner.py ADDED
@@ -0,0 +1,492 @@
1
+ """Orchestrates a comparison run: connect to everything, run the reference, then each
2
+ candidate in turn.
3
+
4
+ Every server is contacted before any real work starts, so a typo in a model tag fails in
5
+ a second instead of after minutes of reference generation. Candidates then run one at a
6
+ time on purpose: they usually share a GPU, and running them in parallel would distort
7
+ throughput numbers and the servers' own caching. A server error on one request is
8
+ recorded and never discards results that were already measured.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import logging
14
+ from collections.abc import Callable, Sequence
15
+ from dataclasses import dataclass, field, replace
16
+ from datetime import datetime, timezone
17
+
18
+ from quantdiff._version import __version__
19
+ from quantdiff.backends import open_backend
20
+ from quantdiff.backends._common import forced_texts
21
+ from quantdiff.backends.base import Backend
22
+ from quantdiff.cache import ReferenceCache, ReferenceOutputs, cache_key
23
+ from quantdiff.errors import BackendError, CapabilityError
24
+ from quantdiff.metrics import (
25
+ agreement_metrics,
26
+ evaluate_case,
27
+ logit_metrics,
28
+ perf_metrics,
29
+ summarize_tasks,
30
+ )
31
+ from quantdiff.metrics.tasks import SCORED_KINDS
32
+ from quantdiff.preflight import run_preflight
33
+ from quantdiff.suites import suite_digest
34
+ from quantdiff.types import (
35
+ CandidateResult,
36
+ CandidateSpec,
37
+ CaseOutcome,
38
+ ChatResult,
39
+ LogitMetrics,
40
+ Message,
41
+ PreflightFinding,
42
+ ProgressEvent,
43
+ ProgressPhase,
44
+ ReferenceTrace,
45
+ Report,
46
+ RunSettings,
47
+ ScoringPrompt,
48
+ ServerInfo,
49
+ TaskCase,
50
+ TokenStep,
51
+ TopK,
52
+ )
53
+
54
+ logger = logging.getLogger(__name__)
55
+
56
+ ProgressFn = Callable[[ProgressEvent], None]
57
+ BackendFactory = Callable[[CandidateSpec], Backend]
58
+
59
+ # Progress units are weighted by measured cost so the ETA is honest. One unit is one
60
+ # teacher-forced position (a single-token request on a cached prefix). On a consumer GPU a
61
+ # chat case costs about 30 of those, a reference scoring prompt about one per generated
62
+ # token, and the pre-flight long-prompt probe about as much as six cases.
63
+ _CASE_UNITS = 30
64
+ _PREFLIGHT_UNITS = 180
65
+ _WARMUP_UNITS = 10
66
+ _WARMUP_MESSAGES = (Message(role="user", content="Hi"),)
67
+
68
+
69
+ @dataclass(frozen=True, slots=True)
70
+ class RunPlan:
71
+ """Everything a run needs, fully resolved and validated."""
72
+
73
+ reference: CandidateSpec
74
+ candidates: tuple[CandidateSpec, ...]
75
+ cases: tuple[TaskCase, ...]
76
+ scoring: tuple[ScoringPrompt, ...]
77
+ settings: RunSettings
78
+ title: str
79
+ hf_repo: str | None = None
80
+ offline: bool = False
81
+ preflight: bool = True
82
+ context_probe_tokens: int = 6000
83
+ code_timeout_seconds: float = 10.0
84
+
85
+
86
+ @dataclass(slots=True)
87
+ class _Progress:
88
+ """Counts weighted work units and forwards ProgressEvents."""
89
+
90
+ sink: ProgressFn
91
+ total: int
92
+ completed: int = 0
93
+ model: str = ""
94
+
95
+ def step(self, phase: ProgressPhase, units: int = 1, detail: str = "") -> None:
96
+ self.completed = min(self.completed + units, self.total)
97
+ self.sink(ProgressEvent(phase, self.model, self.completed, self.total, detail))
98
+
99
+ def finish(self) -> None:
100
+ self.completed = self.total
101
+ self.model = ""
102
+ self.sink(ProgressEvent("done", "", self.total, self.total))
103
+
104
+
105
+ @dataclass(slots=True)
106
+ class _Answers:
107
+ """Chat answers by case id, plus the cases the server failed to answer."""
108
+
109
+ results: dict[str, ChatResult] = field(default_factory=dict)
110
+ failures: list[str] = field(default_factory=list)
111
+
112
+
113
+ def execute(
114
+ plan: RunPlan,
115
+ *,
116
+ backend_factory: BackendFactory = open_backend,
117
+ cache: ReferenceCache | None = None,
118
+ progress: ProgressFn | None = None,
119
+ ) -> Report:
120
+ """Run the plan and return a Report. Pass cache=None to always rerun the reference.
121
+
122
+ Raises BackendError before doing any work if a server cannot be reached or does not
123
+ serve the requested model.
124
+ """
125
+ tracker = _Progress(sink=progress or _log_progress, total=_total_units(plan))
126
+ backends = _connect(plan, backend_factory, tracker)
127
+ try:
128
+ reference, *candidates = backends
129
+ ref_backend, ref_info = reference
130
+ tracker.model = plan.reference.label
131
+ outputs, cached = _reference_outputs(plan, ref_backend, ref_info, cache, tracker)
132
+ ref_result = _reference_result(plan, ref_backend, ref_info, outputs, tracker, cached=cached)
133
+ results = tuple(
134
+ _run_candidate(
135
+ plan,
136
+ spec,
137
+ backend=backend,
138
+ info=info,
139
+ reference=ref_backend,
140
+ ref_info=ref_info,
141
+ outputs=outputs,
142
+ tracker=tracker,
143
+ )
144
+ for spec, (backend, info) in zip(plan.candidates, candidates, strict=True)
145
+ )
146
+ finally:
147
+ for backend, _ in backends:
148
+ backend.close()
149
+ tracker.finish()
150
+
151
+ return Report(
152
+ quantdiff_version=__version__,
153
+ created_at=datetime.now(timezone.utc).isoformat(timespec="seconds"),
154
+ title=plan.title,
155
+ settings=plan.settings,
156
+ reference=ref_result,
157
+ candidates=results,
158
+ )
159
+
160
+
161
+ def _total_units(plan: RunPlan) -> int:
162
+ preflight = _PREFLIGHT_UNITS if plan.preflight else 0
163
+ models = 1 + len(plan.candidates)
164
+ cases = _WARMUP_UNITS + _CASE_UNITS * len(plan.cases)
165
+ scoring = len(plan.scoring) * _scoring_units(plan)
166
+ return models + (cases + scoring + preflight) * models
167
+
168
+
169
+ def _scoring_units(plan: RunPlan) -> int:
170
+ """Units for one scoring prompt: one per generated or teacher-forced token."""
171
+ return max(plan.settings.score_tokens, 1)
172
+
173
+
174
+ def _connect(
175
+ plan: RunPlan, backend_factory: BackendFactory, tracker: _Progress
176
+ ) -> list[tuple[Backend, ServerInfo]]:
177
+ """Open every backend and fetch its info, or raise one error naming every failure."""
178
+ connected: list[tuple[Backend, ServerInfo]] = []
179
+ problems: list[str] = []
180
+ for spec in (plan.reference, *plan.candidates):
181
+ tracker.model = spec.label
182
+ try:
183
+ backend = backend_factory(spec)
184
+ except BackendError as exc:
185
+ problems.append(f"{spec.label}: {exc}")
186
+ continue
187
+ try:
188
+ connected.append((backend, backend.info()))
189
+ except BackendError as exc:
190
+ backend.close()
191
+ problems.append(f"{spec.label}: {exc}")
192
+ tracker.step("connect")
193
+ if problems:
194
+ for backend, _ in connected:
195
+ backend.close()
196
+ raise BackendError("cannot start the comparison:\n " + "\n ".join(problems))
197
+ return connected
198
+
199
+
200
+ # Reference --------------------------------------------------------------------------------
201
+
202
+
203
+ def _reference_outputs(
204
+ plan: RunPlan,
205
+ reference: Backend,
206
+ info: ServerInfo,
207
+ cache: ReferenceCache | None,
208
+ tracker: _Progress,
209
+ ) -> tuple[ReferenceOutputs, bool]:
210
+ """Return the reference outputs and whether they came from the cache."""
211
+ key = cache_key(
212
+ info,
213
+ base_url=plan.reference.base_url,
214
+ suite_digest=suite_digest(plan.cases, plan.scoring),
215
+ top_k=plan.settings.top_k,
216
+ score_tokens=plan.settings.score_tokens,
217
+ seed=plan.settings.seed,
218
+ )
219
+ if cache is not None:
220
+ cached = cache.load(key)
221
+ if cached is not None:
222
+ units = (
223
+ _WARMUP_UNITS
224
+ + _CASE_UNITS * len(plan.cases)
225
+ + _scoring_units(plan) * len(plan.scoring)
226
+ )
227
+ tracker.step("reference", units, "loaded from cache")
228
+ return cached, True
229
+
230
+ answers = _answer_cases(plan, reference, tracker, "reference")
231
+ if answers.failures:
232
+ logger.warning("reference failed %d case(s); they are left out", len(answers.failures))
233
+ traces = _reference_traces(plan, reference, info, tracker)
234
+ outputs = ReferenceOutputs(answers=answers.results, traces=traces)
235
+ if cache is not None and not answers.failures:
236
+ cache.store(key, outputs)
237
+ return outputs, False
238
+
239
+
240
+ def _reference_traces(
241
+ plan: RunPlan, reference: Backend, info: ServerInfo, tracker: _Progress
242
+ ) -> tuple[ReferenceTrace, ...]:
243
+ if not plan.scoring or not info.supports_logprobs:
244
+ tracker.step("scoring", _scoring_units(plan) * len(plan.scoring), "skipped: no logprobs")
245
+ return ()
246
+ traces = []
247
+ for index, prompt in enumerate(plan.scoring, start=1):
248
+ try:
249
+ steps = reference.generate_scored(
250
+ prompt.text, max_tokens=plan.settings.score_tokens, top_k=plan.settings.top_k
251
+ )
252
+ ids = reference.tokenize(prompt.text) if info.exact_token_ids else None
253
+ except CapabilityError as exc:
254
+ logger.warning("reference exposes no logprobs, skipping the logit tier: %s", exc)
255
+ remaining = len(plan.scoring) - index + 1
256
+ tracker.step("scoring", _scoring_units(plan) * remaining, "skipped: no logprobs")
257
+ return ()
258
+ except BackendError as exc:
259
+ logger.warning("reference failed scoring prompt %s: %s", prompt.id, exc)
260
+ else:
261
+ traces.append(ReferenceTrace(prompt.id, ids, tuple(steps)))
262
+ tracker.step("scoring", _scoring_units(plan), f"{index}/{len(plan.scoring)}")
263
+ return tuple(traces)
264
+
265
+
266
+ def _reference_result(
267
+ plan: RunPlan,
268
+ reference: Backend,
269
+ info: ServerInfo,
270
+ outputs: ReferenceOutputs,
271
+ tracker: _Progress,
272
+ *,
273
+ cached: bool,
274
+ ) -> CandidateResult:
275
+ findings = _preflight(plan, reference, None, tracker)
276
+ outcomes = _score_cases(plan, outputs.answers)
277
+ perf = perf_metrics(list(outputs.answers.values()))
278
+ return CandidateResult(
279
+ spec=plan.reference,
280
+ info=info,
281
+ logit=None,
282
+ tasks=summarize_tasks(outcomes),
283
+ agreement=None,
284
+ # Timings from an earlier run stay visible but labelled, since the machine's load
285
+ # may differ from the fresh candidate runs.
286
+ perf=replace(perf, source="cached") if cached else perf,
287
+ preflight=findings,
288
+ outcomes=outcomes,
289
+ )
290
+
291
+
292
+ # Candidates -------------------------------------------------------------------------------
293
+
294
+
295
+ def _run_candidate(
296
+ plan: RunPlan,
297
+ spec: CandidateSpec,
298
+ *,
299
+ backend: Backend,
300
+ info: ServerInfo,
301
+ reference: Backend,
302
+ ref_info: ServerInfo,
303
+ outputs: ReferenceOutputs,
304
+ tracker: _Progress,
305
+ ) -> CandidateResult:
306
+ tracker.model = spec.label
307
+ findings = _preflight(plan, backend, reference, tracker)
308
+ answers = _answer_cases(plan, backend, tracker, spec.label)
309
+ logit, logit_error = _logit_tier(
310
+ plan,
311
+ backend,
312
+ info=info,
313
+ ref_info=ref_info,
314
+ outputs=outputs,
315
+ findings=findings,
316
+ tracker=tracker,
317
+ )
318
+
319
+ errors = []
320
+ if answers.failures:
321
+ errors.append(_case_failure_summary(answers.failures))
322
+ if logit_error is not None:
323
+ errors.append(logit_error)
324
+ chats = [
325
+ (case.id, outputs.answers[case.id].text, answers.results[case.id].text)
326
+ for case in plan.cases
327
+ if case.kind == "chat" and case.id in answers.results and case.id in outputs.answers
328
+ ]
329
+ outcomes = _score_cases(plan, answers.results)
330
+ return CandidateResult(
331
+ spec=spec,
332
+ info=info,
333
+ logit=logit,
334
+ tasks=summarize_tasks(outcomes),
335
+ agreement=agreement_metrics(chats) if chats else None,
336
+ perf=perf_metrics(list(answers.results.values())),
337
+ preflight=findings,
338
+ errors=tuple(errors),
339
+ outcomes=outcomes,
340
+ )
341
+
342
+
343
+ def _logit_tier(
344
+ plan: RunPlan,
345
+ backend: Backend,
346
+ *,
347
+ info: ServerInfo,
348
+ ref_info: ServerInfo,
349
+ outputs: ReferenceOutputs,
350
+ findings: Sequence[PreflightFinding],
351
+ tracker: _Progress,
352
+ ) -> tuple[LogitMetrics | None, str | None]:
353
+ """Return logit metrics, or None plus a reason when the tier cannot run fairly."""
354
+ units = len(plan.scoring) * _scoring_units(plan)
355
+ if not outputs.traces or not info.supports_logprobs:
356
+ tracker.step("scoring", units, "skipped")
357
+ return None, None
358
+ if any(f.check == "tokenizer" and f.severity == "fail" for f in findings):
359
+ tracker.step("scoring", units, "skipped")
360
+ return None, "logit tier skipped: tokenizer differs from the reference"
361
+ return _teacher_force(
362
+ plan, backend, info=info, ref_info=ref_info, outputs=outputs, tracker=tracker
363
+ )
364
+
365
+
366
+ def _teacher_force(
367
+ plan: RunPlan,
368
+ backend: Backend,
369
+ *,
370
+ info: ServerInfo,
371
+ ref_info: ServerInfo,
372
+ outputs: ReferenceOutputs,
373
+ tracker: _Progress,
374
+ ) -> tuple[LogitMetrics | None, str | None]:
375
+ exact = ref_info.exact_token_ids and info.exact_token_ids
376
+ prompts = {prompt.id: prompt.text for prompt in plan.scoring}
377
+ traces = (
378
+ outputs.traces
379
+ if exact
380
+ else tuple(_text_scorable(trace, prompts[trace.prompt_id]) for trace in outputs.traces)
381
+ )
382
+ budget = len(plan.scoring) * _scoring_units(plan)
383
+ spent = 0
384
+ scored: list[list[TopK]] = []
385
+ for index, trace in enumerate(traces, start=1):
386
+ try:
387
+ tops = backend.score_continuation(
388
+ prompts[trace.prompt_id],
389
+ trace.steps,
390
+ top_k=plan.settings.top_k,
391
+ prompt_token_ids=trace.prompt_token_ids if exact else None,
392
+ )
393
+ except BackendError as exc:
394
+ tracker.step("scoring", budget - spent, "failed")
395
+ return None, f"logit tier failed: {exc}"
396
+ scored.append(tops)
397
+ spent += _scoring_units(plan)
398
+ tracker.step("scoring", _scoring_units(plan), f"{index}/{len(traces)}")
399
+ tracker.step("scoring", budget - spent)
400
+ return logit_metrics(traces, scored, exact_token_ids=exact), None
401
+
402
+
403
+ def _text_scorable(trace: ReferenceTrace, prompt: str) -> ReferenceTrace:
404
+ """The trace with every position that text cannot reproduce marked unscored.
405
+
406
+ A reference scored by token ids can stop inside a multi-byte character, and no text
407
+ prompt ends there. Marking those positions unscored (an empty `top`) makes metrics skip
408
+ them for this candidate instead of counting the empty result as a top-1 miss.
409
+ """
410
+ texts = forced_texts(prompt, trace.steps)
411
+ steps = tuple(
412
+ step if text is not None else TokenStep(step.chosen, ())
413
+ for step, text in zip(trace.steps, texts, strict=True)
414
+ )
415
+ return ReferenceTrace(trace.prompt_id, trace.prompt_token_ids, steps)
416
+
417
+
418
+ # Shared steps -----------------------------------------------------------------------------
419
+
420
+
421
+ def _answer_cases(plan: RunPlan, backend: Backend, tracker: _Progress, label: str) -> _Answers:
422
+ _warm_up(plan, backend, tracker, label)
423
+ answers = _Answers()
424
+ total = len(plan.cases)
425
+ for index, case in enumerate(plan.cases, start=1):
426
+ try:
427
+ answers.results[case.id] = backend.chat(
428
+ case.messages,
429
+ max_tokens=case.max_tokens,
430
+ tools=case.tools,
431
+ json_schema=case.json_schema if case.kind == "json" else None,
432
+ seed=plan.settings.seed,
433
+ )
434
+ except BackendError as exc:
435
+ logger.warning("%s failed case %s: %s", label, case.id, exc)
436
+ answers.failures.append(f"{case.id}: {exc}")
437
+ tracker.step("cases", _CASE_UNITS, f"{index}/{total}")
438
+ return answers
439
+
440
+
441
+ def _warm_up(plan: RunPlan, backend: Backend, tracker: _Progress, label: str) -> None:
442
+ """Send one tiny request so model load time never lands in the first case's timing.
443
+
444
+ A failure here is only logged: if the server is really broken, the cases report it.
445
+ """
446
+ try:
447
+ backend.chat(_WARMUP_MESSAGES, max_tokens=1, seed=plan.settings.seed)
448
+ except BackendError as exc:
449
+ logger.debug("%s warm-up request failed: %s", label, exc)
450
+ tracker.step("cases", _WARMUP_UNITS, "warm-up")
451
+
452
+
453
+ def _case_failure_summary(failures: Sequence[str]) -> str:
454
+ first = failures[0]
455
+ more = f" (and {len(failures) - 1} more)" if len(failures) > 1 else ""
456
+ return f"{len(failures)} case(s) failed with server errors and were not scored: {first}{more}"
457
+
458
+
459
+ def _score_cases(plan: RunPlan, answers: dict[str, ChatResult]) -> tuple[CaseOutcome, ...]:
460
+ """Outcomes of the answered json, tools and code cases, in plan order."""
461
+ return tuple(
462
+ evaluate_case(
463
+ case,
464
+ answers[case.id],
465
+ allow_code_exec=plan.settings.allow_code_exec,
466
+ timeout_seconds=plan.code_timeout_seconds,
467
+ )
468
+ for case in plan.cases
469
+ if case.id in answers and case.kind in SCORED_KINDS
470
+ )
471
+
472
+
473
+ def _preflight(
474
+ plan: RunPlan, backend: Backend, reference: Backend | None, tracker: _Progress
475
+ ) -> tuple[PreflightFinding, ...]:
476
+ if not plan.preflight:
477
+ return ()
478
+ findings = run_preflight(
479
+ backend,
480
+ reference=reference,
481
+ hf_repo=plan.hf_repo,
482
+ offline=plan.offline,
483
+ context_probe_tokens=plan.context_probe_tokens,
484
+ )
485
+ tracker.step("preflight", _PREFLIGHT_UNITS)
486
+ return findings
487
+
488
+
489
+ def _log_progress(event: ProgressEvent) -> None:
490
+ logger.info(
491
+ "%s %s %d/%d %s", event.phase, event.model, event.completed, event.total, event.detail
492
+ )
quantdiff/spec.py ADDED
@@ -0,0 +1,154 @@
1
+ """Parse candidate spec strings from the command line into CandidateSpec values.
2
+
3
+ Grammar::
4
+
5
+ spec = [label "="] target
6
+ target = "ollama:" tag
7
+ | "llamacpp:" base_url
8
+ | "openai:" base_url "#" model ["@env:" env_var]
9
+
10
+ label = any non-empty text before the first "=" that contains no ":"
11
+ env_var = name of an environment variable holding an API key ([A-Za-z_][A-Za-z0-9_]*)
12
+
13
+ Examples::
14
+
15
+ ollama:qwen2.5:0.5b-instruct-q8_0
16
+ q4=ollama:qwen2.5:7b-instruct-q4_K_M
17
+ llamacpp:http://127.0.0.1:8080
18
+ openai:http://127.0.0.1:1234/v1#qwen2.5-7b-instruct
19
+ vllm=openai:http://gpu-box:8000/v1#Qwen/Qwen2.5-7B-Instruct@env:VLLM_API_KEY
20
+
21
+ Ollama specs use OLLAMA_HOST (`host`, `host:port` or a full URL) when it is set, and
22
+ http://127.0.0.1:11434 otherwise. Default labels are the Ollama tag, `llamacpp@host:port`
23
+ and the OpenAI model name. Making labels unique across candidates is the caller's job.
24
+ """
25
+
26
+ from __future__ import annotations
27
+
28
+ import dataclasses
29
+ import os
30
+ import re
31
+ import urllib.parse
32
+ from typing import Final
33
+
34
+ from quantdiff._http import validate_base_url
35
+ from quantdiff.errors import BackendError, SpecError
36
+ from quantdiff.types import CandidateSpec
37
+
38
+ DEFAULT_OLLAMA_URL: Final = "http://127.0.0.1:11434"
39
+ _OLLAMA_DEFAULT_PORT: Final = 11434
40
+ _API_KEY_MARKER: Final = "@env:"
41
+ _ENV_NAME: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
42
+ _OLLAMA_TAG: Final = re.compile(r"[A-Za-z0-9][A-Za-z0-9._:/-]*")
43
+ _BIND_ALL_HOSTS: Final = frozenset({"0.0.0.0", "::"}) # noqa: S104 - matched, not bound
44
+
45
+
46
+ def parse_spec(text: str) -> CandidateSpec:
47
+ """Parse one candidate spec string. Raises SpecError if it is malformed."""
48
+ stripped = text.strip()
49
+ if not stripped:
50
+ raise SpecError("empty candidate spec")
51
+ label, target = _split_label(stripped)
52
+ kind, separator, rest = target.partition(":")
53
+ if not separator or not rest:
54
+ raise SpecError(f"{text!r}: expected <kind>:<target>, for example ollama:<tag>")
55
+ if kind == "ollama":
56
+ spec = _ollama_spec(rest)
57
+ elif kind == "llamacpp":
58
+ spec = _llamacpp_spec(rest)
59
+ elif kind == "openai":
60
+ spec = _openai_spec(rest)
61
+ else:
62
+ raise SpecError(f"{text!r}: unknown backend {kind!r}; use ollama, llamacpp or openai")
63
+ return spec if label is None else dataclasses.replace(spec, label=label)
64
+
65
+
66
+ def ollama_base_url(host: str | None) -> str:
67
+ """Resolve an OLLAMA_HOST value the way the Ollama CLI does."""
68
+ value = (host or "").strip()
69
+ if not value:
70
+ return DEFAULT_OLLAMA_URL
71
+ url = value if "://" in value else f"http://{value}"
72
+ parsed = urllib.parse.urlsplit(url)
73
+ hostname = parsed.hostname
74
+ if not hostname:
75
+ raise SpecError(f"OLLAMA_HOST={value!r} has no host")
76
+ # OLLAMA_HOST is often a listen address such as 0.0.0.0, which clients cannot dial.
77
+ if hostname in _BIND_ALL_HOSTS:
78
+ hostname = "127.0.0.1"
79
+ port = _port(parsed, value)
80
+ if port is None and "://" not in value:
81
+ port = _OLLAMA_DEFAULT_PORT
82
+ netloc = f"[{hostname}]" if ":" in hostname else hostname
83
+ if port is not None:
84
+ netloc = f"{netloc}:{port}"
85
+ return _base_url(urllib.parse.urlunsplit((parsed.scheme, netloc, parsed.path, "", "")))
86
+
87
+
88
+ def _split_label(text: str) -> tuple[str | None, str]:
89
+ head, separator, tail = text.partition("=")
90
+ if not separator or ":" in head:
91
+ return None, text
92
+ label = head.strip()
93
+ if not label:
94
+ raise SpecError(f"{text!r}: the label before '=' is empty")
95
+ if not label.isprintable():
96
+ raise SpecError(f"{text!r}: the label contains control characters")
97
+ return label, tail.strip()
98
+
99
+
100
+ def _ollama_spec(tag: str) -> CandidateSpec:
101
+ if not _OLLAMA_TAG.fullmatch(tag):
102
+ raise SpecError(f"invalid Ollama tag {tag!r}")
103
+ return CandidateSpec(
104
+ kind="ollama",
105
+ base_url=ollama_base_url(os.environ.get("OLLAMA_HOST")),
106
+ model=tag,
107
+ label=tag,
108
+ )
109
+
110
+
111
+ def _llamacpp_spec(url: str) -> CandidateSpec:
112
+ base_url = _base_url(url)
113
+ return CandidateSpec(
114
+ kind="llamacpp",
115
+ base_url=base_url,
116
+ model="",
117
+ label=f"llamacpp@{urllib.parse.urlsplit(base_url).netloc}",
118
+ )
119
+
120
+
121
+ def _openai_spec(rest: str) -> CandidateSpec:
122
+ url, separator, model_part = rest.partition("#")
123
+ if not separator:
124
+ raise SpecError(f"openai spec {rest!r} needs a model: openai:<base_url>#<model>")
125
+ model, marker, env_name = model_part.rpartition(_API_KEY_MARKER)
126
+ if not marker:
127
+ model, env_name = model_part, ""
128
+ model = model.strip()
129
+ if not model or any(char.isspace() for char in model):
130
+ raise SpecError(f"invalid model name {model!r} in openai spec")
131
+ if marker and not _ENV_NAME.fullmatch(env_name):
132
+ raise SpecError(f"invalid environment variable name {env_name!r} after {_API_KEY_MARKER}")
133
+ return CandidateSpec(
134
+ kind="openai",
135
+ base_url=_base_url(url),
136
+ model=model,
137
+ label=model,
138
+ api_key_env=env_name or None,
139
+ )
140
+
141
+
142
+ def _base_url(url: str) -> str:
143
+ try:
144
+ return validate_base_url(url.strip())
145
+ except BackendError as exc:
146
+ # The message from validate_base_url never echoes credentials; the raw URL might.
147
+ raise SpecError(f"invalid server URL: {exc}") from None
148
+
149
+
150
+ def _port(parsed: urllib.parse.SplitResult, raw: str) -> int | None:
151
+ try:
152
+ return parsed.port
153
+ except ValueError:
154
+ raise SpecError(f"OLLAMA_HOST={raw!r} has an invalid port") from None