diff --git a/CHANGELOG.md b/CHANGELOG.md index 6370546..cac74fb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,7 @@ ## Unreleased +- Added labels-only Ollama scoring with up to 20 choices through `choosekit/ollama`. - Added an exact `/v1/models` check to the SemIf llama.cpp benchmark, with `--skip-model-check` for unverified runs. diff --git a/README.md b/README.md index 8d9b324..acdf17b 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,6 @@ # choosekit -`choosekit` scores a finite set of choices with a language model and returns a typed decision with a probability distribution. It supports local llama.cpp models and an optional OpenRouter backend. +`choosekit` scores a finite set of choices with a language model and returns a typed decision with a probability distribution. It supports local models through llama.cpp and Ollama, plus an optional OpenRouter backend. ![SuperGPQA direct-choice benchmark](benchmarks/supergpqa-benchmark.svg) @@ -45,7 +45,7 @@ Agents often need to choose from known options: `choosekit` scores choices using the model's conditional log probabilities at the token branches that distinguish them. -The project was inspired by [Jev and the System One model interface](https://typesafe.ai/blog/introducing-system-one-models-and-jev): application state in, typed probabilistic decisions out. Jev is a specialized hosted model. `choosekit` brings the same typed decision interface to general-purpose language models. The llama.cpp backend runs on infrastructure you choose; OpenRouter provides hosted inference. +The project was inspired by [Jev and the System One model interface](https://typesafe.ai/blog/introducing-system-one-models-and-jev): application state in, typed probabilistic decisions out. Jev is a specialized hosted model. `choosekit` brings the same typed decision interface to general-purpose language models. The llama.cpp and Ollama backends run on infrastructure you choose; OpenRouter provides hosted inference. `choosekit` is an independent project with no affiliation to TypeSafe or Jev. @@ -81,6 +81,16 @@ The llama.cpp backend requires its native `/tokenize` and `/completion` endpoint The library has no telemetry. +## Ollama + +```ts +import { fromOllama } from "choosekit/ollama"; + +const choose = fromOllama({ model: "your-model" }); +``` + +`model` is required. The adapter uses `http://127.0.0.1:11434/` by default and supports only `labels` mode with up to 20 choices. It requires Ollama 0.12.11 or newer. + ## OpenRouter ```ts @@ -94,7 +104,7 @@ const choose = fromOpenRouter({ The OpenRouter backend supports models and providers that return first-token `top_logprobs`, with up to 20 choices. It sends the prompt to OpenRouter and requests reasoning to be disabled. -Choices omitted from `top_logprobs` receive zero probability. Returned probabilities are normalized across the supplied choices and are not calibrated correctness estimates. +Returned probabilities are normalized across the supplied choices and are not calibrated correctness estimates. OpenRouter may route the same model through different providers. Set `provider: "provider-slug"` to use only that provider and disable fallback. @@ -102,9 +112,11 @@ OpenRouter may route the same model through different providers. Set `provider: | Mode | Candidate representation | Use when | |---|---|---| -| `labels` | `A`, `B`, `C`, ... | Default. Up to 26 choices with llama.cpp or 20 with OpenRouter. | +| `labels` | `A`, `B`, `C`, ... | Default. Up to 26 choices with llama.cpp or 20 with Ollama or OpenRouter. | | `minimal-prefix` | Original JSON-quoted keys | llama.cpp only. Use when key names should influence the decision. | +Choices for which the backend returns no logprob receive zero probability. + In `labels` mode, choices are shown to the model as `A`, `B`, `C` instead of their original keys. For example, `refund: "Issue the refund"` is shown as `"A": "Issue the refund"`. Each description must therefore make the option clear. `choosekit` maps the selected label back to the original key. `minimal-prefix` walks the token tree until every key is distinguishable. For keys such as `watermelon` and `watermelon juice`, the shared token path is handled once and scoring stops when the paths separate. diff --git a/package.json b/package.json index f6a302f..dc3eac8 100644 --- a/package.json +++ b/package.json @@ -40,6 +40,16 @@ "default": "./dist/cjs/llama-cpp.js" } }, + "./ollama": { + "import": { + "types": "./dist/esm/ollama.d.ts", + "default": "./dist/esm/ollama.js" + }, + "require": { + "types": "./dist/cjs/ollama.d.ts", + "default": "./dist/cjs/ollama.js" + } + }, "./openrouter": { "import": { "types": "./dist/esm/openrouter.d.ts", @@ -79,6 +89,7 @@ "choices", "logprobs", "llama.cpp", + "ollama", "openrouter", "typescript" ] diff --git a/src/ollama.ts b/src/ollama.ts new file mode 100644 index 0000000..3bae4c1 --- /dev/null +++ b/src/ollama.ts @@ -0,0 +1,149 @@ +import { createFormattedChooser } from "./internal-chooser.js"; +import type { Chooser, ChooserOptions, Scorer, Usage } from "./types.js"; +import { isCount, isRecord, requireText, ScoringError } from "./validation.js"; + +const DEFAULT_BASE_URL = "http://127.0.0.1:11434"; +const MAX_CANDIDATES = 20; + +export interface OllamaOptions extends ChooserOptions { + readonly model: string; + readonly baseURL?: string; + readonly fetch?: typeof globalThis.fetch; +} + +function endpoint(value: unknown): string { + const baseURL = value === undefined ? DEFAULT_BASE_URL : value; + requireText(baseURL, "baseURL"); + const url = new URL(baseURL); + if ((url.protocol !== "http:" && url.protocol !== "https:") || url.username || url.password + || url.search || url.hash) { + throw new TypeError("baseURL must be an HTTP(S) URL without credentials, query, or fragment."); + } + const path = url.pathname.replace(/\/api\/?$/, "").replace(/\/$/, ""); + url.pathname = `${path}/api/chat`; + return url.href; +} + +async function post(fetchImpl: typeof globalThis.fetch, url: string, body: unknown, + signal?: AbortSignal): Promise { + signal?.throwIfAborted(); + let response: Response; + try { + response = await fetchImpl(url, { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify(body), + ...(signal ? { signal } : {}), + }); + } catch (error) { + signal?.throwIfAborted(); + throw error; + } + signal?.throwIfAborted(); + if (!response.ok) { + throw new ScoringError(`Ollama returned HTTP ${response.status} for ${new URL(url).pathname}.`); + } + try { + const value: unknown = await response.json(); + signal?.throwIfAborted(); + return value; + } catch { + signal?.throwIfAborted(); + throw new ScoringError(`Ollama returned invalid JSON for ${new URL(url).pathname}.`); + } +} + +function parseUsage(value: Record): Usage { + const promptTokens = value.prompt_eval_count; + const completionTokens = value.eval_count; + if (!isCount(promptTokens) || !isCount(completionTokens)) { + throw new ScoringError("Ollama returned invalid token usage."); + } + const cached = value.prompt_eval_cached_count; + if (cached !== undefined && (!isCount(cached) || cached > promptTokens)) { + throw new ScoringError("Ollama returned invalid cached-token usage."); + } + return Object.freeze({ + promptTokens, + cachedTokens: cached === undefined ? null : cached, + completionTokens, + requests: 1, + }); +} + +function collectLabelScore(value: unknown, expected: ReadonlySet, + found: Map): void { + if (!isRecord(value) || typeof value.token !== "string") { + throw new ScoringError("Ollama returned an invalid logprob entry."); + } + if (!expected.has(value.token)) return; + if (typeof value.logprob !== "number" || !Number.isFinite(value.logprob) + || value.logprob > 0) { + throw new ScoringError(`Ollama returned an invalid logprob for label ${value.token}.`); + } + if (value.bytes !== undefined && value.bytes !== null) { + if (!Array.isArray(value.bytes) || value.bytes.length !== 1 + || value.bytes[0] !== value.token.charCodeAt(0)) { + throw new ScoringError(`Ollama returned invalid bytes for label ${value.token}.`); + } + } + const existing = found.get(value.token); + if (existing !== undefined && existing !== value.logprob) { + throw new ScoringError(`Ollama returned conflicting logprobs for label ${value.token}.`); + } + found.set(value.token, value.logprob); +} + +function parseScores(value: unknown, candidates: readonly string[]): { + readonly logprobs: readonly number[]; + readonly usage: Usage; +} { + if (!isRecord(value) || !Array.isArray(value.logprobs) || value.logprobs.length !== 1 + || !isRecord(value.logprobs[0])) { + throw new ScoringError("Ollama did not return exactly one scored token position."); + } + const position = value.logprobs[0]; + if (!Array.isArray(position.top_logprobs) || position.top_logprobs.length === 0) { + throw new ScoringError("Ollama did not return top logprobs."); + } + const expected = new Set(candidates); + const found = new Map(); + collectLabelScore(position, expected, found); + for (const entry of position.top_logprobs) collectLabelScore(entry, expected, found); + if (found.size === 0) { + throw new ScoringError("Ollama did not return logprobs for any choice label."); + } + return { + logprobs: Object.freeze(candidates.map((candidate) => found.get(candidate) ?? -Infinity)), + usage: parseUsage(value), + }; +} + +export function fromOllama(options: OllamaOptions): Chooser { + if (!isRecord(options)) throw new TypeError("options must be an object."); + const { model, formatPrompt } = options; + requireText(model, "model"); + if (options.fetch !== undefined && typeof options.fetch !== "function") { + throw new TypeError("fetch must be a function."); + } + const fetchImpl = options.fetch ?? globalThis.fetch; + if (typeof fetchImpl !== "function") throw new TypeError("A fetch implementation is required."); + const url = endpoint(options.baseURL); + + const score: Scorer = async ({ prompt, candidates, signal }) => { + if (candidates.length > MAX_CANDIDATES) { + throw new TypeError(`Ollama supports at most ${MAX_CANDIDATES} choices.`); + } + const response = await post(fetchImpl, url, { + model, + messages: [{ role: "user", content: prompt }], + stream: false, + think: false, + logprobs: true, + top_logprobs: MAX_CANDIDATES, + options: { num_predict: 1 }, + }, signal); + return parseScores(response, candidates); + }; + return createFormattedChooser(score, formatPrompt === undefined ? {} : { formatPrompt }, "labels"); +} diff --git a/tests/ollama.test.mjs b/tests/ollama.test.mjs new file mode 100644 index 0000000..0fa4efa --- /dev/null +++ b/tests/ollama.test.mjs @@ -0,0 +1,201 @@ +import assert from "node:assert/strict"; +import { test } from "node:test"; +import { fromOllama } from "../dist/esm/ollama.js"; +import { ScoringError } from "../dist/esm/index.js"; + +const request = Object.freeze({ + context: "A customer says their payout failed.", + question: "Which team should handle this message?", + choices: Object.freeze({ + sales: "Pricing, upgrades, and new accounts.", + technical: "Bugs, outages, integrations, and API errors.", + billing: "Payments, payouts, invoices, and refunds.", + }), +}); + +const alternatives = Object.freeze([ + { token: "B", bytes: [66], logprob: -2.2 }, + { token: "C", bytes: [67], logprob: -0.1 }, + { token: "x", bytes: [120], logprob: -4 }, +]); + +function response(topLogprobs = alternatives, overrides = {}) { + return { + model: "test-model", + message: { role: "assistant", content: "A" }, + done: true, + done_reason: "length", + prompt_eval_count: 120, + prompt_eval_cached_count: 80, + eval_count: 1, + logprobs: [{ + token: "A", + bytes: [65], + logprob: -1.2, + top_logprobs: topLogprobs, + }], + ...overrides, + }; +} + +function json(value, status = 200) { + return new Response(JSON.stringify(value), { + status, + headers: { "content-type": "application/json" }, + }); +} + +function fixture(value = response(), status = 200) { + const calls = []; + return { + calls, + fetch: async (url, init) => { + calls.push({ url, init, body: JSON.parse(init.body) }); + return json(value, status); + }, + }; +} + +function chooser(f, extra = {}) { + return fromOllama({ model: "test-model", fetch: f.fetch, ...extra }); +} + +test("scores labels through the native chat endpoint", async () => { + const f = fixture(); + const decision = await chooser(f)(request); + + assert.equal(decision.choice, "billing"); + assert.deepEqual(decision.scores, { sales: -1.2, technical: -2.2, billing: -0.1 }); + assert.deepEqual(decision.usage, { + promptTokens: 120, + cachedTokens: 80, + completionTokens: 1, + requests: 1, + }); + + assert.equal(f.calls.length, 1); + const call = f.calls[0]; + assert.equal(call.url, "http://127.0.0.1:11434/api/chat"); + assert.equal(call.init.method, "POST"); + assert.equal(new Headers(call.init.headers).get("content-type"), "application/json"); + const { messages, ...body } = call.body; + assert.deepEqual(body, { + model: "test-model", + stream: false, + think: false, + logprobs: true, + top_logprobs: 20, + options: { num_predict: 1 }, + }); + assert.deepEqual(messages.map(({ role }) => role), ["user"]); + assert.ok(messages[0].content.startsWith(request.context)); +}); + +test("normalizes an explicit API base URL", async () => { + const f = fixture(); + await chooser(f, { baseURL: "https://example.test/ollama/api" })(request); + assert.equal(f.calls[0].url, "https://example.test/ollama/api/chat"); +}); + +test("rejects invalid configuration before sending a request", () => { + const f = fixture(); + assert.throws(() => fromOllama({ model: "", fetch: f.fetch }), /model/i); + assert.throws(() => chooser(f, { baseURL: "https://user:secret@example.test" }), /baseURL/i); + assert.equal(f.calls.length, 0); +}); + +test("supports at most 20 choices", async () => { + const labels = [..."ABCDEFGHIJKLMNOPQRST"]; + const choices = Object.fromEntries(labels.map((label) => [label, `Choice ${label}`])); + const top = labels.slice(1).map((token, index) => ({ + token, + bytes: [token.charCodeAt(0)], + logprob: -(index + 2), + })); + const f = fixture(response(top, { + logprobs: [{ token: "A", bytes: [65], logprob: -1, top_logprobs: top }], + })); + + const decision = await chooser(f)({ ...request, choices }); + assert.equal(Object.keys(decision.distribution).length, 20); + + const tooMany = { ...choices, U: "Choice U" }; + await assert.rejects(chooser(f)({ ...request, choices: tooMany }), /at most 20/i); + assert.equal(f.calls.length, 1); +}); + +test("assigns zero probability to omitted labels and rejects no label scores", async () => { + const oneMissing = fixture(response(alternatives.filter(({ token }) => token !== "B"), { + prompt_eval_cached_count: undefined, + })); + const decision = await chooser(oneMissing)(request); + assert.equal(decision.scores.technical, -Infinity); + assert.equal(decision.distribution.technical, 0); + assert.equal(decision.usage.cachedTokens, null); + + const noLabels = fixture(response([ + { token: "x", bytes: [120], logprob: -0.1 }, + ], { + logprobs: [{ + token: "y", bytes: [121], logprob: -0.2, + top_logprobs: [{ token: "x", bytes: [120], logprob: -0.1 }], + }], + })); + await assert.rejects(chooser(noLabels)(request), + (error) => error instanceof ScoringError && /choice label/i.test(error.message)); +}); + +test("passes AbortSignal to fetch and preserves its reason", async () => { + const controller = new AbortController(); + const reason = new Error("cancelled by caller"); + const f = fixture(); + f.fetch = async (_url, init) => { + assert.equal(init.signal, controller.signal); + controller.abort(reason); + throw reason; + }; + + await assert.rejects(chooser(f)({ ...request, signal: controller.signal }), + (error) => error === reason); +}); + +const malformed = [ + ["null response", null, /scored token/i], + ["missing logprobs", response(undefined, { logprobs: undefined }), /scored token/i], + ["multiple scored positions", response(undefined, { + logprobs: [response().logprobs[0], response().logprobs[0]], + }), /one|position|logprobs/i], + ["missing top logprobs", response(undefined, { + logprobs: [{ token: "A", bytes: [65], logprob: -1.2 }], + }), /top.*logprobs/i], + ["conflicting selected token", response([ + ...alternatives, { token: "A", bytes: [65], logprob: -1.3 }, + ]), /conflict|duplicate/i], + ["invalid logprob", response(alternatives.map((entry) => + entry.token === "C" ? { ...entry, logprob: 0.1 } : entry)), /logprob/i], + ["mismatched bytes", response(alternatives.map((entry) => + entry.token === "C" ? { ...entry, bytes: [99] } : entry)), /bytes/i], + ["invalid cached usage", response(undefined, { prompt_eval_cached_count: 121 }), /cache|usage/i], +]; + +for (const [name, value, pattern] of malformed) { + test(`rejects ${name}`, async () => { + const f = fixture(value); + await assert.rejects(chooser(f)(request), + (error) => error instanceof ScoringError && pattern.test(error.message)); + }); +} + +test("reports HTTP and JSON failures without exposing the prompt", async () => { + const failed = fixture({ error: request.context }, 500); + await assert.rejects(chooser(failed)(request), (error) => { + assert.match(error.message, /500/); + assert.doesNotMatch(error.message, new RegExp(request.context)); + return true; + }); + + const invalidJson = fixture(); + invalidJson.fetch = async () => new Response("{", { status: 200 }); + await assert.rejects(chooser(invalidJson)(request), + (error) => error instanceof ScoringError && /invalid JSON/i.test(error.message)); +}); diff --git a/tests/package.test.mjs b/tests/package.test.mjs index cbccead..55ace96 100644 --- a/tests/package.test.mjs +++ b/tests/package.test.mjs @@ -14,6 +14,8 @@ test("ESM and CommonJS public exports implement the same API", async () => { assert.deepEqual(a, b); assert.equal(typeof (await import("choosekit/llama-cpp")).fromLlamaCpp, "function"); assert.equal(typeof createRequire(import.meta.url)("choosekit/llama-cpp").fromLlamaCpp, "function"); + assert.equal(typeof (await import("choosekit/ollama")).fromOllama, "function"); + assert.equal(typeof createRequire(import.meta.url)("choosekit/ollama").fromOllama, "function"); assert.equal(typeof (await import("choosekit/openrouter")).fromOpenRouter, "function"); assert.equal(typeof createRequire(import.meta.url)("choosekit/openrouter").fromOpenRouter, "function"); }); @@ -34,6 +36,7 @@ test("importing public entrypoints does not call fetch or log anything", () => { globalThis.fetch = () => { throw new Error("Unexpected network call"); }; await import("choosekit"); await import("choosekit/llama-cpp"); + await import("choosekit/ollama"); await import("choosekit/openrouter"); `; const result = spawnSync(process.execPath, ["--input-type=module", "-e", code], diff --git a/tests/types/api.cts b/tests/types/api.cts index 18c78ea..19e56bf 100644 --- a/tests/types/api.cts +++ b/tests/types/api.cts @@ -1,8 +1,10 @@ import { createChooser, type Scorer } from "choosekit"; import { fromLlamaCpp } from "choosekit/llama-cpp"; +import { fromOllama } from "choosekit/ollama"; import { fromOpenRouter } from "choosekit/openrouter"; const score: Scorer = ({ candidates }) => ({ logprobs: candidates.map(() => -1) }); const choose = createChooser(score); void choose({ context: "", question: "Next?", choices: { test: "Test", done: "Done" } }); void fromLlamaCpp({ baseURL: "http://127.0.0.1:8080" }); +void fromOllama({ model: "test-model" }); void fromOpenRouter({ apiKey: "test-key", model: "test/model" }); diff --git a/tests/types/api.mts b/tests/types/api.mts index 290b0fe..e697ffa 100644 --- a/tests/types/api.mts +++ b/tests/types/api.mts @@ -1,5 +1,6 @@ import { createChooser, type Scorer, type Decision } from "choosekit"; import { fromLlamaCpp } from "choosekit/llama-cpp"; +import { fromOllama } from "choosekit/ollama"; import { fromOpenRouter } from "choosekit/openrouter"; const scorer: Scorer = async ({ candidates, signal }) => { @@ -31,6 +32,8 @@ void numericKey; const local = fromLlamaCpp({ baseURL: "http://127.0.0.1:8080", mode: "labels" }); void local({ context: "", question: "?", choices: { yes: "Yes", no: "No" } }); +const ollama = fromOllama({ model: "test-model" }); +void ollama({ context: "", question: "?", choices: { yes: "Yes", no: "No" } }); const remote = fromOpenRouter({ apiKey: "test-key", model: "test/model" }); void remote({ context: "", question: "?", choices: { yes: "Yes", no: "No" } });