feat(local-llm): route file_summarization + diff_checking to a metered local Ollama model #22
212
.omo/notepads/expand-local-llm-usage/learnings.md
Normal file
212
.omo/notepads/expand-local-llm-usage/learnings.md
Normal 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`.
|
||||
@@ -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 |
|
||||
|
||||
42
CLAUDE.md
42
CLAUDE.md
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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);
|
||||
|
||||
26
docs/api.md
26
docs/api.md
@@ -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"`.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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` |
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
337
evals/tasks.yaml
337
evals/tasks.yaml
@@ -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.
|
||||
|
||||
115
src/config.py
115
src/config.py
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
130
src/seed_local_dispatch_core.py
Normal file
130
src/seed_local_dispatch_core.py
Normal 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
320
src/seed_local_dispatch_energy.py
Executable 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
363
src/sum_large_prompt.txt
Normal 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
2257
tests/test_local_dispatch.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
395
tests/test_seed_local_dispatch.py
Normal file
395
tests/test_seed_local_dispatch.py
Normal 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
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user