stuntd 0.1.0__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.
stuntd/__init__.py ADDED
@@ -0,0 +1,5 @@
1
+ from __future__ import annotations
2
+
3
+ __all__ = ["__version__"]
4
+
5
+ __version__ = "0.1.0"
stuntd/__main__.py ADDED
@@ -0,0 +1,5 @@
1
+ from __future__ import annotations
2
+
3
+ from stuntd.cli import main
4
+
5
+ raise SystemExit(main())
stuntd/cli.py ADDED
@@ -0,0 +1,665 @@
1
+ from __future__ import annotations
2
+
3
+ import argparse
4
+ import json
5
+ import math
6
+ import sqlite3
7
+ import sys
8
+ import time
9
+ from collections.abc import Callable
10
+ from dataclasses import asdict, dataclass, fields
11
+ from pathlib import Path
12
+ from typing import TYPE_CHECKING
13
+
14
+ from stuntd.decisions.schema import DecisionSchema, detect_schema
15
+ from stuntd.paths import private_file
16
+ from stuntd.serve.modes import (
17
+ MODE_CHECK,
18
+ MODE_COLLECT,
19
+ MODE_LIVE,
20
+ MODE_SHADOW,
21
+ SiteState,
22
+ site_states,
23
+ write_mode,
24
+ )
25
+ from stuntd.serve.monitor import window_agreement, window_start
26
+ from stuntd.settings import (
27
+ CONFIG_TEMPLATE,
28
+ Settings,
29
+ config_path,
30
+ database_path,
31
+ load_settings,
32
+ models_path,
33
+ )
34
+ from stuntd.train.artifacts import HEAD_FILE, SiteModel, list_models, load_model, site_dir
35
+ from stuntd.train.report import render, render_json
36
+ from stuntd.train.run import NO_CAPTURES, Trainer, TrainResult, train_sites
37
+
38
+ if TYPE_CHECKING:
39
+ from stuntd.serve.runtime import DeciderLike
40
+ from stuntd.store.db import Capture, Store
41
+
42
+ __all__ = ["main"]
43
+
44
+ _COLUMN_WIDTHS = (18, 8, 9, 8, 6)
45
+ _CONFIG_HELP = "settings file to read instead of the one in the data directory"
46
+ _IMPORT_MODEL = "import"
47
+ _COVERAGE_DIGITS = 2
48
+ # Below this share of the holdout a head answers so little that the operator should hear about it.
49
+ _LOW_COVERAGE = 0.10
50
+ _KINDS = ("choice", "boolean", "number")
51
+ _BOOLEAN_ANSWERS = ("true", "false")
52
+
53
+
54
+ @dataclass(frozen=True)
55
+ class _SiteStatus:
56
+ """One row of the status table: how a site serves and what it has answered lately."""
57
+
58
+ site: str
59
+ mode: str
60
+ captures: int
61
+ shadow: int
62
+ live: int
63
+ agreement: float | None
64
+
65
+
66
+ class _LearningOff(Exception):
67
+ """A command that would write captures, models or decisions while learning is off."""
68
+
69
+
70
+ def _require_learning(config: str | None, settings: Settings) -> None:
71
+ if not settings.learn:
72
+ raise _LearningOff(f"learning is off in {_config_file(config)}")
73
+
74
+
75
+ def _make_decider(settings: Settings) -> DeciderLike:
76
+ # The only place serving reaches torch and laya, so every other command runs without the extra.
77
+ from stuntd.serve.decider import Decider
78
+
79
+ return Decider(settings.base_model, settings.device)
80
+
81
+
82
+ def _require_serving() -> None:
83
+ """Refuses the way serve would when the train extra is missing, without loading a model."""
84
+ try:
85
+ import stuntd.serve.decider # noqa: F401
86
+ except ImportError as exc:
87
+ raise ValueError('serving needs the train extra: pip install "stuntd[train]"') from exc
88
+
89
+
90
+ def _serve_decider(settings: Settings, serving: int) -> DeciderLike | None:
91
+ # A local Jev answers its questions zero-shot from the base checkpoint, trained head or not,
92
+ # so it needs the decider even where no site serves one of its own.
93
+ local_jev = settings.jev_upstream == ""
94
+ if serving == 0 and not local_jev:
95
+ return None
96
+ try:
97
+ return _make_decider(settings)
98
+ except ImportError:
99
+ # The proxy still records captures, so a missing extra costs the answers, not the run.
100
+ print('stuntd: serving disabled: pip install "stuntd[train]"', file=sys.stderr)
101
+ return None
102
+
103
+
104
+ def _serving_sites(models: Path) -> list[tuple[str, SiteModel]]:
105
+ return [
106
+ (state.site, state.model)
107
+ for state in site_states(models)
108
+ if state.mode != MODE_COLLECT and state.model is not None
109
+ ]
110
+
111
+
112
+ def _serve(args: argparse.Namespace) -> int:
113
+ settings = _load(args.config, args.upstream, args.port)
114
+ import uvicorn
115
+
116
+ from stuntd.proxy.app import build_app
117
+
118
+ models = models_path(settings)
119
+ # With learning off no head answers, so the sites on disk serve nothing.
120
+ serving = _serving_sites(models) if settings.learn else []
121
+ # The base model loads before anything is printed: it takes seconds, and a checkpoint that
122
+ # will not load must fail the command rather than leave a line promising a proxy that is up.
123
+ decider = _serve_decider(settings, len(serving))
124
+ if decider is not None and serving:
125
+ # The first pass through a freshly loaded model is far slower than the rest, so one site
126
+ # pays for it here rather than the first caller of whichever site asks first.
127
+ site, model = serving[0]
128
+ decider.warm(model, site_dir(models, site) / HEAD_FILE)
129
+ elif decider is not None:
130
+ decider.warm_base()
131
+ # uvicorn.run never returns while the daemon is up, so a piped stdout needs the lines now.
132
+ listening = f"stuntd listening on http://{settings.host}:{settings.port}"
133
+ if settings.upstream:
134
+ print(f"{listening} -> {settings.upstream}", flush=True)
135
+ else:
136
+ print(listening, flush=True)
137
+ print("no upstream: only the Jev routes are served", flush=True)
138
+ if decider is not None:
139
+ served = f"{len(serving)} site(s)" if serving else "jev locally"
140
+ print(f"serving {served} with {settings.base_model}", flush=True)
141
+ if not settings.learn:
142
+ print("learning off", flush=True)
143
+ # The relay is byte-exact, so uvicorn must not put its own Date and Server headers next to
144
+ # the ones the provider sent.
145
+ uvicorn.run(
146
+ build_app(settings, decider=decider),
147
+ host=settings.host,
148
+ port=settings.port,
149
+ log_level="warning",
150
+ server_header=False,
151
+ date_header=False,
152
+ )
153
+ return 0
154
+
155
+
156
+ def _enable(args: argparse.Namespace) -> int:
157
+ settings = _load(args.config, None, None)
158
+ _require_learning(args.config, settings)
159
+ models = models_path(settings)
160
+ site = _site_name(args.site)
161
+ model = _one_model(models, site)
162
+ if model.threshold is None:
163
+ print(
164
+ f"stuntd: {site} never reached the target agreement on its holdout;"
165
+ " retrain with more examples",
166
+ file=sys.stderr,
167
+ )
168
+ return 1
169
+ # The curve always answers at least one holdout row, so a useless operating point shows up
170
+ # as the coverage the report rounds to, not as a bare zero.
171
+ coverage = None if model.coverage is None else round(model.coverage, _COVERAGE_DIGITS)
172
+ if coverage == 0:
173
+ print(
174
+ f"stuntd: {site} covers 0% of the holdout at target"
175
+ f" {model.target_agreement:.2f}: nothing will be answered locally;"
176
+ " lower training.target_agreement or retrain",
177
+ file=sys.stderr,
178
+ )
179
+ return 1
180
+ if coverage is not None and coverage < _LOW_COVERAGE:
181
+ print(
182
+ f"stuntd: {site} covers {coverage:.0%} of the holdout at target"
183
+ f" {model.target_agreement:.2f}; most answers will still go to the provider",
184
+ file=sys.stderr,
185
+ )
186
+ _require_serving()
187
+ # The proxy re-reads the file once it changes, so the running daemon needs no restart.
188
+ write_mode(site_dir(models, site), MODE_LIVE, time.time())
189
+ print(f"{site}: {MODE_LIVE}")
190
+ return 0
191
+
192
+
193
+ def _disable(args: argparse.Namespace) -> int:
194
+ settings = _load(args.config, None, None)
195
+ _require_learning(args.config, settings)
196
+ models = models_path(settings)
197
+ site = _site_name(args.site)
198
+ _one_model(models, site)
199
+ write_mode(site_dir(models, site), MODE_SHADOW, time.time())
200
+ print(f"{site}: {MODE_SHADOW}")
201
+ return 0
202
+
203
+
204
+ def _status(args: argparse.Namespace) -> int:
205
+ settings = _load(args.config, None, None)
206
+ rows = _status_rows(settings) if settings.learn else []
207
+ if args.json:
208
+ print(json.dumps({"learning": settings.learn, "sites": [asdict(row) for row in rows]}))
209
+ return 0
210
+ if not settings.learn:
211
+ print("learning off")
212
+ return 0
213
+ if not rows:
214
+ print("no captures yet")
215
+ return 0
216
+ print(_row("site", "mode", "captures", "shadow", "live", "agreement"))
217
+ for row in rows:
218
+ print(_status_line(row))
219
+ return 0
220
+
221
+
222
+ def _config_init(args: argparse.Namespace) -> int:
223
+ path = Path(args.config) if args.config else config_path()
224
+ if path.exists() and not args.force:
225
+ print(f"{path} exists; use --force to overwrite")
226
+ return 1
227
+ path.parent.mkdir(parents=True, exist_ok=True)
228
+ private_file(path).write_text(CONFIG_TEMPLATE, encoding="utf-8")
229
+ print(f"wrote {path}")
230
+ return 0
231
+
232
+
233
+ def _config_show(args: argparse.Namespace) -> int:
234
+ settings = _load(args.config, args.upstream, args.port)
235
+ flagged = {name for name in ("upstream", "port") if getattr(args, name) is not None}
236
+ defaults = Settings()
237
+ for item in fields(Settings):
238
+ value = getattr(settings, item.name)
239
+ if item.name in flagged:
240
+ source = "flag"
241
+ elif value != getattr(defaults, item.name):
242
+ source = "file"
243
+ else:
244
+ source = "default"
245
+ if item.name == "redaction_patterns":
246
+ shown = ", ".join(settings.redaction_patterns)
247
+ else:
248
+ shown = "" if value is None else str(value)
249
+ print(f"{item.name} = {shown} ({source})")
250
+ return 0
251
+
252
+
253
+ def _make_trainer(settings: Settings) -> Trainer:
254
+ # The only place torch and laya are reached, so every other command runs without the extra.
255
+ from stuntd.train.trainer import LayaTrainer
256
+
257
+ return LayaTrainer(
258
+ settings.base_model,
259
+ settings.device,
260
+ settings.epochs,
261
+ cache_encoder=settings.cache_encoder,
262
+ cache_max_bytes=settings.cache_max_mb * 2**20,
263
+ )
264
+
265
+
266
+ def _load_trainer(settings: Settings) -> Trainer:
267
+ try:
268
+ return _make_trainer(settings)
269
+ except ImportError as exc:
270
+ raise ValueError('training needs the train extra: pip install "stuntd[train]"') from exc
271
+
272
+
273
+ def _base_model_line(settings: Settings) -> str:
274
+ # A checkpoint already on disk is read from there, so only a model name is fetched.
275
+ if Path(settings.base_model).is_dir():
276
+ return f"base model: {settings.base_model}"
277
+ return f"base model: {settings.base_model} (downloaded on first use)"
278
+
279
+
280
+ def _result_line(result: TrainResult) -> str:
281
+ model = result.model
282
+ if model is None:
283
+ return f"{result.site} skipped {result.reason}"
284
+ headline = f"{result.site} trained holdout agreement {model.agreement:.3f}"
285
+ if model.coverage is None or model.covered_agreement is None:
286
+ return f"{headline} coverage - ece {model.ece:.3f}"
287
+ return (
288
+ f"{headline} coverage {model.coverage:.2f}"
289
+ f" at agreement {model.covered_agreement:.3f} ece {model.ece:.3f}"
290
+ )
291
+
292
+
293
+ def _train(args: argparse.Namespace) -> int:
294
+ from stuntd.store.db import Store
295
+ from stuntd.store.redact import Redactor
296
+
297
+ settings = _load(args.config, None, None)
298
+ _require_learning(args.config, settings)
299
+ database = database_path(settings)
300
+ if not database.exists():
301
+ print("no captures yet")
302
+ return 0
303
+ store = Store(database, Redactor())
304
+ try:
305
+ known = {info.site for info in store.sites()}
306
+ sites = [_site_name(site) for site in args.sites]
307
+ if not known and not sites:
308
+ print("no captures yet")
309
+ return 0
310
+ trainable = [site for site in sites if site in known] if sites else list(known)
311
+ if not trainable:
312
+ results = [TrainResult(site, None, NO_CAPTURES) for site in sites]
313
+ else:
314
+ # Building the trainer fetches the checkpoint, which is a long silent wait without
315
+ # the line, so it waits until a site actually needs training.
316
+ print(_base_model_line(settings), flush=True)
317
+ results = train_sites(store, settings, _load_trainer(settings), sites)
318
+ finally:
319
+ store.close()
320
+ for result in results:
321
+ print(_result_line(result))
322
+ return 0
323
+
324
+
325
+ class _InvalidLine(Exception):
326
+ """A row of an import file that cannot be recorded, already named by its line number."""
327
+
328
+
329
+ def _decision_schema(schema: object) -> DecisionSchema:
330
+ # Read back the way a request is, so an imported site carries the canonical schema and the
331
+ # kind the same decision would have got through the proxy.
332
+ found = detect_schema(
333
+ {"response_format": {"type": "json_schema", "json_schema": {"schema": schema}}}
334
+ )
335
+ if found is None:
336
+ raise ValueError("schema must describe an object with one typed field")
337
+ return found
338
+
339
+
340
+ def _built_schema(site: str, kind: str, labels: str | None) -> dict[str, object]:
341
+ if kind != "choice":
342
+ spec: dict[str, object] = {"type": "boolean" if kind == "boolean" else "integer"}
343
+ else:
344
+ # A label repeated in the list would become a class no example can carry.
345
+ options = dict.fromkeys(
346
+ label.strip() for label in (labels or "").split(",") if label.strip()
347
+ )
348
+ if not options:
349
+ raise ValueError("a choice site needs --labels a,b,c or --schema")
350
+ spec = {"type": "string", "enum": list(options)}
351
+ # The field is what the head is asked for at serve time, so a site built here reads the same
352
+ # way a Jev question of that name does rather than asking the model to "choose answer".
353
+ return {"type": "object", "properties": {site: spec}}
354
+
355
+
356
+ def _import_schema(
357
+ site: str, kind: str | None, raw: str | None, labels: str | None
358
+ ) -> DecisionSchema:
359
+ if raw is None:
360
+ return _decision_schema(_built_schema(site, kind or "choice", labels))
361
+ try:
362
+ given = json.loads(raw)
363
+ except json.JSONDecodeError as exc:
364
+ raise ValueError(f"--schema is not valid JSON: {exc.msg}") from exc
365
+ schema = _decision_schema(given)
366
+ if kind is not None and kind != schema.kind:
367
+ raise ValueError(f"--schema describes a {schema.kind} field, not {kind}")
368
+ return schema
369
+
370
+
371
+ def _check_answer(schema: DecisionSchema, answer: str) -> None:
372
+ if schema.kind == "choice":
373
+ if answer not in schema.options:
374
+ raise ValueError(f"answer {answer!r} not in labels")
375
+ elif schema.kind == "boolean":
376
+ if answer not in _BOOLEAN_ANSWERS:
377
+ raise ValueError(f"answer {answer!r} is not true or false")
378
+ else:
379
+ try:
380
+ value = float(answer)
381
+ except ValueError as exc:
382
+ raise ValueError(f"answer {answer!r} is not a number") from exc
383
+ # An infinity or a NaN parses and is then dropped when the dataset is built, which would
384
+ # cost rows the import reported as recorded.
385
+ if not math.isfinite(value):
386
+ raise ValueError(f"answer {answer!r} is not a finite number")
387
+
388
+
389
+ def _labelled_row(line: str, schema: DecisionSchema) -> tuple[str, str]:
390
+ try:
391
+ row = json.loads(line)
392
+ except json.JSONDecodeError as exc:
393
+ raise ValueError(f"invalid JSON: {exc.msg}") from exc
394
+ if not isinstance(row, dict):
395
+ raise ValueError("expected an object with text and answer")
396
+ text, answer = row.get("text"), row.get("answer")
397
+ if not isinstance(text, str) or not text.strip():
398
+ raise ValueError("text must be a non-empty string")
399
+ if not isinstance(answer, str):
400
+ raise ValueError("answer must be a string")
401
+ _check_answer(schema, answer)
402
+ return text, answer
403
+
404
+
405
+ def _import_captures(path: Path, site: str, schema: DecisionSchema) -> list[Capture]:
406
+ from stuntd.store.db import Capture
407
+
408
+ captures = []
409
+ # utf-8-sig so a file written by a Windows editor is not read with its BOM glued to line 1.
410
+ for number, line in enumerate(path.read_text(encoding="utf-8-sig").splitlines(), start=1):
411
+ if not line.strip():
412
+ continue
413
+ try:
414
+ text, answer = _labelled_row(line, schema)
415
+ except ValueError as exc:
416
+ raise _InvalidLine(f"line {number}: {exc}") from exc
417
+ captures.append(
418
+ Capture(site, schema.canonical, schema.kind, text, answer, _IMPORT_MODEL, 0, None, None)
419
+ )
420
+ return captures
421
+
422
+
423
+ def _import(args: argparse.Namespace) -> int:
424
+ from stuntd.store.db import Store
425
+ from stuntd.store.redact import Redactor
426
+
427
+ settings = _load(args.config, None, None)
428
+ _require_learning(args.config, settings)
429
+ site = _site_name(args.site)
430
+ # Nothing is trained here; the name only has to pass the whitelist it would be stored under.
431
+ site_dir(models_path(settings), site)
432
+ schema = _import_schema(site, args.kind, args.schema, args.labels)
433
+ try:
434
+ # The whole file is read before the store is opened, so a rejected row leaves no database.
435
+ captures = _import_captures(Path(args.file), site, schema)
436
+ except _InvalidLine as exc:
437
+ print(exc, file=sys.stderr)
438
+ return 1
439
+ store = Store(
440
+ database_path(settings),
441
+ Redactor(settings.redaction_patterns, builtin=settings.redact),
442
+ max_rows=settings.max_rows,
443
+ max_age_days=settings.max_age_days,
444
+ )
445
+ # Importing the same file twice records its rows twice; training keeps the latest row per text.
446
+ try:
447
+ for capture in captures:
448
+ store.record(capture)
449
+ finally:
450
+ store.close()
451
+ print(f"imported {len(captures)} rows into {site}")
452
+ return 0
453
+
454
+
455
+ def _site_name(site: str) -> str:
456
+ # A Jev question namespaced with a colon lands in a folder named with a dot, so either
457
+ # spelling of the name reaches the same site whichever command is given it.
458
+ return site.replace(":", ".")
459
+
460
+
461
+ def _one_model(models: Path, site: str) -> SiteModel:
462
+ try:
463
+ return load_model(models, site)
464
+ except FileNotFoundError as exc:
465
+ raise ValueError(f"no model for {site}") from exc
466
+
467
+
468
+ def _report(args: argparse.Namespace) -> int:
469
+ settings = _load(args.config, None, None)
470
+ _require_learning(args.config, settings)
471
+ models = models_path(settings)
472
+ site = None if args.site is None else _site_name(args.site)
473
+ found = list_models(models) if site is None else [_one_model(models, site)]
474
+ if args.json:
475
+ print(render_json(found))
476
+ return 0
477
+ if not found:
478
+ print("no trained sites yet")
479
+ return 0
480
+ print("\n\n".join(render(model, args.curve) for model in found))
481
+ return 0
482
+
483
+
484
+ def _row(
485
+ site: object, mode: object, captures: object, shadow: object, live: object, agreement: object
486
+ ) -> str:
487
+ # The space past each pad keeps the columns apart when a value fills its width.
488
+ columns = zip((site, mode, captures, shadow, live), _COLUMN_WIDTHS, strict=True)
489
+ return "".join(f"{value:<{width - 1}} " for value, width in columns) + str(agreement)
490
+
491
+
492
+ def _status_line(row: _SiteStatus) -> str:
493
+ # A site with no model has nothing to compare or to serve, which is not the same as none yet.
494
+ served = row.mode != MODE_COLLECT
495
+ return _row(
496
+ row.site,
497
+ row.mode,
498
+ row.captures,
499
+ row.shadow if served else "-",
500
+ row.live if served else "-",
501
+ "-" if row.agreement is None else f"{row.agreement:.3f}",
502
+ )
503
+
504
+
505
+ def _compared(counts: dict[str, int]) -> int:
506
+ # A checked live request is answered by both sides, so it counts as a comparison too.
507
+ return counts.get(MODE_SHADOW, 0) + counts.get(MODE_CHECK, 0)
508
+
509
+
510
+ def _agreement(store: Store, site: str, state: SiteState | None, window: int) -> float | None:
511
+ if state is None or state.model is None or state.model.threshold is None:
512
+ return None
513
+ rows = store.comparisons(site, window, window_start(state))
514
+ return window_agreement(rows, state.model.threshold).agreement
515
+
516
+
517
+ def _row_without_captures(state: SiteState) -> _SiteStatus:
518
+ return _SiteStatus(state.site, state.mode, 0, 0, 0, None)
519
+
520
+
521
+ def _status_rows(settings: Settings) -> list[_SiteStatus]:
522
+ from stuntd.store.db import Store
523
+ from stuntd.store.redact import Redactor
524
+
525
+ states = {state.site: state for state in site_states(models_path(settings))}
526
+ database = database_path(settings)
527
+ if not database.exists():
528
+ # Opening a store would create the database that status is only here to read.
529
+ return [_row_without_captures(state) for state in states.values()]
530
+ store = Store(database, Redactor())
531
+ try:
532
+ captures = {info.site: info.count for info in store.sites()}
533
+ counts = store.decision_counts()
534
+ rows = []
535
+ for site in sorted(set(captures) | set(states)):
536
+ state = states.get(site)
537
+ decisions = counts.get(site, {})
538
+ rows.append(
539
+ _SiteStatus(
540
+ site=site,
541
+ mode=MODE_COLLECT if state is None else state.mode,
542
+ captures=captures.get(site, 0),
543
+ shadow=_compared(decisions),
544
+ live=decisions.get(MODE_LIVE, 0),
545
+ agreement=_agreement(store, site, state, settings.window),
546
+ )
547
+ )
548
+ return rows
549
+ finally:
550
+ store.close()
551
+
552
+
553
+ def _config_file(explicit: str | None) -> Path:
554
+ if explicit is None:
555
+ return config_path()
556
+ path = Path(explicit)
557
+ if path.is_dir():
558
+ raise ValueError(f"{path} is not a file")
559
+ if not path.is_file():
560
+ raise ValueError(f"{path} does not exist")
561
+ return path
562
+
563
+
564
+ def _load(config: str | None, upstream: str | None, port: int | None) -> Settings:
565
+ return load_settings(_config_file(config), {"upstream": upstream, "port": port})
566
+
567
+
568
+ def _build_parser() -> argparse.ArgumentParser:
569
+ parser = argparse.ArgumentParser(
570
+ prog="stuntd", description="Local proxy that records typed LLM decisions."
571
+ )
572
+ commands = parser.add_subparsers(dest="command", required=True)
573
+
574
+ serve = commands.add_parser("serve", help="run the proxy until it is interrupted")
575
+ serve.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
576
+ serve.add_argument("--upstream", metavar="URL", help="provider base URL to forward to")
577
+ serve.add_argument("--port", type=int, metavar="N", help="port to listen on")
578
+ serve.set_defaults(handler=_serve)
579
+
580
+ status = commands.add_parser("status", help="summarise the captures recorded so far")
581
+ status.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
582
+ status.add_argument("--json", action="store_true", help="print the rows as JSON")
583
+ status.set_defaults(handler=_status)
584
+
585
+ train = commands.add_parser("train", help="train a model for each site with enough captures")
586
+ train.add_argument(
587
+ "sites", nargs="*", metavar="SITE", help="sites to train, every known site by default"
588
+ )
589
+ train.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
590
+ train.set_defaults(handler=_train)
591
+
592
+ report = commands.add_parser("report", help="print how the trained models did")
593
+ report.add_argument(
594
+ "site", nargs="?", metavar="SITE", help="site to report on, every trained site by default"
595
+ )
596
+ report.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
597
+ report.add_argument("--json", action="store_true", help="print the models as JSON")
598
+ report.add_argument("--curve", action="store_true", help="add the whole threshold curve")
599
+ report.set_defaults(handler=_report)
600
+
601
+ enable = commands.add_parser("enable", help="let a trained site answer from its own model")
602
+ enable.add_argument("site", metavar="SITE", help="site to serve locally")
603
+ enable.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
604
+ enable.set_defaults(handler=_enable)
605
+
606
+ disable = commands.add_parser("disable", help="send a site back to the provider")
607
+ disable.add_argument("site", metavar="SITE", help="site to stop serving locally")
608
+ disable.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
609
+ disable.set_defaults(handler=_disable)
610
+
611
+ importer = commands.add_parser(
612
+ "import",
613
+ help="record labelled examples for a site",
614
+ description=(
615
+ "Record labelled examples for a site from a JSONL file. Each text is stored and"
616
+ " served exactly as written, so it has to be spelled the way the site sees its"
617
+ " input: for a chat site the 'user: ...' lines the store holds, for a Jev site the"
618
+ " serialised state."
619
+ ),
620
+ )
621
+ importer.add_argument("site", metavar="SITE", help="site to record the examples under")
622
+ importer.add_argument(
623
+ "file",
624
+ metavar="FILE",
625
+ help='JSONL file of {"text": ..., "answer": ...} rows, each text written the way the'
626
+ " site sees its input: 'user: ...' lines for a chat site, the serialised state for Jev",
627
+ )
628
+ importer.add_argument("--kind", choices=_KINDS, help="what the answers are, choice by default")
629
+ importer.add_argument("--schema", metavar="JSON", help="JSON schema of the answered field")
630
+ importer.add_argument(
631
+ "--labels", metavar="A,B,C", help="answers a choice site may carry, instead of --schema"
632
+ )
633
+ importer.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
634
+ importer.set_defaults(handler=_import)
635
+
636
+ config = commands.add_parser("config", help="create or inspect the settings file")
637
+ config_commands = config.add_subparsers(dest="config_command", required=True)
638
+
639
+ init = config_commands.add_parser("init", help="write a commented settings file")
640
+ init.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
641
+ init.add_argument("--force", action="store_true", help="overwrite an existing file")
642
+ init.set_defaults(handler=_config_init)
643
+
644
+ show = config_commands.add_parser("show", help="print every setting with its source")
645
+ show.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
646
+ show.add_argument("--upstream", metavar="URL", help="provider base URL to forward to")
647
+ show.add_argument("--port", type=int, metavar="N", help="port to listen on")
648
+ show.set_defaults(handler=_config_show)
649
+ return parser
650
+
651
+
652
+ def main(argv: list[str] | None = None) -> int:
653
+ """Parses one command line and runs the command it names, returning the exit code."""
654
+ args = _build_parser().parse_args(argv)
655
+ handler: Callable[[argparse.Namespace], int] = args.handler
656
+ try:
657
+ return handler(args)
658
+ except _LearningOff as exc:
659
+ print(f"stuntd: {exc}", file=sys.stderr)
660
+ return 1
661
+ # A rejected setting, an unwritable path, a database that will not open and a failed training
662
+ # run are all things the user has to fix, so they are reported rather than raised as a traceback.
663
+ except (ValueError, OSError, RuntimeError, sqlite3.DatabaseError) as exc:
664
+ print(f"stuntd: {exc}", file=sys.stderr)
665
+ return 2
File without changes