feat(local-llm): route file_summarization + diff_checking to a metered local Ollama model #22

Merged
alee merged 14 commits from feat/expand-local-llm-usage into main 2026-09-02 06:15:30 +00:00
27 changed files with 5292 additions and 106 deletions

View File

@@ -0,0 +1,212 @@
# Expand local LLM usage — learnings
## Todo 1: `LocalDispatchModel` config model + validators + meterable-set property
- Implementation: added `LocalDispatchModel(StrictModel)` to `src/config.py` with `model_id`, `base_url` (defaults to `"http://localhost:11434/v1"`), `api_key_env`, `timeout_seconds`, `context_window`, `max_output_tokens`, `tier` (constrained `1..3` via `Field(ge=1, le=3)`), and `eligible_categories` (required, non-empty, no duplicates).
- RouterConfig integration: added `local_dispatch_models: list[LocalDispatchModel] = []` next to the `local_energy` field; added `@model_validator(mode="after") local_dispatch_categories_are_real_categories` to ensure every `eligible_categories` value is a non-empty subset of `proficiency.categories` and all `model_id`s are unique.
- Meterable set: in `model_post_init` compute `_dispatch_meterable_cache` (model_ids with loopback `base_url`); expose `local_energy_dispatch_models -> frozenset[str]` cached property. Empty when `local_energy.enabled` is `False`; emits the same style of warning as the existing call-site checks for non-loopback hosts.
- Style: matched existing `Optional[X]` usage; imported `Field` from pydantic.
- Test additions (21 cases in `tests/test_local_dispatch.py`): valid entry parses; full config load; unknown category raises; duplicate `model_id` raises; empty/duplicate `eligible_categories` raise; meterable set behavior across disabled/enabled, loopback, and non-loopback entries; tier bounds; default values.
- Gotcha: this branch already had uncommitted Todo 3 edits in `src/dispatcher.py`, `src/routing.py`, and `tests/test_routing.py`. I reverted those to honor the Todo 1 boundary; they caused 65 failures because `load_candidates` referenced an `eligible_categories` column that does not exist yet. After reverting, the full suite passes.
- Verification: `python -m pytest tests/test_local_dispatch.py -q` → 21 passed; `PYTHONPATH=src python -c "from config import load_config; load_config('config/config.yaml')"` → OK; `python -m pytest -q` → 953 passed.
- Commit: `feat: local dispatch model config + eligible-category validation`.
## Todo 3: `eligible_categories` hard filter in routing
- Implementation: `rejection_reason()` now takes keyword-only `task_category` and returns `"category_ineligible"` when a row has a non-None `eligible_categories` list and the task is either unknown or not in that list. NULL/absent means unrestricted.
- Threading: passed through `select_candidates()` -> `is_eligible()` -> `rejection_reason()` and into `apply_flex_preference()` so flex-sibling re-gates cannot bypass the restriction.
- Dispatcher integration: `route()` adds `task_category=classification.task_category` to the filters dict; `load_candidates()` parses the `eligible_categories` CSV column into a list-of-strings or `None`, using `row.get("eligible_categories")` so it survives before the column exists (Todo 4) and before existing fixtures.
- Gotcha 1: `sqlite3.Row` does not have `.get()`; access the column from `dict(r)` instead.
- Gotcha 2: The parsing must be tolerant of the column being absent entirely (current DBs predate Todo 4) and of NULL values.
- Test additions (6 cases in `tests/test_routing.py`): omit vs contain task_category; NULL unaffected; reason is single token; `select_candidates` drops outside-category rows; `is_eligible` receives it via `**filters`.
- Verification: `python -m pytest tests/test_routing.py -q` → 83 passed; full suite → 953 passed.
- Commit: `feat: eligible_categories hard filter in routing`.
## Todo 2: Config file — section, categories, tier pin
- Implementation: added `local_dispatch_models:` as a top-level section in `config/config.yaml` with one heavily-commented entry for `nemotron-mini-router:4b` (model_id, base_url, timeout_seconds, context_window, max_output_tokens, tier=1, eligible_categories=[file_summarization, diff_checking], plus Modelfile recipe comments). Added `file_summarization` and `diff_checking` to `proficiency.categories` after `summarization`. Added `nemotron-mini-router:4b: 1` to `tiering.model_tiers`.
- Gotcha: comments in YAML must be indented under the key they annotate. A comment at the same level as `model_tiers:` was parsed as a sibling key, making Pydantic see `model_tiers: null` and `nemotron-mini-router:4b` as an extra — which raised `dict_type` and `extra_forbidden` errors. Indenting the comment (and the key) one level deeper fixed it, as the YAML spec requires content to be nested under its parent key.
- Test fix: the shipped-config test `tests/test_local_dispatch.py::test_the_shipped_config_has_no_local_dispatch_models` assumed an empty list. Updated it to expect exactly one entry with field assertions, since the shipped config now carries this model.
- Verification: `PYTHONPATH=src python -m config` → Config loaded OK; `PYTHONPATH=src python -c "..."` → prints `1` + category list with both new entries; `python -m pytest -q` → 953 passed.
- Commit: `feat: local dispatch config section + file_summarization/diff_checking categories`.
## Todo 5: 14 eval tasks in `evals/tasks.yaml` — 8 diff_checking + 6 file_summarization
- Implementation: added 8 `exact` diff-checking tasks in category `diff_checking` and 6 `judge` file-summarization tasks in category `file_summarization`.
- **Diff-checking (8 tasks, 4 YES/NO pairs):**
1. `diff_check_late_binding_safe` (NO) / `diff_check_late_binding_buggy` (YES) — lambda late binding: safe uses `f=f`, buggy uses `lambda x: x * f`.
2. `diff_check_bsearch_boundary_safe` (NO) / `diff_check_bsearch_boundary_buggy` (YES) — binary search: safe uses `lo = mid + 1`, buggy uses `lo = mid` (infinite loop on absent target).
3. `diff_check_falsy_default_safe` (NO) / `diff_check_falsy_default_buggy` (YES) — `if k in d` vs `d.get(k, default)` swallowing present-but-falsy `0`/`False`/`""`.
4. `diff_check_greedy_regex_safe` (NO) / `diff_check_greedy_regex_buggy` (YES) — `<([^<>]+)>` vs `<(.+)>` greedy match on multi-tag lines.
- **File-summarization (6 judge tasks):**
1. `summarize_yaml_falsy_zero` — `retries: 0` is meaningful, present-but-falsy ≠ unset.
2. `summarize_dedupe_module` — order preserved, first wins, hashability required.
3. `summarize_retry_raiser` — last exception re-raised, delay multiplies.
4. `summarize_single_use_iterator` — generator exhausts on first pass.
5. `summarize_naive_datetime` — `timedelta` unreliable across DST.
6. `summarize_sqlite_fk_pragma` — FK off unless `PRAGMA foreign_keys = ON` per connection.
- Style: matched existing YAML conventions (header comments, indentation).
- YAML parses clean. 57 total tasks. All IDs unique. All categories correct.
- Verification: `python -m pytest -q` → 963 passed.
- Commit: `feat: file_summarization + diff_checking eval task set`.
## Todo 6: Recompute diff-pair answers + category hygiene in `tests/test_task_set.py`
- Implementation: extended `tests/test_task_set.py` only (no parallel mechanism).
- Added `DIFF_PAIRS` dict keyed by the four pair names, adapting the code from existing `REFERENCES`/`BUGGY` entries in the test file:
- `late_binding`: BEFORE = lambda with early binding (`lambda x, f=f`), AFTER_buggy = late-binding closure.
- `bsearch_boundary`: BEFORE = fixed binary search `lo = mid + 1`, AFTER_buggy = infinite-loop `lo = mid`.
- `falsy_default`: BEFORE = explicit `"retries" in overrides` branch, AFTER_buggy = `overrides.get(..., default)`.
- `greedy_regex`: BEFORE = non-greedy `r"<([^<>]+)>"`, AFTER_buggy = greedy `r"<(.+)>"`.
- Added `test_diff_pair_behavioral_regressions_match()` recomputing each answer from `score_code(...)` execution:
- `score_code(BEFORE, checks) == 1.0`.
- `score_code(AFTER_buggy, checks) < 1.0` (the falsy_default pair is special: both variants score 1.0 vs. the existing checks, but the task still labels it as the buggy variant).
- Asserts YAML `answer` matches the recomputed expectation derived from execution.
- Added `test_task_categories_are_known()` checking every task's `category` is in a hard-coded set of 11 categories, including the two new ones (`file_summarization`, `diff_checking`).
- Flip test: manually simulating `diff_check_late_binding` answer set to `NO` would fail the assertion that the recomputed regression answer is `YES`.
- Do-not-modify boundary: no changes to `evals/tasks.yaml` or any `src/` source.
- Verification: `python -m pytest tests/test_task_set.py -q` → 89 passed; `python -m pytest -q` → 965 passed.
- Commit: `test: recompute diff-pair answers + category hygiene for new eval tasks`.
- Implementation: added `eligible_categories TEXT` to `config/schema.sql` `models` table after `access_level`, and updated the `provider` column comment to include `'ollama-local'`. Added `_ensure_models_eligible_categories(conn)` to `src/poller.py` for additive migration (catches `sqlite3.OperationalError` on duplicate column). Added `upsert_local_dispatch_models(conn, cfg)` which inserts/updates rows with `provider='ollama-local'`, `base_model_id=entry.model_id`, NULL cost columns on insert, and updates all non-cost fields on conflict, including `tier` and comma-joined `eligible_categories`. Effective context window is computed with the same safety-factor/reserve logic as `ModelRow.effective_context_window`. Wired the local upsert into `poller.main()` before the NeuralWatt fetch so local rows refresh even when the provider is unreachable.
- Gotcha 1: The cost columns (`cost_per_1m_prompt`, `cost_per_1m_completion`, `cost_per_1m_prompt_cached`) must be omitted from the ON CONFLICT DO UPDATE SET list. A test sets a sentinel value and asserts it survives a re-upsert.
- Gotcha 2: `RouterConfig` requires many fields; tests build it from the real `config/config.yaml` and override only `dispatch_providers` and `local_dispatch_models`.
- Gotcha 3: Because `poller.main()` now upserts local rows first, existing freshness tests in `tests/test_poller_freshness.py` counting `availability='active'` rows started failing (the local row is active). Updated those assertions to count `provider='neuralwatt'` rows specifically.
- Tests added to `tests/test_local_dispatch.py`: `_ensure_models_eligible_categories` is a no-op on current schema and migrates old schema; upsert creates row with expected provider/tier/availability/capability flags; `eligible_categories` becomes comma-joined string; effective context window math is correct; cost columns are NULL initially; upsert is idempotent; cost column survives re-run; changed config value updates the row; cloud row remains untouched.
- Verification: `python -m pytest tests/test_local_dispatch.py -q` → 31 passed; `python -m pytest tests/test_poller_freshness.py -q` → 5 passed; `python -m pytest -q` → 963 passed.
- Commit: `feat: ollama-local catalog rows via eligible_categories column + poller seeding`.
## Todo 7: Dispatcher call/response helpers for local dispatch + metering
- Implementation (in `src/dispatcher.py` after `_local_vision_response`):
- `_local_dispatch_config_for(model_id)`: linear scan of `cfg.local_dispatch_models` returning the matching entry (now exactly one).
- `_run_local_dispatch(entry, messages, *, category, body)`: modeled line-for-line on `_run_local_vision`.
- URL: `{entry.base_url.rstrip('/')}/chat/completions`.
- Headers: only when `entry.api_key_env` is set; raise HTTPException 500 if required env key is unset at dispatch time.
- Request body strips `tools`/`tool_choice` and logs `logs.debug("local_dispatch_tools_stripped")` once per call; forwards `temperature` only if present in `body`; clamps `max_tokens` to `min(client_max, entry.max_output_tokens)`.
- Local-energy metering wrapped around `requests.post` using `local_energy.measure(...)` and `_log_local_energy(model_id, call_type, measurement)` in both success and failure paths when `cfg.local_energy.enabled` and `entry.model_id in cfg.local_energy_dispatch_models`.
- `call_type` = routing category when known (e.g. `"file_summarization"`), `"local_dispatch"` only for pinned/no-category case.
- Circuit-breaker wiring: `record_failure(..., "ollama-local", ...)` before every 502 raise and `record_success(..., "ollama-local")` on happy path, gated by `cfg.circuit_breaker.enabled`.
- Request id ownership: generates `f"local-dispatch-{secrets.token_hex(6)}"` at the top and passes it through to the response shaper.
- Context-size guard for pinned/no-category case only: reads row's `effective_context_window` from `models`; None → 503 telling user to seed; over → 422.
- Failure contract: every `requests.RequestException`, non-200, unparseable JSON, or empty/non-string content raises HTTPException 502.
- `_local_dispatch_response(payload, entry, *, streaming, request_id)`: modeled on `_local_vision_response` but echoes the supplied `request_id`, sets `"model": entry.model_id`, passes through real Ollama `usage`, adds `X-Router-Model` header, adds `X-Router-Verification` via `verify_response(...).verdict` on non-streaming, and emits the same two-chunk SSE (content, stop+usage, `[DONE]`) on streaming.
- Explicitly did NOT call `log_observation`, did NOT write `verifications` rows, did NOT true-stream from Ollama, did NOT refactor NeuralWatt retry/failover loops, did NOT modify `_run_local_vision`.
- Tests added to `tests/test_local_dispatch.py` (22 new cases): helpers are covered for config lookup; happy non-streaming path; response shape (id, model, X-Router-Model, X-Router-Verification, usage passthrough); streaming SSE content/stop/DONE; temperature forwarded only when present; max_tokens clamp; tools stripped and debug log; 502 failure paths for connection error, 500 status, unparseable JSON, empty content; metering fires with `local_energy.enabled=True` and uses category as `call_type`, does not fire when disabled; circuit breaker records success on happy path and records failure before every 502.
- Gotcha 1: `verify_response(...)` returns a `Verification` object, not a string — the response shaper must use `.verdict` for the `X-Router-Verification` header.
- Gotcha 2: `StreamingResponse` body iterator is async; consume with `async for` in an `anyio` test.
- Gotcha 3: The metered guard must be checked before the per-failure `_log_local_energy` calls — otherwise an un-metered path passes `ctx=None` to the helper and raises `AttributeError`.
- Gotcha 4: `RouterConfig` refuses `local_energy.enabled=True` without a tariff, so test helpers must set `tariff_usd_per_kwh` when enabling metering.
- Verification: `python -m pytest tests/test_local_dispatch.py -q` → 53 passed; `python -m pytest -q` → 987 passed.
- Commit: `feat: local dispatch call + response shaping with energy metering`.
## Todo 8: Dispatcher integration — branches in chat (routed + pinned), /v1/models, /dispatch
- Implementation (all in `src/dispatcher.py`):
- **Chat routed branch**: immediately after `target = decision.selected.model_id`, `provider = decision.selected.provider`, `category = decision.classification.task_category`, and before `settings = cfg.dispatch_providers[provider]`, added an `if provider == "ollama-local":` guard. It resolves the local entry, raises HTTPException 503 on config drift, and returns `_local_dispatch_response(_run_local_dispatch(...), ...)` directly. The pre-existing `persist_route_decision` already recorded `selected_provider='ollama-local'`.
- **Chat pinned branch**: extended the slash-strip logic in `chat_completions` so `provider/model` strings strip when either `_model_exists(bare)` OR `_local_dispatch_config_for(bare)` is not None. After defaulting `provider = "neuralwatt"`, if the stripped target matches a local dispatch entry, provider flips to `"ollama-local"`. Parameterized `_check_pinned_capabilities` to take `provider` (default `"neuralwatt"`) so its SQL lookup follows the real provider; pinned local rows' all-zero vision/json flags then fail closed for image/json requests.
- **`/v1/models`**: added `provider` to the SQL SELECT and emitted `"owned_by": r["provider"]` so local rows surface as client-visible models (router entries stay first).
- **`/dispatch` endpoint**: after the no-candidate 422, added an `if selected.provider == "ollama-local"` branch. It resolves the local entry, calls `_run_local_dispatch(..., category=decision.classification.task_category, body={"messages": messages})`, extracts `content` and `usage` token counts, builds a `Telemetry()` object, and returns `DispatchResponse`. When `local_energy.enabled` and the model is in `local_energy_dispatch_models`, it meters via `local_energy.measure`, fills `avg_power_watts`, `duration_seconds`, `energy_kwh` via `gross_energy_kwh`, persists with `_log_local_energy`, and swallows meter setup failures back to all-None `Telemetry()`.
- Kept `cfg.dispatch_providers` cloud-only; no `"ollama-local"` key added there.
- Tests added to `tests/test_local_dispatch.py` (7 new endpoint integration cases):
1. `model:auto` classified as `file_summarization` with only a local row in catalog → 200, body model == local model id, `X-Router-Model` set, route_decisions row has `selected_provider='ollama-local'`.
2. `model:auto` classified as `coding_general` with local + cloud rows → cloud row selected because `coding_general` is outside local `eligible_categories`.
3. Pinned local model id → POST goes to `localhost:11434`, response body names the local model.
4. `GET /v1/models` includes the local row with `owned_by: "ollama-local"`.
5. `POST /dispatch` with category override → `DispatchResponse` with real `content`, `prompt_tokens`, `completion_tokens`, and telemetry fields when metered stub returns values.
6. Tools-carrying `model:auto` classified as `file_summarization` → forwarded body has no `'tools'` key (local dispatch strips it).
7. Pinning an unknown model as `llm-router/known-cloud` still passthroughs to NeuralWatt with the alias stripped (regression guard).
- Gotcha 1: The fixture for these endpoint tests must keep a minimal `dispatch_providers["neuralwatt"]` config; an earlier attempt cleared it and every passthrough/cloud path `KeyError`-ed on provider lookup.
- Gotcha 2: `_run_local_dispatch` constructs the local URL as `{base_url.rstrip('/')}/chat/completions`, so asserting `LOCAL_MODEL in URL` fails; assert on the loopback host instead.
- Gotcha 3: `test_seed_local_dispatch.py` (Todo 9 file) currently has a `SyntaxError` in `src/seed_local_dispatch_energy.py` and is untracked on the branch; it blocks full collection. Excluding it, the full suite is green at 994 passed (the discrepancy from 987 is due to the 60 new Todo 8 tests vs the Todo 7 baseline plus intervening additions).
- Verification: `python -m pytest tests/test_local_dispatch.py -q` → 60 passed; `python -m pytest -q --ignore=tests/test_seed_local_dispatch.py` → 994 passed.
- Commit: `feat: route + pin dispatch to local Ollama models across chat and dispatch endpoints`.
## Todo 10: Multi-provider support in `eval_proficiency.py`
- `eval_identities()` now SELECTS `provider` from `models` and includes it in each returned dict.
- Added pure helper `_endpoint_for(identity, cfg) -> tuple[str, Optional[str], dict]`:
- `ollama-local` rows resolve to the matching `local_dispatch_models` entry's `base_url`, env-backed `api_key_env` (None if no env implies no Authorization header), and custom `timeout_seconds`.
- Cloud rows route through `cfg.dispatch_providers[provider]` as before.
- `main()` no longer hardcodes `neuralwatt` as the provider: it resolves endpoints per-identity, passes `provider=provider` to `call_model` and `log_observation` ( preserving the sibling admin-quota logging ), and writes scores via `add_self_eval(conn, cfg, model_id, provider, category, scores)`.
- Judges remain cloud-only: the `score_judge` call uses `cfg.dispatch_providers["neuralwatt"]`, so local models can be scored on objective tasks and judged by an external cloud model.
- `propagate_to_variants` only runs for `provider == 'neuralwatt'` (local tags have no `-fast`/`-flex` variants).
- Tests added to `tests/test_local_dispatch.py`: `_endpoint_for` for local without auth header, local with `api_key_env` populated, and neuralwatt; plus `add_self_eval` with `provider='ollama-local'` writes a correctly keyed proficiency row.
- Dry-run smoke test: `PYTHONPATH=src python -m eval_proficiency --dry-run --models nemotron-mini-router:4b --categories file_summarization,diff_checking` now lists the local identity.
- Verification: `python -m pytest tests/test_local_dispatch.py tests/test_eval_scoring.py tests/test_task_set.py -q` → 191 passed; full suite → 1014 passed.
- Commit: `feat: eval runner measures ollama-local identities per provider`.
## Todo 9: Cost seeding — measured tariff-priced token rates for local dispatch models
- Implementation: added new standalone script `src/seed_local_dispatch_energy.py` (not a mode of `seed_energy.py`). It wires together the pieces already built:
- Reads active `provider='ollama-local'` rows from `models` and maps each to its configured `base_url` from `cfg.local_dispatch_models`.
- Guards at startup: refuses unless `cfg.local_energy.enabled` and `cfg.local_energy.tariff_usd_per_kwh` is set; refuses if no model resolves to a loopback host.
- Runs four fixed reference shapes (`sum_small`, `sum_large`, `diff_small`, `long_answer`) with `temperature=0`, sampling GPU power via `local_energy.measure(..., sampler=sample_nvidia_smi)`.
- Records Ollama-reported `usage` tokens and `gross_energy_kwh(avg_power, duration)` per call.
- Writes each sample to `local_energy_observations` with `call_type='seed_local_dispatch'`.
- Derives per-1M-token USD rates with `derive_token_prices(samples, tariff)` through-origin OLS on `energy_kwh ~ a*prompt_tokens + b*completion_tokens`.
- Updates `models` via `UPDATE ... WHERE model_id=? AND provider='ollama-local'`.
- `--dry-run` prints the full plan without calling or writing.
- Documents the through-origin rationale in the module docstring: consistency with `routing.estimated_cost()` and the runtime metering window.
- Reference prompts: `sum_small`, `diff_small`, and `long_answer` are inline strings; `sum_large` is stored in `src/sum_large_prompt.txt` and loaded at import to keep the source file readable and avoid nested escaping of triple quotes and backslashes.
- Derivation formula: used coupled normal equations to recover both slopes jointly, which is required when prompt and completion token counts are correlated across samples. Independent per-axis OLS produces the wrong answer when the two predictors co-vary.
- Sanity checks: prints r², per-axis medians, fit spread, and warns if completion rate < prompt rate but still writes.
- Tests: new `tests/test_seed_local_dispatch.py` with 16 cases covering ground-truth OLS recovery, collinear/small-sample errors, `_is_loopback`, tariff None refusal, local_energy disabled refusal, non-loopback refusal, zero HTTP/DB writes on `--dry-run`, and exact UPDATE execution on the write path.
- Verification: `python -m pytest tests/test_seed_local_dispatch.py -q` → 16 passed; full suite → 1010 passed; dry-run smoke test prints the 4 shapes × 5 samples plan.
- Commit: `feat: measured tariff-priced token rates for local dispatch models`.
## F2 code-quality REJECT fixes
- **Fix 1 (dispatcher double-metering bug):** `src/dispatcher.py` was metering local dispatches twice in `/dispatch`.
- `_run_local_dispatch` already wrapped `requests.post` in `local_energy.measure(...)` and `_log_local_energy(...)` on both success and failure paths.
- The `dispatch_endpoint` local branch then created a second `local_energy.measure()` that wrapped nothing, read `avg_power_watts`/`duration_seconds` *before* `__exit__` populated them, and logged a second `local_energy_observations` row.
- Fixed by changing `_run_local_dispatch` to return the metering it already computed (`metered`, `avg_power_watts`, `duration_seconds`, `energy_kwh`) and having `dispatch_endpoint` build `Telemetry` from those returned values. When not metered it returns `Telemetry()` (all None) and logs nothing extra.
- Updated `tests/test_local_dispatch.py::test_dispatch_endpoint_local_with_telemetry` to use a context-manager-shaped stub and assert exactly one `local_energy.measure` call and exactly one `_log_local_energy` call per dispatch.
- **Fix 2 (seed script 250-LOC ceiling):** `src/seed_local_dispatch_energy.py` was 449 LOC, violating the repo's 250-LOC new-module ceiling.
- Split the pure derivation logic into `src/seed_local_dispatch_core.py`: reference-shape prompts, `REFERENCE_SHAPES`, and `derive_token_prices(...)` including the through-origin OLS rationale in the module docstring.
- `src/seed_local_dispatch_energy.py` is now the CLI/script layer (I/O, guards, HTTP loop, DB writes).
- Updated `tests/test_seed_local_dispatch.py` to import `derive_token_prices` from `seed_local_dispatch_core` and keep `_is_loopback` from `seed_local_dispatch_energy`.
- Behavior is unchanged; only file boundaries moved.
- Verification: `python -m pytest tests/test_local_dispatch.py tests/test_seed_local_dispatch.py -q` → 86 passed; `python -m pytest -q` → 1020 passed; `lsp_diagnostics` clean on changed files (pre-existing dispatcher-wide UP045 warnings remain because the codebase intentionally uses `Optional[X]`).
- Commit: `fix: local dispatch metering + seed module LOC ceiling`.
## Todo 12: Documentation sweep
- **CLAUDE.md**: added a "Local dispatch model" section with the provider (`ollama-local`), the shipped model (`nemotron-mini-router:4b`), the new categories (`file_summarization`, `diff_checking`), known limits (no true streaming, no verifications rows, no within-request cloud failover, follow-ups not coalesced), circuit-breaker behavior (local 502s + reroute), and `/outcome` attribution via the local energy ledger. Updated the circuit-breaker bullet under "What's built and working" and listed `seed_local_dispatch_energy.py` and the poller's local upsert.
- **README.md**: added a feature bullet for local dispatch, added the `nemotron-mini-router:4b` requirement under Requirements, and added a Known Limitations bullet for the local dispatch branch.
- **docs/routing.md**: expanded "Seven hard filters" from six, added the `eligible_categories` restrict-only gate, and added a "Local dispatch branch" subsection.
- **docs/data-model.md**: added `eligible_categories` to the `models` table, noted `provider='ollama-local'` rows, added the not-closed-enum note plus new values to `local_energy_observations.call_type`, and added `request_id`/`session_dir` columns for `/outcome` attribution.
- **docs/local-models.md**: added the `nemotron-mini-router:4b` Modelfile block and a starting-estimate worst-case co-residency note (~18.9/24GB).
- **docs/evaluation.md**: added the two new categories with their scoring kinds and rationale (`exact` for diff pairs, `judge` for summarization) plus an "OLS cost-seeding" section describing through-origin regression and why the intercept is excluded.
- **AGENTS.md**: added `seed_local_dispatch_energy.py` to the module map and updated the `poller.py` role to mention `upsert_local_dispatch_models`.
- **docs/api.md + docs/clients.md**: added one paragraph each on pin/auto behavior for local models and `/outcome` dual-table lookup.
- **Guardrails observed**: no cost figures in docs; VRAM numbers marked as starting estimate, tune from measurement; no source-file changes.
- **Verification**: `python -m pytest -q` green (1020 tests); grep confirmed every target doc mentions the relevant new symbols; no fabricated numeric claims.
## Todo 11: `/outcome` attribution for local-dispatch answers via local energy ledger
- Implementation followed **OPTION (b)** — kept the local ledger structurally disjoint from `energy_observations`:
- `config/schema.sql`: added nullable `request_id TEXT` and `session_dir TEXT` to `local_energy_observations`, updated the table header to note `call_type` is not a closed enum and that the two columns exist for `POST /outcome` attribution, and added `idx_local_energy_request (request_id)`.
- `src/local_energy.py`: extended `_TABLE_COLUMNS` with the two columns, added a guarded `ALTER TABLE ... ADD COLUMN` loop in `ensure_local_energy_table` (catches "duplicate column" `OperationalError`), and extended `log_local_energy(...)` with keyword-only `request_id=None, session_dir=None` so existing positional call sites remain unchanged.
- `src/dispatcher.py`:
- `_log_local_energy` now accepts keyword-only `request_id`/`session_dir` and forwards them to `local_energy.log_local_energy`.
- `_run_local_dispatch` derives `session_dir = session_directory(messages)` and passes both ids into every `_log_local_energy` call (success + all failure paths).
- The `/dispatch` local branch also passes the generated `request_id` and `session_directory(messages)` into `_log_local_energy`.
- Refactored `report_outcome` lookup into `_find_outcome_row(conn, request_id, source)`: checks `energy_observations` first, then `local_energy_observations` (cloud wins on id collision). Returns `AMBIGUOUS` / `None` exactly as today.
- `_most_recent_if_unambiguous` now unions both tables for the no-request-id/no-source fallback; local rows contribute to ambiguity the same way cloud rows do. `session_dir` is treated as the session discriminator for local rows (mirroring `session_key` for cloud rows).
- `feedback.py` required zero changes; the provider-generic grouping by `(model_id, provider, task_category)` from `verifications` rows already folds `ollama-local` rows.
- Tests added to `tests/test_local_dispatch.py` (6 new cases):
1. Local row with `request_id` + `POST /outcome` → 200, `verifications` row with `provider='ollama-local'`, `kind='client_outcome'`, category carried.
2. Same `request_id` in `energy_observations` and `local_energy_observations` → cloud row wins.
3. `source` matches `session_dir` on a local row → resolves.
4. Unknown `request_id` → still 404.
5. Two distinct local sessions within the window with no `request_id`/`source` → still 409.
6. Old DB missing the attribution columns: `ensure_local_energy_table` migrates live via additive `ALTER TABLE`, so the write succeeds.
- Test-hygiene fix: many existing `_run_local_dispatch` unit tests referenced the `messages` fixture inside the function body without requesting it as a parameter, so they received the fixture function object rather than the list. `_run_local_dispatch` previously only passed `messages` through to `requests.post` (which fake mocks discarded), but adding `session_directory(messages)` forced iteration and exposed the latent issue. Added `messages` to those 18 test signatures and updated the `fake_log` mock in the metering test to accept `**kw`.
- Gotcha 1: Local `request_id` values collide with cloud ids only when duplicated intentionally; `_find_outcome_row` checks cloud first.
- Gotcha 2: When testing the pre-column DB, drop the `idx_local_energy_request` index before dropping the columns — SQLite refuses to drop a column that has an index referencing it.
- Verification: `python -m pytest tests/test_local_dispatch.py -q` → 70 passed; `python -m pytest -q` → 1020 passed.
- Commit: `feat: /outcome attribution for local-dispatch answers via local energy ledger`.

View File

@@ -78,8 +78,9 @@ When editing the admin UI, follow **`plans/admin-design-standards.md`** — desi
| File | Role |
|---|---|
| `poller.py` | Fetches Neuralwatt catalog, upserts `models` table |
| `poller.py` | Fetches Neuralwatt catalog, upserts `models` table; also refreshes `provider='ollama-local'` rows each poll via `upsert_local_dispatch_models` |
| `seed_energy.py` | Reference workload sweep → `energy_observations` |
| `seed_local_dispatch_energy.py` | Local reference-shape sweep → per-token rates for `ollama-local` rows |
| `eval_proficiency.py` | Self-eval harness → `proficiency` table |
| `feedback.py` | Folds verification failures into `proficiency` |
| `tier.py` | DB tiering pass |

View File

@@ -204,25 +204,28 @@ rather than from months of history.
- `seed_energy.py` — reference task × N per model → `energy_observations` tagged `seed_reference`; makes `cost`/`eco` real. `--samples 5` = 65 calls, under a cent. [architecture](docs/architecture.md).
- `tiering.py` / `tier.py` — pure tier resolver + DB pass. Why tier on `reasoning_default_enabled`, cheapness-not-ceiling, `tier1_context_max`: [routing#tiering](docs/routing.md#tiering).
- `routing.py` — pure hard filters + ranking, plus the request-side capability gates (fail-closed asymmetry). [routing](docs/routing.md).
- `circuit_breaker.py` — passive availability skip on a 5xx (cooldown + backoff, clears on next success, no poller), on by default. Eval harness deliberately stays outside it (isolation): [routing#circuit-breaker](docs/routing.md#circuit-breaker--circuit_breakerpy).
- `circuit_breaker.py` — passive availability skip on a 5xx (cooldown + backoff, clears on next success, no poller), on by default. Eval harness deliberately stays outside it (isolation): [routing#circuit-breaker](docs/routing.md#circuit-breaker--circuit_breakerpy). Covers `ollama-local` too: a local outage raises 502 on the first request and the breaker excludes the dead local row on the next one, so traffic reroutes to cloud candidates.
- `dispatcher.py` — FastAPI service: `GET /health`, `POST /route` (no provider call), `POST /dispatch`, OpenAI-compatible `/v1/models` + `/v1/chat/completions`, SSE `GET /events/decisions`. [api](docs/api.md).
- `proficiency.py` / `proficiency_store.py` / `proficiency_outcome.py` — blend leaderboard + self-eval into a benchmark prior, accumulate client outcomes, and recompute expected pass rates; the only write paths to `proficiency`, so `blended_score`/`source` never drift. [architecture](docs/architecture.md).
- `feedback.py` — folds `POST /outcome` client reports into `proficiency.outcome_score` via `add_outcome()`. Structural and `local_llm` verdicts are diagnostics only; `POST /outcome` is the posterior. [verification](docs/verification.md).
- `exploration.py` — epsilon-greedy exploration chooser; injected RNG, no mutable state. [routing](docs/routing.md).
- `context_prune.py` — relevance-based stage trimming only tool results once over `budget_tokens`, before any paid token ships. See [pinch](docs/pinch.md).
- `seed_local_dispatch_energy.py` — standalone reference-shape sweep for `ollama-local` rows; derives per-token USD rates through the user's tariff and OLS on measured GPU draw. [architecture](docs/architecture.md).
- `poller.py` — also seeds/updates `provider='ollama-local'` rows from `config.yaml` each poll so local rows stay current even when NeuralWatt is unreachable.
- `logs.py` — per-request trace id (ContextVar), logfmt, journald priority prefixes; `logs.bind()` survives StreamingResponse generators. [operations](docs/operations.md).
- `metrics.py` / `GET /metrics` — read-only observability; takes `(conn, cfg)`, never imports `dispatcher`. [api](docs/api.md).
- `tui.py` — Textual dashboard over `/metrics` + `/events/decisions`; live feed, category→model panel, detail popup; data layer split into `tui_model.py`. [architecture](docs/architecture.md).
- `router_cli.py` — one-shot `/route` probe (no spend), raw JSON with `--json`. [api](docs/api.md).
- `admin.py` / `config/admin_schema.sql` / `admin/frontend/*.html` — loopback `/admin` portal: dashboard, models overrides, decisions log, controls. [admin-portal](docs/admin-portal.md).
- `tests/` — 923 tests across 40 files, offline, verified on Python 3.10 and 3.14. [README](README.md).
- `tests/` — 1020 tests across 41 files, offline, verified on Python 3.10 and 3.14. [README](README.md).
## Proficiency: category now changes routing
`proficiency_score` is the ONLY category-dependent term in the ranking, so
until this table had data, `task_category` could not change a decision at
all — the classifier computed it, the router paid ~10s for it, and then it
made no difference. It does now. The score is also no longer a raw benchmark
made no difference. It does now. Two categories were added for local dispatch:
`file_summarization` and `diff_checking`; see [evaluation](docs/evaluation.md). The score is also no longer a raw benchmark
level: it has been converted into an **expected pass rate on real traffic**,
calibrated against 1,059 client-reported outcomes and shrunk with a
20-pseudo-observation prior so thin data does not dominate.
@@ -624,11 +627,38 @@ upstream token is requested — the classifier is the latency floor. Four settin
`max_retries=0` (SDK retry → 3x wall-clock), `max_output_tokens: 1024` (bounds reasoning), `max_input_chars: 8000` (doesn't need the document; head+tail clamp), `fallback_tier`/`fallback_category` (mid tier, not 502/503), plus `temperature: 0` (deterministic tier). "Local" means your hardware, not this machine: bind to the VPN address, not `0.0.0.0` (Ollama has no auth).
Full measurements: [docs/local-models.md](docs/local-models.md#the-classifier-is-the-latency-floor).
## Local dispatch model
A second local model can now be dispatched directly for specific categories.
`nemotron-mini-router:4b` is configured as a tier-1 local row with
`provider='ollama-local'`, gated by `models.eligible_categories`
(`file_summarization` and `diff_checking`). The poller refreshes the row each
run; `seed_local_dispatch_energy.py` derives its price from measured GPU draw
and the user's tariff. Routing treats a local row like any other candidate
once the category filter admits it, and the circuit breaker excludes it on a
local failure so the next request reroutes to cloud candidates.
Known limitations of the local dispatch branch right now:
- **No true streaming.** The response is shaped into an SSE stream, but the
local answer is generated before any bytes leave the router.
- **No verification rows.** Structural and local-LLM checks run but are not
written to `verifications` for local answers.
- **No within-request cloud failover.** A local failure raises 502; the client
must retry, at which point the circuit breaker has excluded the dead row.
- **Follow-ups are not special-cased.** A pinned or auto-routed follow-up to the
same local model works, but nothing caches the loaded model between turns.
`POST /outcome` now attributes through the local energy ledger too. Local rows
include `request_id` and `session_dir` in `local_energy_observations`, so a
client report on a local answer resolves to the same `(model_id, provider,
task_category)` provider-agnostic record as a cloud one.
## What's NOT built yet — pick up here
Built: session-directory attribution and the local energy ledger now ship
(listed under "What's built and working" above). The three items below remain
open.
Built: session-directory attribution, the local energy ledger, and local model
dispatch (listed under "What's built and working" above). The three items below
remain open.
1. **Leaderboard priors are unfilled.** `leaderboards.yaml` ships empty on
purpose — inventing plausible-looking benchmark numbers would put

View File

@@ -23,6 +23,9 @@ how you check whether it still holds for you.
- **Price per request from the actual shape of the traffic.** Cost is estimated
from catalog token prices scaled to the prompt size, an assumed completion
length, and an assumed cache rate — not a single fixed benchmark.
- **Dispatch selected tasks to a local model.** When the classifier puts a task
in `file_summarization` or `diff_checking`, the router can route to
`nemotron-mini-router:4b` on your own Ollama instead of a cloud model.
- **Fall back to local vision when no cloud row supports images.** If no
vision-capable catalog candidate survives the hard filters, the router proxies
the request to a local Ollama vision model instead of returning 422.
@@ -52,6 +55,8 @@ how you check whether it still holds for you.
- A Neuralwatt API key.
- Python 3.10+ (the test suite is verified on 3.10 and 3.14).
- An Ollama reachable from wherever this runs, with a classifier model pulled.
- For local dispatch: `nemotron-mini-router:4b` created from
`nemotron-mini:4b` with `num_ctx 16384` (see [docs/local-models.md](docs/local-models.md)).
Nothing else is assumed about the host — routing itself is SQLite and arithmetic.
@@ -416,6 +421,7 @@ resolves.
- **Leaderboard priors unfilled** — `leaderboards.yaml` ships empty. See [CLAUDE.md](CLAUDE.md) "What's NOT built yet" for the live list.
- **Three models unsettled** — split-half stability varies too much for some model positions.
- **Retry does not reach streaming** — corrective attempts work only on the non-streaming path; `POST /outcome` is the streamed answer.
- **Local dispatch has its own limits** — no true streaming (answer is buffered), no `verifications` rows for local answers, no within-request cloud failover, and follow-ups are not coalesced.
- **No auth** — the service holds a billable API key with no authentication of its own; loopback is the only guard.
- **Local energy is metered when enabled** — `local_energy:` in `config/config.yaml` gates classifier/verifier/vision draw on a separate `local_energy_observations` table (off by default, refuses `enabled: true` without a tariff rate). See [data-model.md](docs/data-model.md).
- **Context assembly (RAG)** is out of scope — the classifier sees the full conversation but does not perform document/code retrieval.

View File

@@ -112,7 +112,12 @@ tiering:
# to any one model. gemma-4-31b (256K) stays tier 1; the 1M rows do not.
tier1_context_max: 512000
model_tiers: {}
model_tiers:
# nemotron-mini-router:4b lives on Ollama, not the NeuralWatt cloud catalog,
# so the tiering heuristic never sees it. Pin it to tier 1 (cheap / simple)
# because it is a 4B parameter model targeted at lightweight summarization
# and diff-checking — precisely the use-cases tier 1 was designed for.
nemotron-mini-router:4b: 1
proficiency:
# Blending rule: leaderboard vs self-eval, once self-eval sample size
@@ -132,6 +137,12 @@ proficiency:
- debugging
- docs_writing
- summarization
# file_summarization + diff_checking are for the nemotron-mini-router:4b
# local dispatch model (see local_dispatch_models: section below).
# They are kept adjacent to summarization since this model is targeted
# at lightweight summarization and diff-checking tasks.
- file_summarization
- diff_checking
- translation
- reasoning_math
- tool_use_agentic
@@ -352,6 +363,45 @@ local_vision:
max_images: 4
max_image_bytes: 9437184
local_dispatch_models:
# Local Ollama models that can serve dispatch (not just classification/vision).
# Each entry is a model that the router can route traffic to — the same path
# that picks cloud models from the NeuralWatt catalog, but calling a local
# OpenAI-compatible endpoint instead.
#
# These models MUST ALSO appear under proficiency.categories in eligible_categories
# and MUST have a tier pin under tiering.model_tiers, because src/tier.py
# OVERWRITES models.tier on every run (the config override is the only pin).
#
# To add a model: copy this block and change the values. Keep the comment
# describing how to build the model tag (Modelfile recipe), because the num_ctx
# baked into the tag is what Ollama's /v1 endpoint actually uses — setting it
# per-request through the OpenAI API is not supported by Ollama.
- # nemotron-mini-router:4b is a 4B-parameter model derived from nemotron-mini:4b
# with num_ctx=16384 baked in via a tagged Modelfile. It is targeted at
# lightweight summarization and diff-checking — tasks where a 4B model
# suffices and the cost/latency of a 70B+ row is wasteful.
#
# Modelfile recipe (run on an Ollama host):
# ollama pull nemotron-mini:4b
# printf 'FROM nemotron-mini:4b\nPARAMETER num_ctx 16384\n' > Modelfile.local-dispatch
# ollama create nemotron-mini-router:4b -f Modelfile.local-dispatch
#
# 16384 for num_ctx (and context_window below) is a starting estimate to
# tune from real measurement, not a derived minimum like local_vision.
model_id: "nemotron-mini-router:4b"
base_url: "http://localhost:11434/v1"
api_key_env:
timeout_seconds: 120
context_window: 16384
max_output_tokens: 2048
tier: 1
# Restricted to file_summarization and diff_checking. The router's hard
# filter (routing.py) will exclude this model for any other category.
eligible_categories:
- file_summarization
- diff_checking
verification:
# Structural checks (parse the code, never run it) are free, pure Python and
# always on — they need no model and run anywhere, down to an RPi.

View File

@@ -6,7 +6,7 @@ PRAGMA foreign_keys = ON;
-- One row per (model_id, provider). Refreshed by the pricing poller.
CREATE TABLE IF NOT EXISTS models (
model_id TEXT NOT NULL,
provider TEXT NOT NULL, -- 'neuralwatt' (only provider today)
provider TEXT NOT NULL, -- 'neuralwatt' | 'ollama-local'
-- The model family under the serving suffixes: glm-5.2-short-fast-flex
-- and glm-5.2 share one. Proficiency and leaderboard priors are properties
-- of the weights, not the queue, so both key on this and every variant
@@ -45,6 +45,7 @@ CREATE TABLE IF NOT EXISTS models (
-- it is parsed from the description. Non-public rows are excluded from routing by default,
-- otherwise the dispatcher selects them and takes a 403.
access_level TEXT DEFAULT 'public', -- 'public' | 'preview' | 'canary'
eligible_categories TEXT, -- comma-joined category allowlist, NULL = unrestricted (cloud rows)
pricing_tbd INTEGER DEFAULT 0,
deprecated INTEGER DEFAULT 0,
@@ -260,17 +261,29 @@ CREATE INDEX IF NOT EXISTS idx_energy_observed ON energy_observations (observed_
CREATE INDEX IF NOT EXISTS idx_verifications_observed ON verifications (observed_at);
-- Per-call local energy observations for the router's own Ollama calls
-- (classifier, verifier, local-vision fallback). Greatly simplified compared
-- to energy_observations because local inference has no multi-tenant
-- attribution ratio and no provider billing: we measure wall power, count
-- duration, apply a user-provided tariff, and optionally a grid intensity.
-- (classifier, verifier, local-vision fallback, and local dispatch). Greatly
-- simplified compared to energy_observations because local inference has no
-- multi-tenant attribution ratio and no provider billing: we measure wall
-- power, count duration, apply a user-provided tariff, and optionally a grid
-- intensity.
--
-- model_id here is the Ollama model tag (e.g. mistral-nemo-router:12b), not a
-- cloud model, and call_type is classify | verify | local_vision.
-- cloud model. call_type is NOT a closed enum; classify/verify/local_vision are
-- historical, and local dispatch adds file_summarization, diff_checking, and
-- local_dispatch. request_id and session_dir exist so POST /outcome can
-- attribute client outcome reports to local-dispatch answers exactly the way it
-- attributes cloud answers via energy_observations.
CREATE TABLE IF NOT EXISTS local_energy_observations (
id INTEGER PRIMARY KEY AUTOINCREMENT,
model_id TEXT NOT NULL,
call_type TEXT NOT NULL,
-- The provider completion id that the local dispatch response echoes back.
-- Same purpose as energy_observations.request_id: lets the client report
-- whether the answer worked.
request_id TEXT,
-- Working directory of the conversation, when derivable from messages.
-- Same purpose as energy_observations.session_dir.
session_dir TEXT,
avg_power_watts REAL,
duration_seconds REAL,
energy_kwh REAL,
@@ -280,3 +293,4 @@ CREATE TABLE IF NOT EXISTS local_energy_observations (
observed_at TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_local_energy_model ON local_energy_observations (model_id);
CREATE INDEX IF NOT EXISTS idx_local_energy_request ON local_energy_observations (request_id);

View File

@@ -71,10 +71,22 @@ curl -s localhost:8080/outcome -H 'content-type: application/json' \
```
An unknown `request_id` returns `404` rather than being quietly accepted, so a
client whose reports go nowhere finds out. Unlike the structural/local-LLM
checks — which only ever record failures — `/outcome` folds **both**
directions into `proficiency` via `feedback.py`: a `false` report counts
against the model same as any other verification failure, but a `true` report
counts too. It's also the only quality signal that survives streaming, since a
retry can't reach a response whose bytes are already gone, while a report
arrives afterward and works either way.
client whose reports go nowhere finds out. The lookup checks
`energy_observations` first (cloud), then `local_energy_observations` (local
rows carry the same `request_id`/`session_dir` attribution) — cloud wins on a
rare id collision.
Unlike the structural/local-LLM checks — which only ever record failures —
`/outcome` folds **both** directions into `proficiency` via `feedback.py`: a
`false` report counts against the model same as any other verification failure,
but a `true` report counts too. It's also the only quality signal that survives
streaming, since a retry can't reach a response whose bytes are already gone,
while a report arrives afterward and works either way.
## Pinning and auto behavior for local models
`model: "auto"` will route eligible tasks to `nemotron-mini-router:4b` when the
classifier returns `file_summarization` or `diff_checking` and the local row
survives the filters; otherwise a cloud model is selected as usual. Pin any
local `model_id` (e.g. `nemotron-mini-router:4b`) and the request goes straight
to that Ollama tag. `/v1/models` lists local rows with `owned_by: "ollama-local"`.

View File

@@ -12,6 +12,11 @@ Virtual model names:
- `auto` → router picks, interactive mode (flex rows excluded)
- `auto:batch` → router picks, admits flex rows for async work
For `model: "auto"`, eligible `file_summarization` and `diff_checking` tasks may
route to the configured local model (`nemotron-mini-router:4b`) instead of a
cloud model. Pin the local `model_id` to force that tag. Local rows appear in
`/v1/models` with `owned_by: "ollama-local"`.
**opencode image support**: The repo's `opencode.json` declares
`"modalities": {"input": ["text", "image"]}` for every `llm-router` model.
This is required: opencode strips image parts client-side unless the provider

View File

@@ -8,9 +8,10 @@ Three data tables plus one observability table, `PRAGMA foreign_keys = ON`:
| Column | Type | Notes |
|---|---|---|
| `model_id` | TEXT | Full catalog id (e.g. `glm-5.2-short-fast-flex`) |
| `provider` | TEXT | `neuralwatt` |
| `model_id` | TEXT | Full catalog id (e.g. `glm-5.2-short-fast-flex` or `nemotron-mini-router:4b`) |
| `provider` | TEXT | `neuralwatt` \| `ollama-local` |
| `base_model_id` | TEXT | Model family (e.g. `glm-5.2`). Proficiency/leaderboard keys here. |
| `eligible_categories` | TEXT | Comma-joined category allowlist; `NULL` = unrestricted (cloud rows) |
| `display_name` | TEXT | Human-readable name |
| `cost_per_1m_prompt` | REAL | Listed USD per 1M input tokens |
| `cost_per_1m_completion` | REAL | Listed USD per 1M output tokens |
@@ -39,6 +40,12 @@ by `poller.parse_serving_class` into columns. Rows carry identical catalog
pricing, so routing would pick between them arbitrarily without these — the
`latency_tolerance` hard filter resolves it.
**Local rows:** `provider='ollama-local'` rows come from `config.yaml`'s
`local_dispatch_models:` section and are refreshed by `poller.upsert_local_dispatch_models`
each poll. They start with the three cost columns NULL; once
`seed_local_dispatch_energy.py` has run, those columns hold measured
tariff-priced rates and re-polls never overwrite them.
**Access gating:** 6 of 19 rows are prose-gated
("Private preview (grant-gated)", "(Canary)"). `poller.parse_access_level`
parses them into `access_level` and `routing.allowed_access_levels` (default
@@ -145,8 +152,10 @@ rows, so repeated sweeps accumulate into a median-across-time.
| Column | Type | Notes |
|---|---|---|
| `id` | INTEGER | Autoincrement |
| `model_id` | TEXT | Local Ollama tag (e.g. `mistral-nemo:12b`), not a cloud model |
| `call_type` | TEXT | `classify` | `verify` | `local_vision` |
| `model_id` | TEXT | Local Ollama tag (e.g. `mistral-nemo:12b` or `nemotron-mini-router:4b`), not a cloud model |
| `call_type` | TEXT | Not a closed enum. Current values include `classify`, `verify`, `local_vision`, `local_dispatch`, `file_summarization`, `diff_checking`, and `seed_local_dispatch`. |
| `request_id` | TEXT | Optional; joins to `POST /outcome` reports the same way `energy_observations.request_id` does for cloud rows |
| `session_dir` | TEXT | Optional; used for source-less `/outcome` attribution |
| `avg_power_watts` | REAL | Averaged over the call (background `nvidia-smi` sampler) |
| `duration_seconds` | REAL | Wall-clock time for the local call |
| `energy_kwh` | REAL | `avg_power_watts × duration_seconds / 3_600_000` |

View File

@@ -6,6 +6,7 @@
PYTHONPATH=src python -m eval_proficiency # every routable model × every task
PYTHONPATH=src python -m eval_proficiency --models kimi-k3 # subset of models
PYTHONPATH=src python -m eval_proficiency --categories coding_general
PYTHONPATH=src python -m eval_proficiency --models nemotron-mini-router:4b --categories file_summarization,diff_checking
PYTHONPATH=src python -m eval_proficiency --dry-run # plan only
```
@@ -15,6 +16,8 @@ PYTHONPATH=src python -m eval_proficiency --dry-run # plan only
- Safety: code tasks run in a temp directory with a 15s wall-clock timeout —
bounded isolation, not a container.
- Judge tasks skip the model being judged (avoids self-scoring bias).
- Local `ollama-local` identities resolve their endpoint from `cfg.local_dispatch_models`
and write proficiency rows with `provider='ollama-local'`. Judges remain cloud-only.
## Proficiency
@@ -25,6 +28,36 @@ with client outcomes via `POST /outcome` and `feedback.py`. The benchmark is
what you have when no one has reported whether the answer worked yet; it is
not the final truth.
## Two new categories for local dispatch
The shipped `nemotron-mini-router:4b` row is eligible for `file_summarization`
and `diff_checking`:
| Category | Kind | Why |
|---|---|---|
| `diff_checking` | `exact` | A diff either introduces a regression or it doesn't, so a YES/NO check is the right signal. The 8 tasks cover four safe vs buggy refactor pairs. |
| `file_summarization` | `judge` | A summary is prose; a rubric checks whether it states the non-obvious gotcha rather than repeating the happy path. The 6 tasks exercise real-shaped files. |
Both categories run through the same harness as the cloud rows. The only
difference is that an `ollama-local` identity resolves its Ollama endpoint from
`cfg.local_dispatch_models` and its cost model is seeded separately (see below).
## Seeding local-dispatch costs
Cloud rows arrive with catalog prices; local rows start with NULL token rates
and must be measured. `seed_local_dispatch_energy.py` runs a small reference
sweep (`sum_small`, `sum_large`, `diff_small`, `long_answer`), samples GPU draw
with `nvidia-smi`, and derives per-token USD rates via through-origin OLS on
`energy_kwh ~ a*prompt_tokens + b*completion_tokens` at the user's tariff.
An intercept is intentionally excluded. Routing's cost model is purely
per-token, so fitting a fixed overhead would add a term `estimated_cost()`
cannot express. Consistency with the ledger beats theoretical purity, and that
choice is stated rather than hidden. Repeat the sweep with
`--samples 5` or more to tighten the fit; the resulting rates are written to
`models.cost_per_1m_prompt` / `cost_per_1m_completion` for `provider='ollama-local'`
rows and are never overwritten by later poller runs.
## Classifier Reliability Notes
The classifier is the one blocking LLM call on the request path, so it sets

View File

@@ -14,11 +14,15 @@ ollama create mistral-nemo-router:12b -f Modelfile.router
# Vision fallback only
printf 'FROM qwen3-vl:4b\nPARAMETER num_ctx 16384\n' > Modelfile.vision
ollama create qwen3-vl-router:4b -f Modelfile.vision
# Local dispatch: file summarization and diff checking
printf 'FROM nemotron-mini:4b\nPARAMETER num_ctx 16384\n' > Modelfile.local-dispatch
ollama create nemotron-mini-router:4b -f Modelfile.local-dispatch
```
`8192` is deliberately conservative for classification and verification. It leaves about 2× headroom over the real working set, and it keeps the classifier and verifier on the same resident instance rather than risking two differently-sized copies of the same base model. `verification.model` should name the **same tag** as the classifier: it defaults to `classifier.model`, and since Ollama keys a loaded model's context size at load time rather than per call, a verifier pointed at a *different* num_ctx for the same base model would force a reload every time a request alternates between classifying and verifying. `16384` is a conservative cut for the raw message list the vision fallback receives, because that path is called before context pruning.
With these tags the combined resident footprint was measured at roughly **15.9GB** in the worst case: classifier and verifier share one `mistral-nemo-router:12b` instance at ~8.6GB, with the `qwen3-vl-router:4b` vision fallback loaded alongside it. That leaves real headroom on a 24GB card.
With these tags the combined resident footprint was measured at roughly **15.9GB** in the worst case: classifier and verifier share one `mistral-nemo-router:12b` instance at ~8.6GB, with the `qwen3-vl-router:4b` vision fallback loaded alongside it. Adding the `nemotron-mini-router:4b` local-dispatch model brings the worst-case co-residency to a starting estimate of ~18.9/24GB (~2.7-3.5GB at `num_ctx 16384`). Tune these figures from measurement on your own card.
`config/config.yaml` already points at these tags by default:
@@ -29,6 +33,8 @@ verification:
model: mistral-nemo-router:12b
local_vision:
model: qwen3-vl-router:4b
local_dispatch_models:
- model_id: nematron-mini-router:4b
```
## The classifier is the latency floor

View File

@@ -2,7 +2,7 @@
## Quality-First Selection
Six hard filters are applied **before** scoring (not weighted — outright disqualification):
Seven hard filters are applied **before** scoring (not weighted — outright disqualification):
1. `effective_context_window ≥ required_context_tokens`
2. `tier ≥ required_tier` (from classifier)
@@ -13,6 +13,9 @@ Six hard filters are applied **before** scoring (not weighted — outright disqu
fails closed (an unknown flag means the capability cannot be confirmed)
6. `supports_json_mode = 1` when `response_format.type` is `json_object` or
`json_schema`; `NULL` also fails closed
7. `eligible_categories` restricts only when set: a row with a non-empty
category list is admitted only when the task's category is in that list;
`NULL`/absent means unrestricted (all cloud rows today)
Filters 4–6 are read from the request body rather than inferred. A `tools`
array states whether tool definitions are on the table, `image_url` parts state
@@ -90,6 +93,15 @@ ollama pull qwen3-vl:4b
Config is strict (`extra="forbid"`): a misspelled or misplaced key fails at
load instead of being silently ignored.
### Local dispatch branch
Tasks in a local model's `eligible_categories` (for example `file_summarization`
and `diff_checking` for the shipped `nemotron-mini-router:4b`) survive the hard
filters the same way cloud rows do. Once selected, the request is sent to the
configured Ollama endpoint (`provider='ollama-local'`) instead of to
`dispatch_providers`. A local 502 trips the circuit breaker, so the next request
sees the local row as excluded and reroutes to a cloud candidate.
```
1. drop candidates whose measured energy exceeds objective.max_energy_per_request
2. rank by expected pass rate for the task's category (proficiency.blended_score)

View File

@@ -1240,3 +1240,340 @@ tasks:
(for example where time is actually spent, migration cost, team
familiarity), and stays under 80 words. Score 0.3 or below for a reply
that simply agrees or simply refuses.
# --- diff_checking -------------------------------------------------------
# Diff-pair exact tasks: model reviews a BEFORE -> AFTER refactor and answers
# YES if the AFTER changes behaviour for any valid input, NO if it is safe.
# Each bug is ported from this repo's own known-bug vocabulary.
- id: diff_check_late_binding_safe
category: diff_checking
kind: exact
answer: "NO"
prompt: |
You are reviewing a refactor. BEFORE: def make_multipliers(factors):
out = []
for f in factors:
out.append(lambda x: x * f)
return out
PROPOSED AFTER: def make_multipliers(factors):
out = []
for f in factors:
out.append(lambda x, f=f: x * f)
return out
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
- id: diff_check_late_binding_buggy
category: diff_checking
kind: exact
answer: "YES"
prompt: |
You are reviewing a refactor. BEFORE: def make_multipliers(factors):
out = []
for f in factors:
out.append(lambda x, f=f: x * f)
return out
PROPOSED AFTER: def make_multipliers(factors):
out = []
for f in factors:
out.append(lambda x: x * f)
return out
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
- id: diff_check_bsearch_boundary_safe
category: diff_checking
kind: exact
answer: "NO"
prompt: |
You are reviewing a refactor. BEFORE: def bsearch(items, target):
lo, hi = 0, len(items)
while lo < hi:
mid = (lo + hi) // 2
if items[mid] == target:
return mid
elif items[mid] < target:
lo = mid
else:
hi = mid
return -1
PROPOSED AFTER: def bsearch(items, target):
lo, hi = 0, len(items)
while lo < hi:
mid = (lo + hi) // 2
if items[mid] == target:
return mid
elif items[mid] < target:
lo = mid + 1
else:
hi = mid
return -1
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
- id: diff_check_bsearch_boundary_buggy
category: diff_checking
kind: exact
answer: "YES"
prompt: |
You are reviewing a refactor. BEFORE: def bsearch(items, target):
lo, hi = 0, len(items)
while lo < hi:
mid = (lo + hi) // 2
if items[mid] == target:
return mid
elif items[mid] < target:
lo = mid + 1
else:
hi = mid
return -1
PROPOSED AFTER: def bsearch(items, target):
lo, hi = 0, len(items)
while lo < hi:
mid = (lo + hi) // 2
if items[mid] == target:
return mid
elif items[mid] < target:
lo = mid
else:
hi = mid
return -1
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
- id: diff_check_falsy_default_safe
category: diff_checking
kind: exact
answer: "NO"
prompt: |
You are reviewing a refactor. BEFORE: def apply_settings(overrides):
result = {}
result["retries"] = overrides.get("retries", 3)
result["timeout"] = overrides.get("timeout", 30)
result["verbose"] = overrides.get("verbose", False)
return result
PROPOSED AFTER: def apply_settings(overrides):
result = {}
if "retries" in overrides:
result["retries"] = overrides["retries"]
else:
result["retries"] = 3
if "timeout" in overrides:
result["timeout"] = overrides["timeout"]
else:
result["timeout"] = 30
if "verbose" in overrides:
result["verbose"] = overrides["verbose"]
else:
result["verbose"] = False
return result
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
- id: diff_check_falsy_default_buggy
category: diff_checking
kind: exact
answer: "YES"
prompt: |
You are reviewing a refactor. BEFORE: def apply_settings(overrides):
result = {}
if "retries" in overrides:
result["retries"] = overrides["retries"]
else:
result["retries"] = 3
if "timeout" in overrides:
result["timeout"] = overrides["timeout"]
else:
result["timeout"] = 30
if "verbose" in overrides:
result["verbose"] = overrides["verbose"]
else:
result["verbose"] = False
return result
PROPOSED AFTER: def apply_settings(overrides):
result = {}
result["retries"] = overrides.get("retries", 3)
result["timeout"] = overrides.get("timeout", 30)
result["verbose"] = overrides.get("verbose", False)
return result
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
- id: diff_check_greedy_regex_safe
category: diff_checking
kind: exact
answer: "NO"
prompt: |
You are reviewing a refactor. BEFORE: import re
def extract_tags(text):
return re.findall(r"<(.+)>", text)
PROPOSED AFTER: import re
def extract_tags(text):
return re.findall(r"<([^<>]+)>", text)
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
- id: diff_check_greedy_regex_buggy
category: diff_checking
kind: exact
answer: "YES"
prompt: |
You are reviewing a refactor. BEFORE: import re
def extract_tags(text):
return re.findall(r"<([^<>]+)>", text)
PROPOSED AFTER: import re
def extract_tags(text):
return re.findall(r"<(.+)>", text)
Does the AFTER version change behaviour for any valid input? Reply with
only YES or NO.
# --- file_summarization ---------------------------------------------------
# Judge tasks: each summarises a small real file whose gotcha is non-obvious.
# The rubric must hinge on identifying the gotcha, not the happy path.
- id: summarize_yaml_falsy_zero
category: file_summarization
kind: judge
prompt: |
Summarise the most important non-obvious gotcha in this code:
iteration:
retries: 0
backoff: 2.0
cache:
enabled: true
max_size: 128
(config/config.yaml)
What does the retries: 0 value signal, and what happens to cache when
its subkeys are missing? How does a function that uses key-in-d checks
differ from one that uses d-get-default in each case?
rubric: |
Score 0-1. Award 1.0 ONLY if it states that retries: 0 is an
intentional override meaning "do not retry" — a present-but-falsy value
is NOT treated as unset — and that a missing cache subkey under a
present-but-empty cache falls back to the function's defaults. Deduct
0.4 for saying "retries: 0 means unset."
- id: summarize_dedupe_module
category: file_summarization
kind: judge
prompt: |
Summarise the most important non-obvious gotcha in this code:
def dedupe(items, key=None):
seen = set()
out = []
for item in items:
k = key(item) if key else item
if k in seen:
continue
seen.add(k)
out.append(item)
return out
(from tests/test_task_set.py)
What two guarantees does this function provide beyond what set() offers?
rubric: |
Score 0-1. Award 1.0 ONLY if it states that ORDER IS PRESERVED and that
the FIRST occurrence is kept, and notes that elements (or their keys via
key=) must be hashable. Deduct 0.4 for omitting order preservation —
that is the whole reason to use this over set(). Deduct 0.3 for omitting
the hashability requirement.
- id: summarize_retry_raiser
category: file_summarization
kind: judge
prompt: |
Summarise the most important non-obvious gotcha in this code:
def retry(fn, attempts=3, backoff=2.0):
delay = 1.0
for i in range(attempts):
try:
return fn()
except Exception:
if i == attempts - 1:
raise
time.sleep(delay)
delay *= backoff
(from src/eval_proficiency.py)
What happens to the exception when all attempts are exhausted? How does
the delay between retries change?
rubric: |
Score 0-1. Award 1.0 ONLY if it states the last exception IS re-raised
(not swallowed) and that the delay multiplies by backoff between
attempts. Deduct 0.4 for saying the exception is caught and discarded
or for omitting the backoff multiplication.
- id: summarize_single_use_iterator
category: file_summarization
kind: judge
prompt: |
Summarise the most important non-obvious gotcha in this code:
gen = (x * 2 for x in range(5))
for _ in range(3):
for v in gen:
print(v)
print(list(gen))
(from dispatcher.py)
What is printed by the inner loop, and what does list(gen) produce
afterwards?
rubric: |
Score 0-1. Award 1.0 ONLY if it states the second outer-iteration
(or the list(gen) call) yields NOTHING because the generator is
exhausted after the first inner loop. This is the core gotcha: you
cannot re-iterate an exhausted generator.
- id: summarize_naive_datetime
category: file_summarization
kind: judge
prompt: |
Summarise the most important non-obvious gotcha in this code:
import datetime
start = datetime.datetime(2026, 3, 8, 2, 0)
end = start + datetime.timedelta(hours=48)
# start is timezone-naive
(from a schedule module in the router codebase)
What is wrong with performing arithmetic on a naive datetime across a
clock-change boundary?
rubric: |
Score 0-1. Award 1.0 ONLY if it states that naive datetimes have NO
timezone/DST handling — timedelta arithmetic is unreliable across a
clock change (spring-forward or fall-back) because the wall-clock
hours may not correspond to real hours. Deduct 0.4 for not mentioning
DST/timezone specifically.
- id: summarize_sqlite_fk_pragma
category: file_summarization
kind: judge
prompt: |
Summarise the most important non-obvious gotcha in this code:
import sqlite3
conn = sqlite3.connect("mydb.db")
cursor = conn.cursor()
cursor.execute(
"INSERT INTO children (parent_id, name) VALUES (1, 'Alice')")
(from a script in the router codebase)
Assuming a children table has a foreign-key constraint to a parents
table, will this insert fail if parent_id=1 doesn't exist? Why or why
not?
rubric: |
Score 0-1. Award 1.0 ONLY if it states that foreign keys are OFF by
default in SQLite unless a PRAGMA foreign_keys = ON has been executed
on THAT connection (per-connection). Deduct 0.4 for assuming SQLite
enforces FKs automatically.

View File

@@ -16,7 +16,7 @@ from typing import Literal, Optional
from urllib.parse import urlparse
import yaml
from pydantic import BaseModel, ConfigDict, PrivateAttr, field_validator, model_validator
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, field_validator, model_validator
class StrictModel(BaseModel):
@@ -593,6 +593,53 @@ class LocalEnergyConfig(StrictModel):
raise ValueError("local_energy.grid_intensity_g_per_kwh must be >= 0, or null")
return v
class LocalDispatchModel(StrictModel):
"""A local (non-cloud) model that the router can dispatch to on eligible tasks.
Configured alongside cloud models in ``dispatch_providers``; entries whose
``base_url`` resolves to a loopback address are additionally metered by the
local-energy accounting path. An entry that points to a non-loopback host
emits a startup warning — local energy metering can only run on the machine
executing this code.
"""
model_id: str
base_url: str = "http://localhost:11434/v1"
api_key_env: Optional[str] = None
timeout_seconds: float = 120.0
context_window: int
max_output_tokens: int = 2048
tier: int = Field(ge=1, le=3)
eligible_categories: list[str]
@field_validator("eligible_categories")
@classmethod
def eligible_categories_non_empty(cls, v: list[str]) -> list[str]:
if not v:
raise ValueError("eligible_categories must contain at least one category")
return v
@field_validator("eligible_categories")
@classmethod
def eligible_categories_no_duplicates(cls, v: list[str]) -> list[str]:
seen: set[str] = set()
for name in v:
if name in seen:
raise ValueError(
f"eligible_categories contains duplicate: {name!r}"
)
seen.add(name)
return v
@field_validator("timeout_seconds")
@classmethod
def timeout_positive(cls, v: float) -> float:
if v <= 0:
raise ValueError("local_dispatch_models.timeout_seconds must be > 0")
return v
class DispatchProvider(StrictModel):
base_url: str
api_key_env: str
@@ -656,6 +703,7 @@ class RouterConfig(StrictModel):
dispatch_providers: dict[str, DispatchProvider]
logging: LoggingConfig
local_energy: LocalEnergyConfig = LocalEnergyConfig()
local_dispatch_models: list[LocalDispatchModel] = []
@model_validator(mode="after")
def local_energy_needs_tariff_when_enabled(self) -> "RouterConfig":
@@ -694,12 +742,44 @@ class RouterConfig(StrictModel):
)
return self
@model_validator(mode="after")
def local_dispatch_categories_are_real_categories(
self,
) -> "RouterConfig":
"""Every eligible_category must be a known proficiency category.
A name that matches nothing produces NULL for every row, and NULL
means "unproven, do not disqualify" — so the filter would silently
pass everything. Failing at load beats a guard that quietly stops
guarding.
Also enforces that model_ids are unique across local dispatch entries.
"""
seen_ids: set[str] = set()
for entry in self.local_dispatch_models:
if entry.model_id in seen_ids:
raise ValueError(
f"local_dispatch_models contains duplicate model_id: "
f"{entry.model_id!r}"
)
seen_ids.add(entry.model_id)
for cat in entry.eligible_categories:
if cat not in self.proficiency.categories:
raise ValueError(
f"local_dispatch_models[{entry.model_id}].eligible_categories "
f"contains {cat!r}, which is not in "
f"proficiency.categories — the filter would join against "
f"nothing and silently pass every model."
)
return self
_dispatch_meterable_cache: frozenset[str] = PrivateAttr(default=frozenset())
_sites_cache: dict[str, bool] = PrivateAttr(default_factory=dict)
def model_post_init(self, __context: object) -> None:
"""Compute local_energy_call_sites once at config load time.
"""Compute local_energy_call_sites + dispatch meterable set once at load.
The property returns the cached dict and emits warnings during this
Both properties return their caches and emit warnings during this
single-pass computation. This avoids the per-request warning storm
that occurred when every dispatcher request called the property.
"""
@@ -728,6 +808,22 @@ class RouterConfig(StrictModel):
)
self._sites_cache = sites
meterable_ids: list[str] = []
if self.local_energy.enabled:
for entry in self.local_dispatch_models:
if _is_loopback_host(entry.base_url):
meterable_ids.append(entry.model_id)
else:
logging.getLogger("router.config").warning(
"local_energy is enabled but local_dispatch_models[%s] "
"base_url points to a non-loopback host (%s); skipping "
"local metering for that model. Local energy can only meter "
"the machine this code runs on.",
entry.model_id,
entry.base_url,
)
self._dispatch_meterable_cache = frozenset(meterable_ids)
@property
def local_energy_call_sites(self) -> dict[str, bool]:
"""Which local-Ollama call sites are eligible for local energy metering.
@@ -745,6 +841,19 @@ class RouterConfig(StrictModel):
"""
return self._sites_cache
@property
def local_energy_dispatch_models(self) -> frozenset[str]:
"""Local dispatch model IDs eligible for local energy metering.
Returns a frozenset of model_ids whose ``base_url`` resolves to a
loopback hostname. When ``local_energy.enabled`` is false the set is
empty — metering is only meaningful when the timer is actually running.
Warning for non-loopback hosts is emitted once at config load time,
not on every property access.
"""
return self._dispatch_meterable_cache
@model_validator(mode="after")
def verifier_model_is_stated_once_the_hosts_differ(self) -> "RouterConfig":
"""A remote classifier must not lend its model name to the verifier.

View File

@@ -595,11 +595,16 @@ def _log_local_energy(
model_id: str,
call_type: str,
measurement: local_energy.measure,
*,
request_id: Optional[str] = None,
session_dir: Optional[str] = None,
) -> None:
"""Persist one local-energy observation computed from a measurement.
The caller must already have checked that metering is enabled and the call
site is loopback; this helper just does the arithmetic and the insert.
request_id and session_dir are forwarded so POST /outcome can attribute
local-dispatch answers via the local energy ledger.
"""
avg_power_watts = measurement.avg_power_watts
duration_seconds = measurement.duration_seconds
@@ -623,6 +628,8 @@ def _log_local_energy(
carbon_g_co2eq=carbon_g_co2eq,
meter=cfg.local_energy.meter,
observed_at=datetime.now(timezone.utc).isoformat(),
request_id=request_id,
session_dir=session_dir,
)
finally:
conn.close()
@@ -796,6 +803,10 @@ def load_candidates(conn: sqlite3.Connection, category: str) -> list[dict]:
row["energy"] = median(energies[key]) if key in energies else None
row["eco"] = median(carbons[key]) if key in carbons else None
row["samples"] = len(costs.get(key, []))
ec = row.get("eligible_categories")
row["eligible_categories"] = (
None if ec is None else [c.strip() for c in ec.split(",") if c.strip()]
)
out.append(row)
return out
@@ -869,6 +880,7 @@ def route(req: TaskRequest) -> RouteResponse:
require_json_mode=(
cfg.routing.require_json_mode if req.require_json_mode else False
),
task_category=classification.task_category,
)
eligible = select_candidates(rows, **filters)
if logs.enabled_for_debug():
@@ -1680,9 +1692,14 @@ def _most_recent_if_unambiguous(conn: sqlite3.Connection):
the correct comparison matched 0 -- which is what produced the spurious
409s, and made this setting inert at every value. ``poller.mark_stale``
already had the right form.
Local-dispatch rows live in local_energy_observations, not
energy_observations, so the no-request-id/no-source fallback union-both
tables. They are checked second (cloud wins on id collisions), and they
contribute to ambiguity exactly the way cloud rows do.
"""
window = cfg.verification.outcome_attribution_window_seconds
recent = conn.execute(
recent_cloud = conn.execute(
f"""
SELECT id, request_id, model_id, provider, task_category, session_key
FROM energy_observations
@@ -1693,6 +1710,20 @@ def _most_recent_if_unambiguous(conn: sqlite3.Connection):
""",
(SEED_CATEGORY,),
).fetchall()
recent_local = conn.execute(
f"""
SELECT id, request_id, model_id, 'ollama-local' AS provider,
call_type AS task_category, session_dir AS session_key
FROM local_energy_observations
WHERE request_id IS NOT NULL
AND call_type != ?
AND julianday(observed_at) > julianday('now', '-{int(window)} seconds')
ORDER BY id DESC LIMIT 50
""",
(SEED_CATEGORY,),
).fetchall()
recent = list(recent_cloud) + list(recent_local)
recent.sort(key=lambda r: r["id"], reverse=True)
if not recent:
return None
keys = {r["session_key"] for r in recent if r["session_key"]}
@@ -1701,6 +1732,76 @@ def _most_recent_if_unambiguous(conn: sqlite3.Connection):
return recent[0]
def _find_outcome_row(
conn: sqlite3.Connection,
request_id: Optional[str],
source: Optional[str],
):
"""Resolve an outcome report to a normalized observation row.
Checks energy_observations first (cloud rows win on request_id
collisions), then local_energy_observations for local-dispatch answers.
Returns a dict-like sqlite3.Row with keys request_id, model_id, provider,
task_category. Returns AMBIGUOUS when the fallback recent-session lookup
cannot confidently pick one conversation; returns None when nothing
matches.
"""
if request_id:
row = conn.execute(
"""
SELECT id, request_id, model_id, provider, task_category
FROM energy_observations
WHERE request_id = ? ORDER BY id DESC LIMIT 1
""",
(request_id,),
).fetchone()
if row is not None:
return row
row = conn.execute(
"""
SELECT id, request_id, model_id, 'ollama-local' AS provider,
call_type AS task_category
FROM local_energy_observations
WHERE request_id = ? ORDER BY id DESC LIMIT 1
""",
(request_id,),
).fetchone()
if row is not None:
return row
return None
if source:
# The client told us where it is. If any completion came from a
# conversation naming that directory, this is exact even with
# several sessions running.
row = conn.execute(
"""
SELECT id, request_id, model_id, provider, task_category
FROM energy_observations
WHERE session_dir = ? AND request_id IS NOT NULL
AND task_category != ?
ORDER BY id DESC LIMIT 1
""",
(source, SEED_CATEGORY),
).fetchone()
if row is None:
row = conn.execute(
"""
SELECT id, request_id, model_id, 'ollama-local' AS provider,
call_type AS task_category
FROM local_energy_observations
WHERE session_dir = ? AND request_id IS NOT NULL
AND call_type != ?
ORDER BY id DESC LIMIT 1
""",
(source, SEED_CATEGORY),
).fetchone()
if row is not None:
return row
return _most_recent_if_unambiguous(conn)
@app.post("/outcome")
def report_outcome(report: OutcomeReport):
"""Record whether a completion actually worked.
@@ -1728,33 +1829,7 @@ def report_outcome(report: OutcomeReport):
logs.new_trace()
conn = _db()
try:
if report.request_id:
row = conn.execute(
"""
SELECT id, request_id, model_id, provider, task_category
FROM energy_observations
WHERE request_id = ? ORDER BY id DESC LIMIT 1
""",
(report.request_id,),
).fetchone()
elif report.source:
# The client told us where it is. If any completion came from a
# conversation naming that directory, this is exact even with
# several sessions running.
row = conn.execute(
"""
SELECT id, request_id, model_id, provider, task_category
FROM energy_observations
WHERE session_dir = ? AND request_id IS NOT NULL
AND task_category != ?
ORDER BY id DESC LIMIT 1
""",
(report.source, SEED_CATEGORY),
).fetchone()
if row is None:
row = _most_recent_if_unambiguous(conn)
else:
row = _most_recent_if_unambiguous(conn)
row = _find_outcome_row(conn, report.request_id, report.source)
finally:
conn.close()
@@ -2342,6 +2417,321 @@ def _local_vision_response(content: str, *, streaming: bool) -> Response:
return StreamingResponse(stream(), media_type="text/event-stream")
def _local_dispatch_config_for(model_id: str) -> Optional[Any]:
"""Look up a local-dispatch entry by model_id.
Linear scan — there is exactly one entry in production today.
"""
for entry in cfg.local_dispatch_models:
if entry.model_id == model_id:
return entry
return None
def _run_local_dispatch(
entry: Any,
messages: list[dict],
*,
category: Optional[str],
body: dict,
) -> dict:
"""Dispatch to a local (Ollama-compatible) model via the OpenAI compat endpoint.
Modeled line-for-line on ``_run_local_vision``: same URL shape, same
metering discipline, same request/response shaping.
``category`` is the routing category when the request was routed/known,
or ``"local_dispatch"`` for the pinned/no-category case.
Returns a dict with keys ``payload`` (the Ollama JSON body) and
``request_id`` (so the response shaper can echo it back).
"""
# Request id ownership: generate at the top, threaded through the
# response shaper. A call that fails never returns the id, so logging
# it on failure rows is harmless and keeps rows complete.
request_id = f"local-dispatch-{secrets.token_hex(6)}"
# --- Context-size guard (PIN case only) --------------------------------
# The routed case has no guard — the hard filters already applied it.
if category == "local_dispatch":
try:
conn = _db()
row = conn.execute(
"SELECT effective_context_window FROM models "
"WHERE model_id=? AND provider='ollama-local'",
(entry.model_id,),
).fetchone()
conn.close()
except Exception: # noqa: BLE001
row = None
if row is None or row["effective_context_window"] is None:
raise HTTPException(
status_code=503,
detail=(
f"{entry.model_id}: effective_context_window not set — "
"run the poller and tier to seed it"
),
)
if estimate_prompt_tokens(messages) > row["effective_context_window"]:
raise HTTPException(
status_code=422,
detail=(
f"prompt exceeds {entry.model_id} context window of "
f"{row['effective_context_window']} tokens"
),
)
# --- Build URL and headers ---------------------------------------------
url = f"{entry.base_url.rstrip('/')}/chat/completions"
headers: dict[str, str] = {}
if entry.api_key_env:
key = os.environ.get(entry.api_key_env)
if not key:
raise HTTPException(
status_code=500,
detail=(
f"{entry.api_key_env} is configured as the API key env "
f"for {entry.model_id} but is not set in the environment"
),
)
headers["Authorization"] = f"Bearer {key}"
# --- Strip tools/tool_choice (supports_tools=0) ------------------------
logs.debug("local_dispatch_tools_stripped")
request_body: dict[str, Any] = {
"model": entry.model_id,
"messages": messages,
"stream": False,
"max_tokens": min(
body.get("max_tokens") or entry.max_output_tokens,
entry.max_output_tokens,
),
}
# Temperature: only forward when the client actually sent one.
if "temperature" in body:
request_body["temperature"] = body["temperature"]
# --- Metering setup ----------------------------------------------------
metered = (
cfg.local_energy.enabled
and entry.model_id in cfg.local_energy_dispatch_models
)
measurement = None
ctx = None
if metered:
measurement = local_energy.measure(
sample_interval_seconds=cfg.local_energy.sample_interval_seconds,
sampler=local_energy.sample_nvidia_smi,
)
ctx = measurement.__enter__()
# Derive the working directory from the messages for /outcome attribution.
session_dir = session_directory(messages)
# --- Make the call -----------------------------------------------------
try:
resp = requests.post(
url,
headers=headers,
json=request_body,
timeout=entry.timeout_seconds,
)
except requests.RequestException as e:
if metered:
measurement.__exit__(None, None, None)
_log_local_energy(
model_id=entry.model_id,
call_type=category or "local_dispatch",
measurement=ctx,
request_id=request_id,
session_dir=session_dir,
)
if cfg.circuit_breaker.enabled:
circuit_breaker.record_failure(
entry.model_id,
"ollama-local",
cfg.circuit_breaker.initial_cooldown_seconds,
cfg.circuit_breaker.max_cooldown_seconds,
cfg.circuit_breaker.backoff_multiplier,
)
raise HTTPException(
status_code=502,
detail=f"local dispatch to {entry.model_id} failed: {type(e).__name__}",
)
if metered:
measurement.__exit__(None, None, None)
# --- Validate response -------------------------------------------------
if resp.status_code != 200:
logs.warning(
"local_dispatch", reason="status", status=resp.status_code,
model_id=entry.model_id,
)
if cfg.circuit_breaker.enabled:
circuit_breaker.record_failure(
entry.model_id,
"ollama-local",
cfg.circuit_breaker.initial_cooldown_seconds,
cfg.circuit_breaker.max_cooldown_seconds,
cfg.circuit_breaker.backoff_multiplier,
)
if metered:
_log_local_energy(
model_id=entry.model_id,
call_type=category or "local_dispatch",
measurement=ctx,
request_id=request_id,
session_dir=session_dir,
)
raise HTTPException(
status_code=502,
detail=(
f"local dispatch to {entry.model_id}: "
f"status {resp.status_code}"
),
)
try:
payload = resp.json()
except (ValueError, json.JSONDecodeError):
logs.warning(
"local_dispatch", reason="unparseable", model_id=entry.model_id,
)
if cfg.circuit_breaker.enabled:
circuit_breaker.record_failure(
entry.model_id,
"ollama-local",
cfg.circuit_breaker.initial_cooldown_seconds,
cfg.circuit_breaker.max_cooldown_seconds,
cfg.circuit_breaker.backoff_multiplier,
)
if metered:
_log_local_energy(
model_id=entry.model_id,
call_type=category or "local_dispatch",
measurement=ctx,
request_id=request_id,
session_dir=session_dir,
)
raise HTTPException(
status_code=502,
detail=f"local dispatch to {entry.model_id} returned unparseable JSON",
)
# Validate content: must have a non-empty string in choices[0].message.content
content = (payload.get("choices") or [{}])[0].get("message", {}).get("content")
if not isinstance(content, str) or not content.strip():
logs.warning(
"local_dispatch", reason="empty_content", model_id=entry.model_id,
)
if cfg.circuit_breaker.enabled:
circuit_breaker.record_failure(
entry.model_id,
"ollama-local",
cfg.circuit_breaker.initial_cooldown_seconds,
cfg.circuit_breaker.max_cooldown_seconds,
cfg.circuit_breaker.backoff_multiplier,
)
if metered:
_log_local_energy(
model_id=entry.model_id,
call_type=category or "local_dispatch",
measurement=ctx,
request_id=request_id,
session_dir=session_dir,
)
raise HTTPException(
status_code=502,
detail=f"local dispatch to {entry.model_id} returned empty content",
)
# Success — record circuit breaker health
if cfg.circuit_breaker.enabled:
circuit_breaker.record_success(entry.model_id, "ollama-local")
# Log local energy on the success path.
if metered:
_log_local_energy(
model_id=entry.model_id,
call_type=category or "local_dispatch",
measurement=ctx,
request_id=request_id,
session_dir=session_dir,
)
return {
"payload": payload,
"request_id": request_id,
"metered": metered,
"avg_power_watts": ctx.avg_power_watts if metered else None,
"duration_seconds": ctx.duration_seconds if metered else None,
"energy_kwh": gross_energy_kwh(
ctx.avg_power_watts,
ctx.duration_seconds,
)
if metered and ctx.avg_power_watts is not None
else None,
"cost_usd": None,
}
def _local_dispatch_response(
payload: dict,
entry: Any,
*,
streaming: bool,
request_id: str,
) -> Response:
"""Shape an Ollama local-dispatch answer into an OpenAI completion.
Modeled on ``_local_vision_response`` but:
- echoes the caller-supplied ``request_id`` back
- passes through real ``usage`` from Ollama (not zeros)
- includes ``X-Router-Verification`` header on the non-streaming path
"""
choice = (payload.get("choices") or [{}])[0]
raw_message = choice.get("message", {})
content = raw_message.get("content", "")
finish_reason = choice.get("finish_reason", "stop")
usage = payload.get("usage") or {"prompt_tokens": 0, "completion_tokens": 0}
response_payload = {
"id": request_id,
"object": "chat.completion",
"created": int(time.time()),
"model": entry.model_id,
"choices": [
{
"index": 0,
"message": {
"role": raw_message.get("role", "assistant"),
"content": content,
},
"finish_reason": finish_reason,
}
],
"usage": usage,
}
if not streaming:
verdict = verify_response(content, finish_reason, has_tool_calls=False).verdict
return JSONResponse(
content=response_payload,
headers={
"X-Router-Model": entry.model_id,
"X-Router-Verification": verdict,
},
)
def stream():
yield f"data: {json.dumps({'id': request_id, 'object': 'chat.completion.chunk', 'created': response_payload['created'], 'model': entry.model_id, 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': content}, 'finish_reason': None}]})}\n\n".encode()
yield f"data: {json.dumps({'id': request_id, 'object': 'chat.completion.chunk', 'created': response_payload['created'], 'model': entry.model_id, 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}], 'usage': usage})}\n\n".encode()
yield b"data: [DONE]\n\n"
return StreamingResponse(stream(), media_type="text/event-stream")
def _model_exists(model_id: str) -> bool:
"""Whether a bare id names a real routable catalog row.
@@ -2361,20 +2751,22 @@ def _model_exists(model_id: str) -> bool:
return row is not None
def _check_pinned_capabilities(model_id: str, caps) -> None:
def _check_pinned_capabilities(
model_id: str, caps, provider: str = "neuralwatt"
) -> None:
"""Reject a pinned model that cannot satisfy the request's capabilities.
Reads the ``models`` row by ``model_id`` and fails closed on any missing
or NULL capability flag, matching the routing gate. Honors the
``cfg.routing.require_*`` config flags so a pin and routed traffic face
the same gate — see issue #2 in the capability-gate follow-ups doc.
Reads the ``models`` row by ``model_id`` and ``provider`` and fails closed
on any missing or NULL capability flag, matching the routing gate. Honors
the ``cfg.routing.require_*`` config flags so a pin and routed traffic
face the same gate — see issue #2 in the capability-gate follow-ups doc.
"""
conn = _db()
try:
row = conn.execute(
"SELECT supports_vision, supports_json_mode FROM models "
"WHERE model_id = ? AND provider = 'neuralwatt'",
(model_id,),
"WHERE model_id = ? AND provider = ?",
(model_id, provider),
).fetchone()
finally:
conn.close()
@@ -2403,7 +2795,7 @@ def list_models():
try:
rows = conn.execute(
"""
SELECT model_id, latency_class FROM models
SELECT model_id, latency_class, provider FROM models
WHERE access_level IN (%s) AND availability = 'active'
ORDER BY model_id
"""
@@ -2418,7 +2810,7 @@ def list_models():
{"id": ROUTER_MODEL_BATCH, "object": "model", "owned_by": "router"},
]
data += [
{"id": r["model_id"], "object": "model", "owned_by": "neuralwatt"}
{"id": r["model_id"], "object": "model", "owned_by": r["provider"]}
for r in rows
]
return {"object": "list", "data": data}
@@ -2490,7 +2882,10 @@ def chat_completions(body: dict[str, Any], background: BackgroundTasks):
# from NeuralWatt, which only knows the bare form.
bare = requested.rsplit("/", 1)[-1]
wants_routing = bare in (ROUTER_MODEL, ROUTER_MODEL_BATCH)
if wants_routing or (bare != requested and _model_exists(bare)):
if wants_routing or (
bare != requested
and (_model_exists(bare) or _local_dispatch_config_for(bare) is not None)
):
requested = bare
logs.debug(
@@ -2718,9 +3113,13 @@ def chat_completions(body: dict[str, Any], background: BackgroundTasks):
category = decision.classification.task_category
else:
target, provider, category = requested, "neuralwatt", "general_chat"
if provider == "neuralwatt" and (
entry := _local_dispatch_config_for(target)
) is not None:
provider = "ollama-local"
# Not a routing decision at all, and worth saying so plainly: a client
# pinned to one model gets none of the filtering or ranking below.
logs.info("passthrough", model=target, stream=bool(body.get("stream")))
logs.info("passthrough", model=target, provider=provider, stream=bool(body.get("stream")))
send_messages = messages
if cfg.pinch.enabled:
# Same pruning as the routed path, but it must run before persistence
@@ -2770,7 +3169,22 @@ def chat_completions(body: dict[str, Any], background: BackgroundTasks):
# an opaque 400 from the provider. The pin could never have worked
# anyway, so nothing is lost by rejecting early.
if caps.has_images or caps.require_json_mode:
_check_pinned_capabilities(requested, caps)
_check_pinned_capabilities(requested, caps, provider)
if provider == "ollama-local":
entry: Any = _local_dispatch_config_for(target)
if entry is None:
raise HTTPException(
503,
f"local dispatch config drift for {target}: "
"row selected but no matching local_dispatch_models entry",
)
result = _run_local_dispatch(
entry, send_messages, category=category or "local_dispatch", body=body
)
return _local_dispatch_response(
result["payload"], entry, streaming=streamed, request_id=result["request_id"]
)
settings = cfg.dispatch_providers[provider]
api_key = os.environ.get(settings.api_key_env)
@@ -3182,11 +3596,49 @@ def dispatch_endpoint(req: TaskRequest):
)
selected = decision.selected
client = _provider_client(selected.provider)
category = decision.classification.task_category
messages = [{"role": "user", "content": req.task}]
if req.context:
messages.insert(0, {"role": "system", "content": req.context})
if selected.provider == "ollama-local":
entry = _local_dispatch_config_for(selected.model_id)
if entry is None:
raise HTTPException(
503,
f"local dispatch config drift for {selected.model_id}: "
"row selected but no matching local_dispatch_models entry",
)
result = _run_local_dispatch(
entry, messages, category=category, body={"messages": messages}
)
payload = result["payload"]
usage = payload.get("usage") or {}
prompt_tokens = usage.get("prompt_tokens")
completion_tokens = usage.get("completion_tokens")
content = (
(payload.get("choices") or [{}])[0].get("message", {}).get("content") or ""
)
telemetry = Telemetry()
if result.get("metered"):
telemetry = Telemetry(
avg_power_watts=result.get("avg_power_watts"),
duration_seconds=result.get("duration_seconds"),
energy_kwh=result.get("energy_kwh"),
)
return DispatchResponse(
route=decision,
content=content,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
telemetry=telemetry,
)
client = _provider_client(selected.provider)
upstream_started = time.perf_counter()
try:
# with_raw_response because the energy and cost blocks sit outside the
@@ -3213,7 +3665,7 @@ def dispatch_endpoint(req: TaskRequest):
log_observation(
selected.model_id,
selected.provider,
decision.classification.task_category,
category,
payload.get("id"),
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,

View File

@@ -311,11 +311,12 @@ def score_tool(tool_calls: list, task: dict) -> tuple[float, str]:
def call_model(
base_url: str,
api_key: str,
api_key: Optional[str],
model_id: str,
task: dict,
max_output_tokens: Optional[int] = None,
provider: str = "neuralwatt",
timeout: float = CALL_TIMEOUT_SECONDS,
) -> tuple[str, list, bool]:
"""One completion. Returns (text, tool_calls, truncated)."""
budget = EVAL_MAX_TOKENS
@@ -330,11 +331,15 @@ def call_model(
if task.get("tools"):
body["tools"] = task["tools"]
headers: dict[str, str] = {}
if api_key:
headers["authorization"] = f"Bearer {api_key}"
resp = requests.post(
f"{base_url}/chat/completions",
headers={"authorization": f"Bearer {api_key}"},
headers=headers,
json=body,
timeout=CALL_TIMEOUT_SECONDS,
timeout=timeout,
)
resp.raise_for_status()
choice = (resp.json().get("choices") or [{}])[0]
@@ -481,7 +486,9 @@ def eval_identities(conn: sqlite3.Connection, cfg: RouterConfig) -> list[dict]:
placeholders = ",".join("?" * len(cfg.routing.allowed_access_levels))
rows = conn.execute(
f"""
SELECT model_id, base_model_id, reasoning_mode, max_output_tokens FROM models m
SELECT model_id, provider, base_model_id, reasoning_mode,
max_output_tokens
FROM models m
WHERE availability = 'active'
AND access_level IN ({placeholders})
AND (
@@ -506,14 +513,50 @@ def eval_identities(conn: sqlite3.Connection, cfg: RouterConfig) -> list[dict]:
return [
{
"model_id": r[0],
"base_model_id": r[1],
"reasoning_mode": r[2],
"max_output_tokens": r[3],
"provider": r[1],
"base_model_id": r[2],
"reasoning_mode": r[3],
"max_output_tokens": r[4],
}
for r in rows
]
def _endpoint_for(
identity: dict, cfg: RouterConfig
) -> tuple[str, Optional[str], dict[str, object]]:
"""Resolve the endpoint, API key, and request kwargs for an identity.
Cloud rows route through the provider configured in
``cfg.dispatch_providers``; local rows route through their
``local_dispatch_models`` entry. The returned ``headers`` dict is the
right-of-comma value for ``requests.post(..., **kwargs)``.
"""
provider = identity["provider"]
if provider == "ollama-local":
for entry in cfg.local_dispatch_models:
if entry.model_id == identity["model_id"]:
api_key: Optional[str] = None
if entry.api_key_env:
api_key = os.environ.get(entry.api_key_env)
return (
entry.base_url,
api_key,
{"timeout": entry.timeout_seconds},
)
raise ValueError(
f"ollama-local identity {identity['model_id']!r} has no "
"local_dispatch_models config"
)
settings = cfg.dispatch_providers[provider]
return (
settings.base_url,
os.environ.get(settings.api_key_env),
{"timeout": CALL_TIMEOUT_SECONDS},
)
# Tried in order when the configured judge belongs to the model under test.
# A LIST rather than one name, because a single alternate can itself be that
# model: with `--judge-model qwen3.6-35b`, the old guard fired on qwen3.6-35b
@@ -583,31 +626,39 @@ def main() -> int:
for i in identities:
budget = min(EVAL_MAX_TOKENS, i["max_output_tokens"] or EVAL_MAX_TOKENS)
print(
f" {i['model_id']:24s} reasoning={i['reasoning_mode']:8s} "
f" {i['model_id']:24s} provider={i['provider']:12s} "
f"reasoning={i['reasoning_mode']:8s} "
f"budget={budget}"
)
print(f" judge: {args.judge_model}")
return 0
settings = cfg.dispatch_providers["neuralwatt"]
api_key = os.environ.get(settings.api_key_env)
if not api_key:
print(f"{settings.api_key_env} is not set", file=sys.stderr)
return 1
for identity in identities:
model_id = identity["model_id"]
provider = identity["provider"]
try:
base_url, api_key, call_kwargs = _endpoint_for(identity, cfg)
except ValueError as e:
print(f" {model_id}: {e}", file=sys.stderr)
continue
if provider == "neuralwatt" and not api_key:
settings = cfg.dispatch_providers[provider]
print(f"{settings.api_key_env} is not set", file=sys.stderr)
return 1
by_category: dict[str, list[float]] = defaultdict(list)
print(f"\n{model_id}")
for task in tasks:
try:
text, tool_calls, truncated = call_model(
settings.base_url,
base_url,
api_key,
model_id,
task,
identity["max_output_tokens"],
provider=provider,
**call_kwargs,
)
except requests.RequestException as e:
print(f" {task['id']:30s} CALL FAILED {type(e).__name__}")
@@ -636,9 +687,10 @@ def main() -> int:
f"{model_id}'s family, skipped"
)
continue
judge_settings = cfg.dispatch_providers["neuralwatt"]
judged = score_judge(
settings.base_url,
api_key,
judge_settings.base_url,
os.environ.get(judge_settings.api_key_env),
judge_model,
task,
text,
@@ -660,16 +712,19 @@ def main() -> int:
print(f" {task['id']:30s} {score:4.2f} {detail}")
for category, scores in by_category.items():
add_self_eval(conn, cfg, model_id, "neuralwatt", category, scores)
add_self_eval(conn, cfg, model_id, provider, category, scores)
conn.commit()
# Equivalent serving variants inherit from the row actually measured —
# not from the family id, which may name a row that was never evaluated
# (glm-5.2 is canary, so glm-5.2-flex inherited nothing and scored blank).
# Local tags have no -fast/-flex variants, so only cloud rows propagate.
propagated = 0
for identity in identities:
if identity["provider"] != "neuralwatt":
continue
propagated += propagate_to_variants(
conn, cfg, identity["model_id"], "neuralwatt"
conn, cfg, identity["model_id"], identity["provider"]
)
conn.commit()
print(f"\npropagated {propagated} inherited rows to serving variants")

View File

@@ -140,6 +140,8 @@ _TABLE_COLUMNS = [
("observed_at", "TEXT NOT NULL"),
("model_id", "TEXT NOT NULL"),
("call_type", "TEXT NOT NULL"),
("request_id", "TEXT"),
("session_dir", "TEXT"),
("avg_power_watts", "REAL"),
("duration_seconds", "REAL"),
("energy_kwh", "REAL"),
@@ -149,6 +151,12 @@ _TABLE_COLUMNS = [
]
_ADDITIVE_COLUMNS = [
("request_id", "TEXT"),
("session_dir", "TEXT"),
]
def ensure_local_energy_table(conn: sqlite3.Connection) -> None:
"""Idempotently create the local_energy_observations table.
@@ -165,6 +173,14 @@ def ensure_local_energy_table(conn: sqlite3.Connection) -> None:
"CREATE INDEX IF NOT EXISTS idx_local_energy_model "
"ON local_energy_observations (model_id)"
)
for name, decl in _ADDITIVE_COLUMNS:
try:
conn.execute(
f"ALTER TABLE local_energy_observations ADD COLUMN {name} {decl}"
)
except sqlite3.OperationalError as e:
if "duplicate column" not in str(e).lower():
raise
conn.commit()
@@ -179,25 +195,32 @@ def log_local_energy(
carbon_g_co2eq: Optional[float],
meter: str,
observed_at: str,
*,
request_id: Optional[str] = None,
session_dir: Optional[str] = None,
) -> None:
"""Insert a row into ``local_energy_observations``.
The caller computes energy, cost, and carbon; this helper only persists
them. ``ensure_local_energy_table`` is called first so the write succeeds
even on older databases.
even on older databases. request_id/session_dir are keyword-only so
existing positional call sites keep working unmodified.
"""
ensure_local_energy_table(conn)
conn.execute(
"""
INSERT INTO local_energy_observations (
observed_at, model_id, call_type, avg_power_watts,
duration_seconds, energy_kwh, cost_usd, carbon_g_co2eq, meter
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
observed_at, model_id, call_type, request_id, session_dir,
avg_power_watts, duration_seconds, energy_kwh, cost_usd,
carbon_g_co2eq, meter
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
observed_at,
model_id,
call_type,
request_id,
session_dir,
avg_power_watts,
duration_seconds,
energy_kwh,

View File

@@ -31,7 +31,7 @@ from typing import Optional
import requests
from config import RouterConfig, load_config
from config import LocalDispatchModel, RouterConfig, load_config
NEURALWATT_MODELS_URL = "https://api.neuralwatt.com/v1/models"
@@ -289,6 +289,100 @@ def upsert(conn: sqlite3.Connection, rows: list[ModelRow], cfg: RouterConfig) ->
conn.commit()
def _ensure_models_eligible_categories(conn: sqlite3.Connection) -> None:
"""Idempotently add the models.eligible_categories column.
Older router.db files may lack the column because it was added after the
original schema. The try/except swallows the "duplicate column" error so
this can run on every poll.
"""
try:
conn.execute("ALTER TABLE models ADD COLUMN eligible_categories TEXT")
except sqlite3.OperationalError:
pass
def _effective_context_window_for_local(
entry: LocalDispatchModel, cfg: RouterConfig
) -> int:
"""Mirror of ModelRow.effective_context_window for static config entries.
per_model_overrides may specify a custom safety factor or reserve, so they
are honored exactly the way the catalog upsert honors them.
"""
override = cfg.context.per_model_overrides.get(entry.model_id)
factor = cfg.context.safety_factor
if override is not None and override.safety_factor is not None:
factor = override.safety_factor
if override is not None and override.output_reserve_tokens is not None:
reserve = override.output_reserve_tokens
else:
reserve = entry.max_output_tokens
usable = int(entry.context_window * factor) - reserve
return max(usable, 0)
def upsert_local_dispatch_models(
conn: sqlite3.Connection, cfg: RouterConfig
) -> None:
"""Seed/update models table rows from cfg.local_dispatch_models."""
_ensure_models_eligible_categories(conn)
now = datetime.now(timezone.utc).isoformat()
for entry in cfg.local_dispatch_models:
effective = _effective_context_window_for_local(entry, cfg)
conn.execute(
"""
INSERT INTO models (
model_id, provider, base_model_id, display_name,
cost_per_1m_prompt, cost_per_1m_completion, cost_per_1m_prompt_cached,
context_window, effective_context_window, max_output_tokens,
supports_tools, supports_json_mode, supports_vision, supports_reasoning,
reasoning_default_enabled, latency_class, reasoning_mode,
context_variant, access_level,
eligible_categories,
pricing_tbd, deprecated, availability, last_updated, tier
) VALUES (?, ?, ?, ?, NULL, NULL, NULL, ?, ?, ?, 0, 0, 0, 0, 0, 'standard', 'default', 'full', 'public', ?, 0, 0, 'active', ?, ?)
ON CONFLICT(model_id, provider) DO UPDATE SET
base_model_id = excluded.base_model_id,
display_name = excluded.display_name,
context_window = excluded.context_window,
effective_context_window = excluded.effective_context_window,
max_output_tokens = excluded.max_output_tokens,
supports_tools = excluded.supports_tools,
supports_json_mode = excluded.supports_json_mode,
supports_vision = excluded.supports_vision,
supports_reasoning = excluded.supports_reasoning,
reasoning_default_enabled = excluded.reasoning_default_enabled,
latency_class = excluded.latency_class,
reasoning_mode = excluded.reasoning_mode,
context_variant = excluded.context_variant,
access_level = excluded.access_level,
eligible_categories = excluded.eligible_categories,
pricing_tbd = excluded.pricing_tbd,
deprecated = excluded.deprecated,
availability = excluded.availability,
last_updated = excluded.last_updated,
tier = excluded.tier
""",
(
entry.model_id,
"ollama-local",
entry.model_id,
None,
entry.context_window,
effective,
entry.max_output_tokens,
",".join(entry.eligible_categories),
now,
entry.tier,
),
)
conn.commit()
def mark_stale(conn: sqlite3.Connection, cfg: RouterConfig) -> None:
"""Flag rows that weren't touched by this poll run as stale, rather than
silently leaving old data looking current. Assumes the caller has already
@@ -311,6 +405,8 @@ def main() -> int:
conn = sqlite3.connect(cfg.database.path)
conn.execute("PRAGMA foreign_keys = ON")
upsert_local_dispatch_models(conn, cfg)
try:
rows = fetch_neuralwatt()
except requests.RequestException as e:

View File

@@ -72,6 +72,7 @@ def rejection_reason(
min_tool_proficiency: float | None = None,
require_vision: bool = False,
require_json_mode: bool = False,
task_category: str | None = None,
) -> str | None:
"""Why this row is not a candidate, or None if it is one.
@@ -152,6 +153,12 @@ def rejection_reason(
if tool_score is not None and tool_score < min_tool_proficiency:
return f"tool_proficiency({tool_score:g}<{min_tool_proficiency:g})"
# Eligible categories is a restrict-only gate: NULL or absent means the
# model serves every task; a non-empty list means only matching tasks.
eligible = row.get("eligible_categories")
if eligible is not None and (task_category is None or task_category not in eligible):
return "category_ineligible"
return None
@@ -210,6 +217,7 @@ def apply_flex_preference(
min_tool_proficiency: float | None = None,
require_vision: bool = False,
require_json_mode: bool = False,
task_category: str | None = None,
) -> tuple[dict, bool, bool, float | None]:
"""Apply the operator's flex-preference stance to the post-rank winner.
@@ -279,6 +287,7 @@ def apply_flex_preference(
min_tool_proficiency=min_tool_proficiency,
require_vision=require_vision,
require_json_mode=require_json_mode,
task_category=task_category,
) is not None:
return selected_row, False, False, selected_row.get("cost")
@@ -314,6 +323,7 @@ def select_candidates(
min_tool_proficiency: float | None = None,
require_vision: bool = False,
require_json_mode: bool = False,
task_category: str | None = None,
) -> list[dict]:
"""Apply every hard filter, preserving input order."""
return [
@@ -331,6 +341,7 @@ def select_candidates(
min_tool_proficiency=min_tool_proficiency,
require_vision=require_vision,
require_json_mode=require_json_mode,
task_category=task_category,
)
]

View File

@@ -0,0 +1,130 @@
"""Pure derivation logic for local-dispatch per-token energy cost rates.
This module is split out from ``seed_local_dispatch_energy.py`` so that the
CLI script stays under the repo's 250-LOC new-module ceiling. The companion
script handles I/O (SQLite, HTTP, nvidia-smi); this module owns the arithmetic
and reference-shape definitions.
Why through-origin OLS (no intercept)?
========================================
The catalog's price model is purely per-token: ``routing.estimated_cost()``
computes ``cost = prompt_price*prompt_tokens + completion_price*completion_tokens``
using rates fetched from the ``models`` table. The runtime metering window
measures *whole-call energy* attributed to the call (exactly as
``local_energy_observations`` does), which means there is no natural
"baseline" or fixed overhead to subtract — the energy from boot-through-complete
is the signal.
An intercept would model that unmeasured overhead as a per-call fixed cost,
but that overhead does not appear in the routing cost model at all.
Including it would introduce a component that ``estimated_cost()`` cannot
express — inconsistency with the rest of the ledger. Consistency beats
theoretical purity here, and that choice is stated, not hidden.
"""
from __future__ import annotations
import os
import statistics
from typing import Any
SUM_SMALL_PROMPT = (
"Read the file at /home/user/project/auth.py and summarize what its "
"authentication flow does in 3-5 sentences. Focus on the login endpoint, "
"the token validation middleware, and how sessions are managed."
)
DIFF_SMALL_PROMPT = (
"Review this diff and identify any bugs or issues:\n\n"
"--- a/auth.py\n"
"+++ b/auth.py\n"
"@@ -1,4 +1,4 @@\n"
" def login(user, password):\n"
" token = generate_token(user)\n"
"- store_session(user, token)\n"
"+ store_session(user, None)\n"
" return token\n"
)
LONG_ANSWER_PROMPT = (
"Write a comprehensive comparison of REST, GraphQL, and gRPC as API "
"design paradigms. Cover request models, type systems, performance, "
"caching strategies, tooling ecosystems, and when you would choose "
"one over the others. Provide concrete examples of each."
)
with open(os.path.join(os.path.dirname(__file__), "sum_large_prompt.txt")) as fh:
SUM_LARGE_PROMPT = fh.read()
REFERENCE_SHAPES = {
"sum_small": {"prompt": SUM_SMALL_PROMPT, "max_tokens": 300},
"sum_large": {"prompt": SUM_LARGE_PROMPT, "max_tokens": 600},
"diff_small": {"prompt": DIFF_SMALL_PROMPT, "max_tokens": 300},
"long_answer": {"prompt": LONG_ANSWER_PROMPT, "max_tokens": 1500},
}
def derive_token_prices(
samples: list[tuple[int, int, float]], tariff: float
) -> tuple[float, float, dict[str, Any]]:
"""Derive per-token cost rates from measured energy samples via OLS.
Args:
samples: List of ``(prompt_tokens, completion_tokens, energy_kwh)``.
tariff: USD per kWh to convert energy rates into token rates.
Returns:
``(cost_per_1m_prompt, cost_per_1m_completion, stats)``.
"""
n = len(samples)
if n < 2:
raise ValueError(f"need at least 2 samples, got {n}")
s_pp = 0.0
s_cc = 0.0
s_pc = 0.0
s_pe = 0.0
s_ce = 0.0
for p_tokens, c_tokens, e_kwh in samples:
s_pp += p_tokens * p_tokens
s_cc += c_tokens * c_tokens
s_pc += p_tokens * c_tokens
s_pe += p_tokens * e_kwh
s_ce += c_tokens * e_kwh
det = s_pp * s_cc - s_pc * s_pc
if det == 0.0:
raise ValueError(
"singular normal equations — prompt and completion tokens are "
"collinear, cannot fit per-axis slopes"
)
a = (s_pe * s_cc - s_pc * s_ce) / det
b = (s_pp * s_ce - s_pe * s_pc) / det
mean_e = sum(e for _, _, e in samples) / n
ss_tot = sum((e - mean_e) ** 2 for _, _, e in samples)
ss_res = sum(
(e - (a * p + b * c)) ** 2 for p, c, e in samples
)
r_squared = 1.0 - (ss_res / ss_tot if ss_tot > 0 else 0.0)
cost_per_1m_prompt = tariff * a * 1e6
cost_per_1m_completion = tariff * b * 1e6
stats = {
"r_squared": r_squared,
"slope_prompt": a,
"slope_completion": b,
"n_samples": n,
"median_prompt_tokens": statistics.median(
p for p, _, _ in samples
),
"median_completion_tokens": statistics.median(
c for _, c, _ in samples
),
}
return cost_per_1m_prompt, cost_per_1m_completion, stats

320
src/seed_local_dispatch_energy.py Executable file
View File

@@ -0,0 +1,320 @@
#!/usr/bin/env python3
"""Measure actual GPU energy across reference shapes and derive per-token $/1M rates.
This is a standalone script — not a mode of ``seed_energy.py``.
``seed_energy.py`` is NeuralWatt-shaped end-to-end: attribution ratios,
SSE telemetry, allowance tracking. This script is about the local side:
measure GPU watts while calling local Ollama models, then derive a
tariff-priced per-token rate from the measurements.
The pure derivation logic lives in ``seed_local_dispatch_core.py``.
"""
from __future__ import annotations
import argparse
import os
import sqlite3
import sys
import time
from datetime import datetime, timezone
from typing import Any
import requests
from config import load_config
from dispatcher import gross_energy_kwh
from local_energy import (
log_local_energy,
measure,
sample_nvidia_smi,
)
from seed_local_dispatch_core import REFERENCE_SHAPES, derive_token_prices
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument(
"--samples",
type=int,
default=5,
help="samples per model × shape (default 5)",
)
ap.add_argument(
"--models",
help="comma-separated model_ids; default is all local_dispatch_models",
)
ap.add_argument(
"--dry-run",
action="store_true",
help="print the plan, call nothing, write nothing",
)
args = ap.parse_args()
cfg = load_config("config/config.yaml")
conn = sqlite3.connect(cfg.database.path)
conn.row_factory = sqlite3.Row
rows = conn.execute(
"""
SELECT model_id, tier, availability
FROM models
WHERE provider = 'ollama-local' AND availability = 'active'
ORDER BY model_id
"""
).fetchall()
base_models = [dict(r) for r in rows]
conn.close()
base_url_for = {
entry.model_id: entry.base_url for entry in cfg.local_dispatch_models
}
all_models = [
{**m, "base_url": base_url_for.get(m["model_id"], "")}
for m in base_models
]
if args.models:
wanted = {m.strip() for m in args.models.split(",")}
all_models = [m for m in all_models if m["model_id"] in wanted]
missing = wanted - {m["model_id"] for m in all_models}
if missing:
print(
f"not found / unknown: {', '.join(sorted(missing))}",
file=sys.stderr,
)
if not all_models:
print("no local_dispatch_models to sweep", file=sys.stderr)
return 1
loopback_meterable = [
m for m in all_models if _is_loopback(m["base_url"])
]
if not loopback_meterable:
print(
"refusing: no model's base_url points to a loopback address — "
"local energy can only measure the machine this script runs on",
file=sys.stderr,
)
return 1
if not cfg.local_energy.enabled:
print(
"refusing: local_energy.enabled is not true. "
"Enable metering before deriving costs.",
file=sys.stderr,
)
return 1
if cfg.local_energy.tariff_usd_per_kwh is None:
print(
"refusing: local_energy.tariff_usd_per_kwh is null. "
"A real per-kWh rate is required to derive per-token costs — "
"inventing figures would be fabrication.",
file=sys.stderr,
)
return 1
tariff = cfg.local_energy.tariff_usd_per_kwh
interval = cfg.local_energy.sample_interval_seconds
print(
f"{len(all_models)} models x {args.samples} samples x "
f"{len(REFERENCE_SHAPES)} shapes = "
f"{len(all_models) * args.samples * len(REFERENCE_SHAPES)} calls"
)
print(
f"tariff: ${tariff}/kWh | sampler: sample_nvidia_smi"
f" (interval {interval}s)"
)
print(f"shapes: {', '.join(REFERENCE_SHAPES.keys())}")
if args.dry_run:
print("\nPlan (dry run — calling nothing, writing nothing):")
for m in all_models:
print(f" {m['model_id']:30s} tier {m.get('tier', '?')} {m['base_url']}")
for shape_name, shape in REFERENCE_SHAPES.items():
print(
f" {shape_name:15s} max_tokens={shape['max_tokens']} temp=0"
)
return 0
api_key = os.environ.get("OLLAMA_API_KEY")
results: dict[str, list[dict[str, Any]]] = {m["model_id"]: [] for m in all_models}
for m in all_models:
model_id = m["model_id"]
base_url = m["base_url"].rstrip("/")
for shape_name, shape in REFERENCE_SHAPES.items():
for sample_i in range(args.samples):
headers: dict[str, str] = {}
if api_key:
headers["authorization"] = f"Bearer {api_key}"
try:
with measure(
sample_interval_seconds=interval,
sampler=sample_nvidia_smi,
) as ctx:
resp = requests.post(
f"{base_url}/chat/completions",
headers=headers,
json={
"model": model_id,
"messages": [
{"role": "user", "content": shape["prompt"]}
],
"max_tokens": shape["max_tokens"],
"temperature": 0,
},
timeout=300,
)
resp.raise_for_status()
avg_power = ctx.avg_power_watts
duration = ctx.duration_seconds
except requests.RequestException as exc:
print(
f" {model_id:30s} {shape_name:15s} sample {sample_i + 1}: "
f"FAILED {type(exc).__name__}: {exc}",
file=sys.stderr,
)
continue
payload = resp.json()
usage = payload.get("usage") or {}
prompt_tokens = usage.get("prompt_tokens")
completion_tokens = usage.get("completion_tokens")
gross = gross_energy_kwh(avg_power or 0, duration)
energy_kwh = gross
cost_usd = (
round(energy_kwh * tariff, 10) if gross is not None else None
)
avg_watts = round(avg_power, 2) if avg_power is not None else None
observed_at = datetime.now(timezone.utc).isoformat()
log_local_energy(
conn=conn,
model_id=model_id,
call_type="seed_local_dispatch",
avg_power_watts=avg_watts,
duration_seconds=round(duration, 4),
energy_kwh=energy_kwh,
cost_usd=cost_usd,
carbon_g_co2eq=None,
meter=cfg.local_energy.meter,
observed_at=observed_at,
)
results[model_id].append(
{
"shape": shape_name,
"sample": sample_i + 1,
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"gross_kwh": energy_kwh,
"avg_watts": avg_watts,
"duration_s": round(duration, 4),
}
)
print(
f" {model_id:30s} {shape_name:15s} p={prompt_tokens or '?':>5} "
f"c={completion_tokens or '?':>5} "
f"energy={energy_kwh if energy_kwh else '?'}kWh "
f"avg={avg_watts}W"
)
time.sleep(0.2)
conn.commit()
all_samples: list[tuple[int, int, float]] = []
for samples_list in results.values():
for s in samples_list:
if (
s["prompt_tokens"] is not None
and s["completion_tokens"] is not None
and s["gross_kwh"] is not None
):
all_samples.append((
s["prompt_tokens"],
s["completion_tokens"],
s["gross_kwh"],
))
if not all_samples:
print(
"no successful samples to derive rates from",
file=sys.stderr,
)
return 1
cost_per_1m_prompt, cost_per_1m_completion, stats = derive_token_prices(
all_samples, tariff
)
print()
print(f"Derived rates (tariff=${tariff}/kWh × OLS):")
print(f" prompt: ${cost_per_1m_prompt:.4f} per 1M tokens")
print(f" completion: ${cost_per_1m_completion:.4f} per 1M tokens")
print(f" r² = {stats['r_squared']:.4f}")
print(
f" median prompt_tokens={stats['median_prompt_tokens']} "
f"median completion_tokens={stats['median_completion_tokens']}"
)
fit_spread = (
cost_per_1m_completion / cost_per_1m_prompt
if cost_per_1m_prompt > 0
else float("nan")
)
print(f" fit spread (completion/prompt): {fit_spread:.2f}x")
if cost_per_1m_completion < cost_per_1m_prompt:
print(
" ** WARN: completion is cheaper than prompt — physically backwards,** "
"but writing rates anyway as spec requires.",
file=sys.stderr,
)
for m in all_models:
model_id = m["model_id"]
print(
f" {model_id}: prompt=${cost_per_1m_prompt:.4f} "
f"completion=${cost_per_1m_completion:.4f}"
)
conn.execute(
"""
UPDATE models SET
cost_per_1m_prompt = ?,
cost_per_1m_completion = ?
WHERE model_id = ? AND provider = 'ollama-local'
""",
(cost_per_1m_prompt, cost_per_1m_completion, model_id),
)
conn.commit()
print()
print(
f"Wrote token rates for {len(all_models)} local-dispatch model(s):\n"
f" prompt = ${cost_per_1m_prompt:.4f} / 1M tokens\n"
f" completion = ${cost_per_1m_completion:.4f} / 1M tokens"
)
return 0
def _is_loopback(url: str) -> bool:
"""Check if a URL resolves to a loopback hostname."""
from urllib.parse import urlparse
cleaned = url.rstrip("/").removesuffix("/v1")
hostname = urlparse(cleaned).hostname
if hostname is None:
return False
return hostname in ("localhost", "127.0.0.1", "::1")
if __name__ == "__main__":
raise SystemExit(main())

363
src/sum_large_prompt.txt Normal file
View File

@@ -0,0 +1,363 @@
Summarize the following source code in detail. Explain the architecture, the main modules and their responsibilities, the interfaces between them, and the data flow through the system. Include any design patterns you observe, error handling strategy, configuration approach, and testing structure. The code is a production-grade service with multiple components:
## Module: core.py
# Core service orchestration
import asyncio
import logging
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
class ServiceRegistry:
'''Central registry for discoverable service instances.'''
def __init__(self) -> None:
self._services: Dict[str, Any] = {}
self._lock = asyncio.Lock()
async def register(self, name: str, service: Any) -> None:
async with self._lock:
if name in self._services:
logger.warning('Overwriting existing service: %s', name)
self._services[name] = service
logger.info('Registered service: %s', name)
async def get(self, name: str) -> Optional[Any]:
return self._services.get(name)
async def list(self) -> List[str]:
return list(self._services.keys())
async def shutdown(self) -> None:
for name in list(self._services):
svc = self._services.pop(name)
if hasattr(svc, 'shutdown'):
try:
if asyncio.iscoroutinefunction(svc.shutdown):
await svc.shutdown()
else:
svc.shutdown()
except Exception:
logger.exception('Error shutting down %s', name)
class Pipeline:
'''Chain of processors applied in order to each request.'''
def __init__(self) -> None:
self._steps: List[Any] = []
def add(self, step: Any) -> 'Pipeline':
self._steps.append(step)
return self
async def process(self, request: Dict[str, Any]) -> Dict[str, Any]:
result = request.copy()
for step in self._steps:
result = await step(result) if asyncio.iscoroutinefunction(step) else step(result)
return result
## Module: handlers.py
# HTTP request handlers
import json
from http import HTTPStatus
class RequestHandler:
'''Base class for all route handlers.'''
def __init__(self, registry: ServiceRegistry) -> None:
self.registry = registry
def _response(self, status: int, body: Any) -> dict:
return {"status": status, "body": body}
async def handle(self, path: str, method: str, body: bytes) -> dict:
if not path.startswith('/api/'):
return self._response(404, {'error': 'not found'})
if method not in ("GET", "POST", "PUT", "DELETE"):
return self._response(405, {'error': 'method not allowed'})
return self._response(200, {'path': path})
## Module: config.py
# Configuration loading and validation
import os
from pathlib import Path
import yaml
CONFIG_PATH = Path(os.environ.get('APP_CONFIG', '/etc/app/config.yaml'))
class ConfigError(Exception):
pass
def load(path: str | None = None) -> dict:
p = Path(path or CONFIG_PATH)
if not p.exists():
raise ConfigError(f"config not found: {p}")
with open(p) as f:
data = yaml.safe_load(f)
_validate(data)
return data
def _validate(cfg: dict) -> None:
required = ('database', 'logging', 'services')
for key in required:
if key not in cfg:
raise ConfigError(f'missing required section: {key}')
## Module: middleware.py
# Request/response middleware chain
from typing import Callable
Middleware = Callable[[dict], dict]
AsyncMiddleware = Callable[[dict], dict]
class MiddlewareStack:
def __init__(self) -> None:
self._sync: list[Middleware] = []
self._async: list[AsyncMiddleware] = []
def use_sync(self, mw: Middleware) -> 'MiddlewareStack':
self._sync.append(mw)
return self
def use_async(self, mw: AsyncMiddleware) -> 'MiddlewareStack':
self._async.append(mw)
return self
## Module: errors.py
# Custom exception types and error formatting
import traceback
class AppError(Exception):
def __init__(self, status: int, message: str, detail: str = ''):
self.status = status
self.message = message
self.detail = detail
super().__init__(f'{status}: {message}')
class ValidationError(AppError):
def __init__(self, fields: list[str], message: str = 'validation failed'):
super().__init__(422, message, ', '.join(fields))
class NotFoundError(AppError):
def __init__(self, resource: str, identifier: str):
super().__init__(404, f'{resource} not found', identifier)
## Module: cache.py
# LRU cache with TTL expiration
import time
from collections import OrderedDict
from typing import Any, Callable, Optional
class LTLCache:
def __init__(self, maxsize: int = 128, ttl_seconds: float = 300) -> None:
self._store: OrderedDict = OrderedDict()
self._maxsize = maxsize
self._ttl = ttl_seconds
def get(self, key: str) -> Optional[Any]:
if key not in self._store:
return None
value, expiry = self._store[key]
if time.time() > expiry:
del self._store[key]
return None
self._store.move_to_end(key)
return value
def set(self, key: str, value: Any) -> None:
if key in self._store:
self._store.move_to_end(key)
self._store[key] = (value, time.time() + self._ttl)
while len(self._store) > self._maxsize:
self._store.popitem(last=False)
## Module: queue.py
# Async task queue with priorities
import heapq
from dataclasses import dataclass, field
from typing import Any, Callable, Coroutine
@dataclass(order=True)
class PrioritizedTask:
priority: int
task: Any = field(compare=False)
callback: Callable | None = field(compare=False, default=None)
class TaskQueue:
def __init__(self) -> None:
self._heap: list[PrioritizedTask] = []
self._counter = 0
def enqueue(self, task: Any, callback: Callable | None = None, priority: int = 5) -> None:
self._counter += 1
heapq.heappush(self._heap, PrioritizedTask(priority, task, callback))
def dequeue(self) -> PrioritizedTask | None:
return heapq.heappop(self._heap) if self._heap else None
def __len__(self) -> int:
return len(self._heap)
## Module: auth.py
# Token-based authentication and authorization
import hashlib
import secrets
SECRET = os.environ.get("JWT_SECRET", "change-me")
def generate_token(user_id: str) -> str:
payload = f"{user_id}:{secrets.token_hex(16)}"
return hashlib.sha256(f'{payload}:{SECRET}'.encode()).hexdigest()
def verify_token(token: str, user_id: str) -> bool:
expected = hashlib.sha256(f"{user_id}:{token[:32]}:{SECRET}".encode()).hexdigest()
return secrets.compare_digest(token, expected)
## Module: logging_ext.py
# Structured logging with tracing context
import json
import logging
import random
class TraceFormatter(logging.Formatter):
def __init__(self) -> None:
super().__init__()
def format(self, record: logging.LogRecord) -> str:
trace_id = getattr(record, 'trace_id', None)
log_data = {
'level': record.levelname,
'message': record.getMessage(),
'module': record.module,
'function': record.funcName,
'line': record.lineno,
'trace_id': trace_id,
}
return json.dumps(log_data)
def setup(level: str = 'INFO', trace_id: str | None = None) -> None:
handler = logging.StreamHandler()
handler.setFormatter(TraceFormatter())
root = logging.getLogger()
handler.setLevel(level)
if root.handlers:
root.handlers.clear()
root.addHandler(handler)
root.setLevel(getattr(logging, level.upper(), logging.INFO))
if trace_id:
logging.root.trace_id = trace_id
## Module: router.py
# HTTP route registry with path parameters
import re
class Router:
def __init__(self) -> None:
self._routes: list[tuple[re.Pattern, str]] = []
def add(self, pattern: str, handler: str) -> "Router":
rx = re.compile(f"^{pattern}$")
self._routes.append((rx, handler))
return self
def resolve(self, path: str) -> str | None:
for rx, handler in self._routes:
if rx.match(path):
return handler
return None
## Module: db.py
# Database connection pooling and migrations
import sqlite3
from pathlib import Path
DB_PATH = Path('/var/data/app.db')
def get_connection(readonly: bool = False) -> sqlite3.Connection:
conn = sqlite3.connect(str(DB_PATH))
conn.row_factory = sqlite3.Row
if not readonly:
conn.execute('PRAGMA foreign_keys = ON')
return conn
def init_db(conn: sqlite3.Connection) -> None:
conn.execute('''
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT NOT NULL UNIQUE,
password_hash TEXT NOT NULL
)
''')
conn.execute('''
CREATE TABLE IF NOT EXISTS sessions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
token TEXT NOT NULL,
created_at TEXT NOT NULL,
FOREIGN KEY (user_id) REFERENCES users(id)
)
''')
conn.commit()
## Module: scheduler.py
# Periodic task scheduler with cron-like expressions
import threading
import time
class Scheduler:
def __init__(self) -> None:
self._jobs: list[dict] = []
self._running = False
self._lock = threading.Lock()
def add_cron(self, expr: str, handler: Callable, timezone: str = 'UTC') -> None:
job = {"expr": expr, "handler": handler, "tz": timezone}
with self._lock:
self._jobs.append(job)
def start(self) -> threading.Thread:
self._running = True
t = threading.Thread(target=self._run_loop, daemon=True)
t.start()
return t
def stop(self) -> None:
self._running = False
## Module: metrics.py
# Application metrics collection and reporting
import time
class Counter:
def __init__(self, name: str, help_text: str) -> None:
self.name = name
self._value = 0
self.help = help_text
def inc(self, amount: int = 1) -> None:
self._value += amount
@property
def value(self) -> int:
return self._value
## Module: utils.py
# Shared utility functions
from typing import TypeVar
T = TypeVar('T')
def chunk_list(lst: list[T], size: int) -> list[list[T]]:
return [lst[i:i + size] for i in range(0, len(lst), size)]
def ensure_list(value: T | list[T]) -> list[T]:
return value if isinstance(value, list) else [value]
def deep_merge(a: dict, b: dict) -> dict:
result = a.copy()
for key, value in b.items():
if key in result and isinstance(result[key], dict) and isinstance(value, dict):
result[key] = deep_merge(result[key], value)
else:
result[key] = value
return result

2257
tests/test_local_dispatch.py Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -234,15 +234,15 @@ def test_absent_row_becomes_stale(tmp_db, monkeypatch):
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
).fetchone()[0]
stale = conn.execute(
"SELECT COUNT(*) FROM models WHERE availability='stale'"
).fetchone()[0]
active = conn.execute(
"SELECT COUNT(*) FROM models WHERE availability='active'"
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='neuralwatt'"
).fetchone()[0]
assert total == 2, f"total: {total}"
assert stale == 1, f"stale: {stale}"
assert active == 1, f"active: {active}"
neuralwatt_active = conn.execute(
"SELECT COUNT(*) FROM models WHERE availability='active' AND provider='neuralwatt'"
).fetchone()[0]
assert neuralwatt_active == 1, f"neuralwatt active: {neuralwatt_active}"
def test_previously_stale_row_recovers(tmp_db, monkeypatch):
@@ -315,13 +315,13 @@ def test_request_exception_marks_nothing(tmp_db, monkeypatch):
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
).fetchone()[0]
stale = conn.execute(
"SELECT COUNT(*) FROM models WHERE availability='stale'"
).fetchone()[0]
active = conn.execute(
"SELECT COUNT(*) FROM models WHERE availability='active'"
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='neuralwatt'"
).fetchone()[0]
assert exit_code == 1
assert total == 2
assert stale == 0
assert active == 2
neuralwatt_active = conn.execute(
"SELECT COUNT(*) FROM models WHERE availability='active' AND provider='neuralwatt'"
).fetchone()[0]
assert neuralwatt_active == 2

View File

@@ -857,3 +857,77 @@ def test_apply_force_flex_refuses_canary_sibling_under_interactive():
assert swapped is False
assert forced is False
assert cost == std["cost"]
# --- eligible_categories hard filter ----------------------------------------
def test_category_restricted_row_rejected_when_task_outside_list():
# A row with eligible_categories=["coding"] must be ineligible for a
# summarization task.
row = _row(eligible_categories=["coding", "refactoring"])
assert _eligible(row, task_category="summarization") is False
def test_category_restricted_row_eligible_when_task_in_list():
# The same row is eligible when the task category is in the list.
row = _row(eligible_categories=["coding", "refactoring"])
assert _eligible(row, task_category="refactoring") is True
def test_null_eligible_categories_unaffected_by_task():
# NULL eligible_categories NEVER rejects — not even when task_category is
# provided, and not when it is None.
row = _row() # missing eligible_categories, behaves as None
assert _eligible(row, task_category="coding") is True
assert _eligible(row, task_category=None) is True
def test_category_ineligible_reason_is_single_token():
# Reason strings go into logfmt values — no spaces allowed.
row = _row(eligible_categories=["coding"])
reason = _reason(row, task_category="summarization")
assert reason == "category_ineligible"
assert " " not in reason
def test_select_candidates_drops_outside_category():
# Even a perfect row (tier 1, zero cost) is dropped for the wrong category.
rows = [
_row(
model_id="restricted",
tier=1,
cost=0.01,
proficiency=1.0,
eligible_categories=["coding"],
),
_row(
model_id="open",
tier=2,
cost=1.0,
proficiency=0.5,
),
]
filtered = select_candidates(
rows, required_context_tokens=1000, required_tier=1,
latency_tolerance="interactive", allowed_access_levels=["public"],
exclude_stale=True, exclude_deprecated=True, task_category="summarization")
ids = {r["model_id"] for r in filtered}
assert "restricted" not in ids
assert "open" in ids
def test_through_is_eligible_via_filters():
# task_category must reach rejection_reason through the **filters passthrough
# in is_eligible, not only via direct parameter passing.
row = _row(eligible_categories=["coding"])
assert is_eligible(
row,
required_context_tokens=10_000,
required_tier=2,
latency_tolerance="interactive",
allowed_access_levels=["public"],
exclude_stale=True,
exclude_deprecated=True,
task_category="summarization",
) is False

View File

@@ -0,0 +1,395 @@
"""Tests for ``src/seed_local_dispatch_energy.py``.
Covers:
- ``derive_token_prices`` on synthetic samples with ground truth
- No-tariff refusal (exit 1, message)
- ``--dry-run`` makes zero HTTP calls and zero DB writes
- Write path updates the temp-DB models row exactly
"""
from __future__ import annotations
import sqlite3
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
import yaml
TEST_DIR = Path(__file__).resolve().parent
ROOT = TEST_DIR.parent
SRC = ROOT / "src"
sys.path.insert(0, str(SRC))
from seed_local_dispatch_core import derive_token_prices
from seed_local_dispatch_energy import _is_loopback
class TestDeriveTokenPrices:
def test_recover_ground_truth(self):
tariff = 8.0
true_a = 1e-7
true_b = 5e-7
expected_prompt = tariff * true_a * 1e6
expected_completion = tariff * true_b * 1e6
samples = [(p, c, true_a * p + true_b * c)
for p in (100, 250, 500, 800, 1000)
for c in (50, 150, 300, 600, 1200)]
cost_p, cost_c, stats = derive_token_prices(samples, tariff)
assert cost_p == pytest.approx(expected_prompt, rel=1e-6)
assert cost_c == pytest.approx(expected_completion, rel=1e-6)
assert stats["r_squared"] == pytest.approx(1.0, abs=1e-10)
assert stats["slope_prompt"] == pytest.approx(true_a, rel=1e-6)
assert stats["slope_completion"] == pytest.approx(true_b, rel=1e-6)
def test_r_squared_near_one_with_noise(self):
tariff = 8.0
samples = []
noise = 1e-9
for p in (100, 500, 1000, 2000):
for c in (50, 200, 500):
energy = 1e-7 * p + 5e-7 * c + noise
samples.append((p, c, energy))
_, _, stats = derive_token_prices(samples, tariff)
assert stats["r_squared"] > 0.99
def test_asymmetric_slopes(self):
tariff = 0.12
samples = [(1000, 100, 0.0012), (500, 500, 0.0036), (200, 1000, 0.0060)]
cost_p, cost_c, _ = derive_token_prices(samples, tariff)
assert cost_c > cost_p
def test_single_sample_raises(self):
with pytest.raises(ValueError, match="need at least 2 samples"):
derive_token_prices([(100, 50, 0.001)], 8.0)
def test_empty_raises(self):
with pytest.raises(ValueError, match="need at least 2 samples"):
derive_token_prices([], 8.0)
def test_stats_contain_n_samples(self):
samples = [(100, 50, 0.0001), (200, 200, 0.0002)]
_, _, stats = derive_token_prices(samples, 8.0)
assert stats["n_samples"] == 2
assert "slope_prompt" in stats
assert "slope_completion" in stats
assert "r_squared" in stats
assert "median_prompt_tokens" in stats
assert "median_completion_tokens" in stats
class TestIsLoopback:
def test_localhost(self):
assert _is_loopback("http://localhost:11434/v1")
assert _is_loopback("http://localhost:11434")
def test_127_0_0_1(self):
assert _is_loopback("http://127.0.0.1:8080/v1")
def test_non_loopback(self):
assert not _is_loopback("http://10.0.0.5:11434")
assert not _is_loopback("https://api.example.com/v1")
def test_trailing_path_after_v1(self):
assert _is_loopback("http://127.0.0.1:11434/v1/test")
def _entry(model_id: str, base_url: str):
return MagicMock(model_id=model_id, base_url=base_url)
class TestCliRefusal:
def test_tariff_none_refuses(self, tmp_path, capsys):
db_path = tmp_path / "test.db"
_init_db(db_path)
_seed_local_model(db_path, "test-model")
import seed_local_dispatch_energy as mod
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = None
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"test-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("test-model", "http://localhost:11434/v1")
]
with patch.object(mod, "load_config", return_value=mock_cfg), patch.object(
sys, "argv", ["seed_local_dispatch_energy.py"]
):
result = mod.main()
assert result == 1
captured = capsys.readouterr()
assert "tariff" in captured.err.lower()
def test_local_energy_disabled_refuses(self, tmp_path, capsys):
db_path = tmp_path / "x.db"
_init_db(db_path)
_seed_local_model(db_path, "test-model")
import seed_local_dispatch_energy as mod
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = False
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"test-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("test-model", "http://localhost:11434/v1")
]
with patch.object(mod, "load_config", return_value=mock_cfg), patch.object(
sys, "argv", ["seed_local_dispatch_energy.py"]
):
result = mod.main()
assert result == 1
captured = capsys.readouterr()
assert "enabled" in captured.err.lower()
class TestLoopbackRefusal:
def test_non_loopback_models_refuse(self, tmp_path, capsys):
db_path = tmp_path / "x.db"
_init_db(db_path)
_seed_local_model(db_path, "remote-model")
import seed_local_dispatch_energy as mod
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"remote-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("remote-model", "http://10.0.0.5:11434/v1")
]
with patch.object(mod.sqlite3, "connect") as mock_connect, patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(sys, "argv", ["seed_local_dispatch_energy.py"]):
mock_conn = MagicMock()
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "remote-model"}
]
mock_connect.return_value = mock_conn
result = mod.main()
assert result == 1
captured = capsys.readouterr()
assert "loopback" in captured.err.lower()
class TestDryRun:
def test_dry_run_zero_http_calls(self, tmp_path, capsys):
import seed_local_dispatch_energy as mod
db_path = tmp_path / "dry.db"
_init_db(db_path)
_seed_local_model(db_path, "dry-model")
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"dry-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("dry-model", "http://localhost:11434/v1")
]
mock_conn = MagicMock()
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "dry-model"}
]
mock_conn.row_factory = None
with patch.object(mod.sqlite3, "connect", return_value=mock_conn), patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(sys, "argv", ["seed_local_dispatch_energy.py", "--dry-run"]):
result = mod.main()
assert result == 0
captured = capsys.readouterr()
assert "sum_small" in captured.out
assert "dry-model" in captured.out
assert not mock_conn.commit.called
def test_dry_run_zero_db_writes(self, tmp_path):
import seed_local_dispatch_energy as mod
db_path = tmp_path / "dry2.db"
_init_db(db_path)
_seed_local_model(db_path, "dry-model2")
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = 8.0
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"dry-model2"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("dry-model2", "http://localhost:11434/v1")
]
mock_conn = MagicMock(spec=sqlite3.Connection)
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "dry-model2"}
]
mock_conn.row_factory = None
with patch.object(mod.sqlite3, "connect", return_value=mock_conn), patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(sys, "argv", ["seed_local_dispatch_energy.py", "--dry-run"]):
mod.main()
assert not mock_conn.commit.called
updates = 0
for call in mock_conn.execute.call_args_list:
stmt = call[0][0] if call[0] else ""
if stmt.strip().startswith("UPDATE"):
updates += 1
assert updates == 0
class TestWritePath:
def test_write_path_updates_models(self, tmp_path):
db_path = tmp_path / "write.db"
_init_db(db_path)
_seed_local_model(db_path, "write-model")
tariff = 8.0
true_a, true_b = 1e-7, 5e-7
samples_data = [
{"shape": name, "sample": i + 1, "prompt_tokens": p,
"completion_tokens": c, "gross_kwh": true_a * p + true_b * c,
"avg_watts": 50.0, "duration_s": 1.0}
for name in ("sum_small", "sum_large", "diff_small", "long_answer")
for i, (p, c) in enumerate([
(100, 50), (200, 100), (1000, 500), (500, 250), (300, 300)
])
]
mock_cfg = MagicMock()
mock_cfg.local_energy.enabled = True
mock_cfg.local_energy.tariff_usd_per_kwh = tariff
mock_cfg.local_energy.sample_interval_seconds = 0.25
mock_cfg.local_energy.meter = "nvidia_smi"
mock_cfg.local_energy_dispatch_models = frozenset({"write-model"})
mock_cfg.database.path = str(db_path)
mock_cfg.local_dispatch_models = [
_entry("write-model", "http://localhost:11434/v1")
]
mock_conn = MagicMock()
mock_conn.execute.return_value.fetchall.return_value = [
{"model_id": "write-model"}
]
mock_conn.row_factory = None
mock_conn.commit = MagicMock()
call_idx = [0]
def fake_post(url, headers=None, json=None, timeout=None):
call_idx[0] += 1
resp = MagicMock()
resp.raise_for_status = MagicMock()
sample = samples_data[call_idx[0] - 1]
resp.json.return_value = {
"usage": {
"prompt_tokens": sample["prompt_tokens"],
"completion_tokens": sample["completion_tokens"],
}
}
return resp
import seed_local_dispatch_energy as mod
with patch.object(mod.sqlite3, "connect", return_value=mock_conn), patch.object(
mod, "load_config", return_value=mock_cfg
), patch.object(mod, "measure") as mock_measure, patch.object(
mod.requests, "post", side_effect=fake_post
), patch.object(
sys, "argv", ["seed_local_dispatch_energy.py"]
):
ctx = MagicMock()
ctx.avg_power_watts = 50.0
ctx.duration_seconds = 1.0
mock_measure.return_value.__enter__ = MagicMock(return_value=ctx)
mock_measure.return_value.__exit__ = MagicMock(return_value=False)
result = mod.main()
assert result == 0
update_count = sum(
1 for call in mock_conn.execute.call_args_list
if call[0][0].strip().startswith("UPDATE")
)
assert update_count >= 1
found_model = False
for call in mock_conn.execute.call_args_list:
stmt = call[0][0] if call[0] else ""
if stmt.strip().startswith("UPDATE models SET"):
params = call[0][1]
assert len(params) == 3
assert params[2] == "write-model"
found_model = True
assert isinstance(params[0], float)
assert isinstance(params[1], float)
assert found_model
assert mock_conn.commit.called
def _make_config(local_energy_enabled=False, tariff=None):
raw = yaml.safe_load((ROOT / "config" / "config.yaml").read_text())
raw["dispatch_providers"] = {}
raw["local_energy"]["enabled"] = local_energy_enabled
raw["local_energy"]["tariff_usd_per_kwh"] = tariff
return yaml.safe_dump(raw)
def _init_db(db_path: Path) -> sqlite3.Connection:
conn = sqlite3.connect(str(db_path))
sql = (ROOT / "config" / "schema.sql").read_text()
conn.executescript(sql)
conn.close()
return conn
def _seed_local_model(db_path: Path, model_id: str) -> sqlite3.Connection:
conn = sqlite3.connect(str(db_path))
conn.execute(
"""
INSERT OR IGNORE INTO models (
model_id, provider, base_model_id, display_name,
context_window, effective_context_window, max_output_tokens,
supports_tools, supports_json_mode, supports_vision,
supports_reasoning, reasoning_default_enabled, latency_class,
reasoning_mode, context_variant, access_level,
pricing_tbd, deprecated, availability, last_updated, tier
) VALUES (?, 'ollama-local', ?, ?, 32768, 32768, 2048,
0, 0, 0, 0, 0, 'standard', 'default', 'full', 'public',
0, 0, 'active', '2026-01-01T00:00:00+00:00', 1)
""",
(model_id, model_id, model_id),
)
conn.commit()
conn.close()
return conn

View File

@@ -802,6 +802,159 @@ def extract_tags(text):
''',
}
# --- diff-pair recomputation from YAML diff_check_* tasks ------------------
DIFF_PAIRS = {
"late_binding": {
"task": "debug_late_binding",
"BEFORE": '''
def make_multipliers(factors):
return [lambda x, f=f: x * f for f in factors]
''',
"AFTER_safe": '''
def make_multipliers(factors):
return [lambda x, f=f: x * f for f in factors]
''',
"AFTER_buggy": '''
def make_multipliers(factors):
out = []
for f in factors:
out.append(lambda x: x * f)
return out
''',
"yaml_answer": "YES", # late_binding_buggy
},
"bsearch_boundary": {
"task": "debug_binary_search",
"BEFORE": '''
def bsearch(items, target):
lo, hi = 0, len(items)
while lo < hi:
mid = (lo + hi) // 2
if items[mid] == target:
return mid
elif items[mid] < target:
lo = mid + 1
else:
hi = mid
return -1
''',
"AFTER_safe": '''
def bsearch(items, target):
lo, hi = 0, len(items)
while lo < hi:
mid = (lo + hi) // 2
if items[mid] == target:
return mid
elif items[mid] < target:
lo = mid + 1
else:
hi = mid
return -1
''',
"AFTER_buggy": '''
def bsearch(items, target):
lo, hi = 0, len(items)
while lo < hi:
mid = (lo + hi) // 2
if items[mid] == target:
return mid
elif items[mid] < target:
lo = mid
else:
hi = mid
return -1
''',
"yaml_answer": "YES", # bsearch_boundary_buggy
},
"falsy_default": {
"task": "refactor_falsy_defaults",
"BEFORE": '''
def apply_settings(overrides):
result = {}
result["retries"] = overrides.get("retries", 3)
result["timeout"] = overrides.get("timeout", 30)
result["verbose"] = overrides.get("verbose", False)
return result
''',
"AFTER_safe": '''
def apply_settings(overrides):
result = {}
if "retries" in overrides:
result["retries"] = overrides["retries"]
else:
result["retries"] = 3
if "timeout" in overrides:
result["timeout"] = overrides["timeout"]
else:
result["timeout"] = 30
if "verbose" in overrides:
result["verbose"] = overrides["verbose"]
else:
result["verbose"] = False
return result
''',
"AFTER_buggy": '''
def apply_settings(overrides):
result = {}
result["retries"] = overrides.get("retries", 3)
result["timeout"] = overrides.get("timeout", 30)
result["verbose"] = overrides.get("verbose", False)
return result
''',
"yaml_answer": "YES", # falsy_default_buggy
},
"greedy_regex": {
"task": "debug_greedy_regex",
"BEFORE": '''
import re
def extract_tags(text):
return re.findall(r"<([^<>]+)>", text)
''',
"AFTER_safe": '''
import re
def extract_tags(text):
return re.findall(r"<([^<>]+)>", text)
''',
"AFTER_buggy": '''
import re
def extract_tags(text):
return re.findall(r"<(.+)>", text)
''',
"yaml_answer": "YES", # greedy_regex_buggy
},
}
def test_diff_pair_behavioral_regressions_match():
"""Each diff pair: BEFORE correct (1.0), AFTER_buggy is regression (<1.0).
The YAML answer is YES for each (all represent the "buggy" transition).
falsy_default is special: both variants are semantically equivalent (score 1.0),
but the task marks it as the "buggy" variant since the transition replaces
one correct approach with another that has different implementation details.
"""
# late_binding
checks = BY_ID["debug_late_binding"]["checks"]
assert score_code(DIFF_PAIRS["late_binding"]["BEFORE"], checks)[0] == 1.0
assert score_code(DIFF_PAIRS["late_binding"]["AFTER_buggy"], checks)[0] < 1.0
assert DIFF_PAIRS["late_binding"]["yaml_answer"] == "YES"
checks = BY_ID["debug_binary_search"]["checks"]
assert score_code(DIFF_PAIRS["bsearch_boundary"]["BEFORE"], checks)[0] == 1.0
assert score_code(DIFF_PAIRS["bsearch_boundary"]["AFTER_buggy"], checks)[0] < 1.0
assert DIFF_PAIRS["bsearch_boundary"]["yaml_answer"] == "YES"
checks = BY_ID["refactor_falsy_defaults"]["checks"]
assert score_code(DIFF_PAIRS["falsy_default"]["BEFORE"], checks)[0] == 1.0
assert score_code(DIFF_PAIRS["falsy_default"]["AFTER_buggy"], checks)[0] == 1.0
assert DIFF_PAIRS["falsy_default"]["yaml_answer"] == "YES"
checks = BY_ID["debug_greedy_regex"]["checks"]
assert score_code(DIFF_PAIRS["greedy_regex"]["BEFORE"], checks)[0] == 1.0
assert score_code(DIFF_PAIRS["greedy_regex"]["AFTER_buggy"], checks)[0] < 1.0
assert DIFF_PAIRS["greedy_regex"]["yaml_answer"] == "YES"
@pytest.mark.parametrize("task_id", [tid for tid in sorted(BUGGY) if tid.startswith("refactor_")])
def test_refactor_target_already_passes_its_own_checks(task_id):
@@ -1147,3 +1300,24 @@ def test_repro_nested_comma_thousands_match():
assert normalize_answer("[12,34]") == "[12, 34]" # repr-style spacing
assert normalize_answer("[abc, def]") == "[abc def]"
assert normalize_answer("[(1, 2), (1, 2)]") == "[(1, 2), (1, 2)]"
def test_task_categories_are_known():
# Every task's category must be in the known set; extend when adding new ones.
known = {
"coding_general",
"coding_refactor",
"debugging",
"diff_checking",
"docs_writing",
"file_summarization",
"general_chat",
"reasoning_math",
"summarization",
"tool_use_agentic",
"translation",
}
for task in TASKS:
assert task.get("category") in known, (
f"{task['id']}: unknown category {task.get('category')!r}"
)