diff --git a/CHANGELOG.md b/CHANGELOG.md index cac74fb..b6c24b9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,8 @@ ## Unreleased +- Added image inputs for vision-capable llama.cpp, Ollama, and OpenRouter models. + llama.cpp image scoring uses `labels` mode. - 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 acdf17b..3d0a666 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 models through llama.cpp and Ollama, plus an optional OpenRouter backend. +`choosekit` scores a finite set of choices and returns a typed decision with a probability distribution. It accepts `text and images` through llama.cpp, Ollama, and OpenRouter. ![SuperGPQA direct-choice benchmark](benchmarks/supergpqa-benchmark.svg) @@ -108,6 +108,29 @@ Returned probabilities are normalized across the supplied choices and are not ca OpenRouter may route the same model through different providers. Set `provider: "provider-slug"` to use only that provider and disable fallback. +## Image inputs + +llama.cpp, Ollama, and OpenRouter can score choices from images when the selected model supports vision. Pass raw base64 data with its media type: + +```ts +import { readFile } from "node:fs/promises"; + +const decision = await choose({ + context: "Inspect the attached screenshot.", + question: "Which state is the interface in?", + choices: { + ready: "The interface is ready for input.", + loading: "The interface is still loading.", + }, + images: [{ + mediaType: "image/png", + base64: (await readFile("screenshot.png")).toString("base64"), + }], +}); +``` + +Supported media types are PNG, JPEG, and WebP. llama.cpp image inputs currently support `labels` mode only. + ## Scoring modes | Mode | Candidate representation | Use when | diff --git a/src/index.ts b/src/index.ts index 0414630..e3a00fd 100644 --- a/src/index.ts +++ b/src/index.ts @@ -3,7 +3,7 @@ import type { Chooser, ChooserOptions, Scorer } from "./types.js"; export type { Choices, ChoiceKey, ChoiceRequest, Chooser, ChooserOptions, Decision, - ScoreRequest, Scorer, Scores, Usage, PromptInput, + ImageInput, ImageMediaType, ScoreRequest, Scorer, Scores, Usage, PromptInput, } from "./types.js"; export { ScoringError } from "./validation.js"; diff --git a/src/internal-chooser.ts b/src/internal-chooser.ts index 375199a..0a987fa 100644 --- a/src/internal-chooser.ts +++ b/src/internal-chooser.ts @@ -1,11 +1,28 @@ import type { - Choices, ChoiceKey, ChoiceRequest, Chooser, ChooserOptions, Decision, Scorer, Usage, + Choices, ChoiceKey, ChoiceRequest, Chooser, ChooserOptions, Decision, ImageInput, Scorer, Usage, } from "./types.js"; import { isCount, isLogprob, isRecord, requireText, ScoringError } from "./validation.js"; export type CandidateFormat = "keys" | "labels"; const LABELS = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"; +const IMAGE_MEDIA_TYPES = new Set(["image/png", "image/jpeg", "image/webp"]); + +function snapshotImages(value: unknown): readonly ImageInput[] | undefined { + if (value === undefined) return undefined; + if (!Array.isArray(value)) throw new TypeError("images must be an array."); + return Object.freeze(value.map((image, index) => { + if (!isRecord(image)) throw new TypeError(`images[${index}] must be an object.`); + if (!IMAGE_MEDIA_TYPES.has(image.mediaType as string)) { + throw new TypeError(`images[${index}].mediaType is not supported.`); + } + requireText(image.base64, `images[${index}].base64`); + return Object.freeze({ + mediaType: image.mediaType as ImageInput["mediaType"], + base64: image.base64, + }); + })); +} function snapshot(choices: unknown): [string, string][] { if (!isRecord(choices) || Object.getOwnPropertySymbols(choices).length !== 0) { @@ -109,6 +126,7 @@ export function createFormattedChooser(score: Scorer, options: ChooserOptions, if (typeof context !== "string") throw new TypeError("context must be a string."); requireText(question, "question"); const entries = snapshot(choices); + const images = snapshotImages(request.images); const keys = entries.map(([key]) => key) as ChoiceKey[]; const prepared = instruction(question, entries, candidateFormat); const prompt = formatPrompt @@ -122,7 +140,9 @@ export function createFormattedChooser(score: Scorer, options: ChooserOptions, } signal?.throwIfAborted(); const scored = await score(Object.freeze({ - prompt, candidates: prepared.candidates, ...(signal ? { signal } : {}), + prompt, candidates: prepared.candidates, + ...(images === undefined ? {} : { images }), + ...(signal ? { signal } : {}), })); signal?.throwIfAborted(); if (!isRecord(scored) || !Array.isArray(scored.logprobs) diff --git a/src/llama-cpp.ts b/src/llama-cpp.ts index 90112a5..fe62fa1 100644 --- a/src/llama-cpp.ts +++ b/src/llama-cpp.ts @@ -12,13 +12,19 @@ export interface LlamaCppOptions extends ChooserOptions { readonly model?: string; readonly headers?: Readonly>; readonly fetch?: typeof globalThis.fetch; - /** Must match the agent's tokenizer setting. Defaults to false for serialized context. */ + /** Controls add_special for llama.cpp tokenization. Defaults to false; image inputs require false. */ readonly addSpecialTokens?: boolean; } interface Endpoints { readonly tokenize: string; + readonly detokenize: string; readonly completion: string; + readonly props: string; +} + +interface ImageSupport { + readonly marker: string; } interface Probe { @@ -36,7 +42,7 @@ interface Branch { readonly score: number; } -function endpoints(baseURL: string): Endpoints { +function endpoints(baseURL: string, model?: string): Endpoints { requireText(baseURL, "baseURL"); const url = new URL(baseURL); if ((url.protocol !== "http:" && url.protocol !== "https:") || url.username || url.password @@ -46,8 +52,13 @@ function endpoints(baseURL: string): Endpoints { const path = url.pathname.replace(/\/v1\/?$/, "").replace(/\/$/, ""); url.pathname = `${path}/tokenize`; const tokenize = url.href; + url.pathname = `${path}/detokenize`; + const detokenize = url.href; url.pathname = `${path}/completion`; - return { tokenize, completion: url.href }; + const completion = url.href; + url.pathname = `${path}/props`; + if (model !== undefined) url.searchParams.set("model", model); + return { tokenize, detokenize, completion, props: url.href }; } function requestHeaders(value: unknown): Readonly> { @@ -91,12 +102,54 @@ async function post(fetchImpl: typeof globalThis.fetch, url: string, } } +async function get(fetchImpl: typeof globalThis.fetch, url: string, + headers: Readonly>, signal?: AbortSignal): Promise { + signal?.throwIfAborted(); + let response: Response; + try { + response = await fetchImpl(url, { method: "GET", headers, ...(signal ? { signal } : {}) }); + } catch (error) { + signal?.throwIfAborted(); + throw error; + } + signal?.throwIfAborted(); + if (!response.ok) { + throw new ScoringError(`llama.cpp 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(`llama.cpp returned invalid JSON for ${new URL(url).pathname}.`); + } +} + function parseTokenization(value: unknown): number[] { if (!isRecord(value)) throw new ScoringError("llama.cpp returned an invalid tokenization."); return tokenIds(value.tokens); } -function parseProbe(value: unknown, prompt: readonly number[], targetTokenId: number): Probe { +function parseDetokenization(value: unknown): string { + if (!isRecord(value) || typeof value.content !== "string") { + throw new ScoringError("llama.cpp returned an invalid detokenization."); + } + return value.content; +} + +function parseImageSupport(value: unknown): ImageSupport { + if (!isRecord(value) || !isRecord(value.modalities) || value.modalities.vision !== true) { + throw new ScoringError("llama.cpp does not advertise vision support for this model."); + } + if (typeof value.media_marker !== "string" || value.media_marker.length === 0) { + throw new ScoringError("llama.cpp did not return a multimodal media marker."); + } + return Object.freeze({ marker: value.media_marker }); +} + +function parseProbe(value: unknown, expectedPromptTokens: number | null, + targetTokenId: number): Probe { if (!isRecord(value)) throw new ScoringError("llama.cpp returned an invalid completion."); if (value.truncated === true) { throw new ScoringError("llama.cpp truncated the token prefix while scoring a candidate."); @@ -144,7 +197,10 @@ function parseProbe(value: unknown, prompt: readonly number[], targetTokenId: nu } const promptTokens: unknown = value.tokens_evaluated; - if (!isCount(promptTokens) || promptTokens !== prompt.length) { + if (!isCount(promptTokens)) { + throw new ScoringError("llama.cpp returned invalid prompt-token usage."); + } + if (expectedPromptTokens !== null && promptTokens !== expectedPromptTokens) { throw new ScoringError("llama.cpp did not evaluate the supplied numeric token prefix as sent."); } let completionTokens = output.length; @@ -178,11 +234,20 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser { } const fetchImpl = options.fetch ?? globalThis.fetch; if (typeof fetchImpl !== "function") throw new TypeError("A fetch implementation is required."); - const urls = endpoints(baseURL); + const urls = endpoints(baseURL, model); const headers = requestHeaders(options.headers); - - const score: Scorer = async ({ prompt, candidates, signal }) => { + const score: Scorer = async ({ prompt, candidates, images, signal }) => { + const hasImages = images !== undefined && images.length > 0; + if (hasImages && mode !== "labels") { + throw new TypeError("llama.cpp image inputs require labels mode."); + } + if (hasImages && addSpecialTokens) { + throw new TypeError("llama.cpp image inputs require addSpecialTokens to be false."); + } signal?.throwIfAborted(); + const imageSupport = hasImages + ? parseImageSupport(await get(fetchImpl, urls.props, headers, signal)) + : undefined; const encoded: number[][] = []; for (const content of [prompt, ...candidates.map((candidate) => prompt + candidate)]) { const response = await post(fetchImpl, urls.tokenize, headers, { @@ -199,7 +264,7 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser { let promptTokens = 0; let cachedTokens: number | null = 0; let completionTokens = 0; - let requests = encoded.length; + let requests = encoded.length + (hasImages ? 1 : 0); let root = treeRoot; const rootSuffix: number[] = []; @@ -209,6 +274,33 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser { rootSuffix.push(first[0]); root = first[1]; } + const materializeImagePrompt = async (numericPrefix: readonly number[]) => { + const response = await post(fetchImpl, urls.detokenize, headers, { + tokens: numericPrefix, ...(model === undefined ? {} : { model }), + }, signal); + requests++; + const detokenized = parseDetokenization(response); + if (detokenized.includes(imageSupport!.marker)) { + throw new ScoringError("The formatted prompt contains llama.cpp's multimodal media marker."); + } + + const roundTripResponse = await post(fetchImpl, urls.tokenize, headers, { + content: detokenized, add_special: false, + ...(model === undefined ? {} : { model }), + }, signal); + requests++; + const roundTrip = parseTokenization(roundTripResponse); + if (roundTrip.length !== numericPrefix.length + || roundTrip.some((tokenId, index) => tokenId !== numericPrefix[index])) { + throw new ScoringError("llama.cpp could not preserve the image prompt token prefix."); + } + + const value = Object.freeze({ + prompt_string: `${images!.map(() => imageSupport!.marker).join("\n")}\n${detokenized}`, + multimodal_data: Object.freeze(images!.map(({ base64 }) => base64)), + }); + return value; + }; const work: Branch[] = [{ node: root, indices: candidates.map((_, index) => index), suffix: rootSuffix, score: 0, @@ -229,12 +321,13 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser { const children = [...branch.node.children]; const siblingLogprobs = new Map(); const numericPrefix = [...base, ...branch.suffix]; + const promptValue = hasImages ? await materializeImagePrompt(numericPrefix) : numericPrefix; for (const [targetTokenId] of children) { if (siblingLogprobs.has(targetTokenId)) continue; const collectSiblings = siblingLogprobs.size === 0 && children.length > 1; signal?.throwIfAborted(); const response = await post(fetchImpl, urls.completion, headers, { - prompt: numericPrefix, + prompt: promptValue, ...(model === undefined ? {} : { model }), n_predict: 1, n_probs: collectSiblings ? 64 : 1, @@ -249,7 +342,7 @@ export function fromLlamaCpp(options: LlamaCppOptions): Chooser { cache_prompt: true, }, signal); signal?.throwIfAborted(); - const probe = parseProbe(response, numericPrefix, targetTokenId); + const probe = parseProbe(response, hasImages ? null : numericPrefix.length, targetTokenId); promptTokens += probe.promptTokens; cachedTokens = cachedTokens === null || probe.cachedTokens === null ? null : cachedTokens + probe.cachedTokens; diff --git a/src/ollama.ts b/src/ollama.ts index 3bae4c1..379c21a 100644 --- a/src/ollama.ts +++ b/src/ollama.ts @@ -130,13 +130,19 @@ export function fromOllama(options: OllamaOptions): Chooser { if (typeof fetchImpl !== "function") throw new TypeError("A fetch implementation is required."); const url = endpoint(options.baseURL); - const score: Scorer = async ({ prompt, candidates, signal }) => { + const score: Scorer = async ({ prompt, candidates, images, 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 }], + messages: [{ + role: "user", + content: prompt, + ...(images === undefined || images.length === 0 + ? {} + : { images: images.map(({ base64 }) => base64) }), + }], stream: false, think: false, logprobs: true, diff --git a/src/openrouter.ts b/src/openrouter.ts index ce517b8..96a8f44 100644 --- a/src/openrouter.ts +++ b/src/openrouter.ts @@ -154,13 +154,22 @@ export function fromOpenRouter(options: OpenRouterOptions): Chooser { const fetchImpl = options.fetch ?? globalThis.fetch; if (typeof fetchImpl !== "function") throw new TypeError("A fetch implementation is required."); - const score: Scorer = async ({ prompt, candidates, signal }) => { + const score: Scorer = async ({ prompt, candidates, images, signal }) => { if (candidates.length > MAX_CANDIDATES) { throw new TypeError(`OpenRouter supports at most ${MAX_CANDIDATES} choices.`); } const response = await post(fetchImpl, apiKey, { model, - messages: [{ role: "user", content: prompt }], + messages: [{ + role: "user", + content: images === undefined || images.length === 0 ? prompt : [ + { type: "text", text: prompt }, + ...images.map(({ mediaType, base64 }) => ({ + type: "image_url", + image_url: { url: `data:${mediaType};base64,${base64}` }, + })), + ], + }], max_tokens: 1, stream: false, temperature: 1, diff --git a/src/types.ts b/src/types.ts index e5aea0b..d8e1240 100644 --- a/src/types.ts +++ b/src/types.ts @@ -1,11 +1,20 @@ export type Choices = Readonly>; export type ChoiceKey = `${Extract}`; +export type ImageMediaType = "image/png" | "image/jpeg" | "image/webp"; + +export interface ImageInput { + readonly mediaType: ImageMediaType; + /** Raw base64 data without a data URL prefix. */ + readonly base64: string; +} + export interface ChoiceRequest { /** Existing context, copied unchanged to the beginning of the scoring prompt. */ readonly context: string; readonly question: string; readonly choices: C; + readonly images?: readonly ImageInput[]; readonly signal?: AbortSignal; } @@ -35,6 +44,7 @@ export interface Decision { export interface ScoreRequest { readonly prompt: string; readonly candidates: readonly string[]; + readonly images?: readonly ImageInput[]; readonly signal?: AbortSignal; } diff --git a/tests/core.test.mjs b/tests/core.test.mjs index 15352e7..bc65e6b 100644 --- a/tests/core.test.mjs +++ b/tests/core.test.mjs @@ -59,6 +59,24 @@ test("scores one complete escaped and terminated key per candidate", async () => assert.equal(called, 1); }); +test("snapshots image inputs before an asynchronous boundary", async () => { + let release; + const gate = new Promise((resolve) => { release = resolve; }); + const image = { mediaType: "image/png", base64: "aW1hZ2U=" }; + const images = [image]; + const choose = createChooser(({ images }) => { + assert.deepEqual(images, [{ mediaType: "image/png", base64: "aW1hZ2U=" }]); + return { logprobs: [-1, -2, -3] }; + }, { + formatPrompt: async ({ context, instruction }) => { await gate; return context + instruction; }, + }); + const pending = choose({ ...request, images }); + image.base64 = "Y2hhbmdlZA=="; + images.push({ mediaType: "image/jpeg", base64: "YWRkZWQ=" }); + release(); + await pending; +}); + test("empty context is supported", async () => { const d = await fixed([-1, -2, -3])({ ...request, context: "" }); assert.equal(d.choice, "edit"); @@ -105,6 +123,9 @@ for (const [name, change] of [ ["blank description", { choices: { a: " ", b: "B" } }], ["numeric description", { choices: { a: 1, b: "B" } }], ["symbol key", { choices: { a: "A", b: "B", [Symbol("c")]: "C" } }], + ["non-array images", { images: {} }], + ["unsupported image type", { images: [{ mediaType: "image/bmp", base64: "YQ==" }] }], + ["empty image data", { images: [{ mediaType: "image/png", base64: "" }] }], ]) { test(`rejects ${name} before invoking the scorer`, async () => { let called = false; diff --git a/tests/llama-cpp.test.mjs b/tests/llama-cpp.test.mjs index 56854a2..3b68c5c 100644 --- a/tests/llama-cpp.test.mjs +++ b/tests/llama-cpp.test.mjs @@ -10,14 +10,22 @@ const choiceRequest = Object.freeze({ choices: Object.freeze({ wait: "Wait for approval.", deploy: "Deploy immediately." }), }); -function fixture({ tokenizer = encode, logprob = () => -1, - topTokenIds = (target) => [target], transform } = {}) { +const decode = (tokens) => new TextDecoder().decode( + Uint8Array.from(tokens, (tokenId) => tokenId - 1), +); + +function fixture({ tokenizer = encode, detokenizer = decode, logprob = () => -1, + topTokenIds = (target) => [target], transform, + props = { modalities: { vision: true }, media_marker: "<__media_test__>" } } = {}) { const calls = []; const fetch = async (url, init) => { - const path = new URL(url).pathname; - const body = JSON.parse(init.body); - calls.push({ path, body, headers: init.headers }); + const parsedURL = new URL(url); + const path = parsedURL.pathname; + const body = init.body === undefined ? undefined : JSON.parse(init.body); + calls.push({ url: parsedURL, path, body, headers: init.headers, signal: init.signal }); + if (path === "/props") return json(props); if (path === "/tokenize") return json({ tokens: tokenizer(body.content, body.add_special) }); + if (path === "/detokenize") return json({ content: detokenizer(body.tokens) }); if (path !== "/completion") return json({}, 404); const target = body.logit_bias[0][0]; @@ -29,7 +37,7 @@ function fixture({ tokenizer = encode, logprob = () => -1, top_logprobs: topTokenIds(target, body).map((id) => ({ id, logprob: logprob(id) })), }], generation_settings: { post_sampling_probs: false, backend_sampling: false }, - tokens_evaluated: body.prompt.length, + tokens_evaluated: Array.isArray(body.prompt) ? body.prompt.length : 123, tokens_predicted: 1, timings: { cache_n: 0 }, }; @@ -67,6 +75,212 @@ test("maps label scores back to choice keys and reuses sibling logprobs", async assert.equal(completions(f)[0].body.model, "local-model"); }); +test("scores image labels through native multimodal completion", async () => { + const a = encode("A")[0]; + const b = encode("B")[0]; + const scores = new Map([[a, -0.2], [b, -1.5]]); + const f = fixture({ + logprob: (id) => scores.get(id) ?? -10, + topTokenIds: (target) => [target], + }); + const controller = new AbortController(); + const choose = fromLlamaCpp({ + baseURL: "http://localhost:8080/v1", + model: "vision/model", + headers: { authorization: "Bearer test" }, + fetch: f.fetch, + }); + + const decision = await choose({ + ...choiceRequest, + images: [ + { mediaType: "image/png", base64: "Zmlyc3Q=" }, + { mediaType: "image/jpeg", base64: "c2Vjb25k" }, + ], + signal: controller.signal, + }); + + assert.equal(decision.choice, "wait"); + assert.deepEqual(decision.scores, { wait: -0.2, deploy: -1.5 }); + assert.equal(decision.usage.promptTokens, 246); + assert.equal(decision.usage.requests, f.calls.length); + const propsCall = f.calls.find((call) => call.path === "/props"); + assert.equal(propsCall.url.searchParams.get("model"), "vision/model"); + assert.equal(propsCall.signal, controller.signal); + assert.equal(propsCall.headers.authorization, "Bearer test"); + assert.equal(completions(f).length, 2); + const detokenizeCall = f.calls.find((call) => call.path === "/detokenize"); + const promptString = `<__media_test__>\n<__media_test__>\n${decode(detokenizeCall.body.tokens)}`; + for (const call of completions(f)) { + assert.deepEqual(call.body.prompt, { + prompt_string: promptString, + multimodal_data: ["Zmlyc3Q=", "c2Vjb25k"], + }); + assert.equal(call.body.model, "vision/model"); + assert.equal(call.body.post_sampling_probs, false); + assert.equal(call.body.backend_sampling, false); + assert.equal(call.signal, controller.signal); + } + assert.equal(completions(f)[1].body.n_probs, 1); +}); + +test("preserves a tokenization-boundary rollback for image labels", async () => { + const pairs = (text) => { + const bytes = new TextEncoder().encode(text); + const ids = []; + for (let i = 0; i < bytes.length; i += 2) { + ids.push(i + 1 < bytes.length ? 257 + bytes[i] * 256 + bytes[i + 1] : bytes[i] + 1); + } + return ids; + }; + const a = pairs("abcA").at(-1); + const b = pairs("abcB").at(-1); + const scores = new Map([[a, -0.1], [b, -1.1]]); + const f = fixture({ + tokenizer: pairs, + detokenizer: (tokens) => { + assert.deepEqual(tokens, pairs("ab")); + return "ab"; + }, + logprob: (id) => scores.get(id) ?? -10, + topTokenIds: () => [a, b], + }); + const choose = fromLlamaCpp({ + baseURL: "http://localhost:8080", + fetch: f.fetch, + formatPrompt: ({ context }) => context, + }); + + const decision = await choose({ + context: "abc", + question: "Which?", + choices: { first: "A", second: "B" }, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }); + + assert.equal(decision.choice, "first"); + assert.equal(decision.boundaryTokens, 1); + assert.equal(completions(f)[0].body.prompt.prompt_string, "<__media_test__>\nab"); +}); + +test("scores image labels that share an initial token", async () => { + const tokenizations = new Map([ + ["p", [1]], + ["pA", [1, 100, 101]], + ["pB", [1, 100, 102]], + ["pC", [1, 200]], + ["px", [1, 100]], + ]); + const textByTokens = new Map([["1", "p"], ["1,100", "px"]]); + const scores = new Map([[100, -0.2], [200, -2], [101, -0.3], [102, -1]]); + const f = fixture({ + tokenizer: (text) => tokenizations.get(text), + detokenizer: (tokens) => textByTokens.get(tokens.join(",")), + logprob: (id) => scores.get(id) ?? -10, + }); + const choose = fromLlamaCpp({ + baseURL: "http://localhost:8080", + fetch: f.fetch, + formatPrompt: ({ context }) => context, + }); + + const decision = await choose({ + context: "p", + question: "Which?", + choices: { first: "First", second: "Second", third: "Third" }, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }); + + assert.equal(decision.choice, "first"); + assert.deepEqual(decision.scores, { first: -0.5, second: -1.2, third: -2 }); + assert.deepEqual( + [...new Set(completions(f).map((call) => call.body.prompt.prompt_string))], + ["<__media_test__>\np", "<__media_test__>\npx"], + ); +}); + +test("rejects an image prefix that does not survive a tokenization round trip", async () => { + const f = fixture({ detokenizer: () => "different text" }); + const choose = fromLlamaCpp({ baseURL: "http://localhost:8080", fetch: f.fetch }); + + await assert.rejects(choose({ + ...choiceRequest, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }), /could not preserve the image prompt token prefix/i); + assert.equal(completions(f).length, 0); +}); + +test("rejects a formatted prompt containing llama.cpp's media marker", async () => { + const marker = "<__media_test__>"; + const f = fixture({ detokenizer: () => `prompt ${marker}` }); + const choose = fromLlamaCpp({ baseURL: "http://localhost:8080", fetch: f.fetch }); + + await assert.rejects(choose({ + ...choiceRequest, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }), /formatted prompt contains llama\.cpp's multimodal media marker/i); + assert.equal(completions(f).length, 0); +}); + +test("rejects image inputs when llama.cpp does not advertise vision", async () => { + const f = fixture({ props: { modalities: { vision: false }, media_marker: "" } }); + const choose = fromLlamaCpp({ baseURL: "http://localhost:8080", fetch: f.fetch }); + + await assert.rejects(choose({ + ...choiceRequest, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }), /does not advertise vision support/i); + assert.deepEqual(f.calls.map((call) => call.path), ["/props"]); +}); + +test("rejects image inputs when llama.cpp omits the media marker", async () => { + const f = fixture({ props: { modalities: { vision: true } } }); + const choose = fromLlamaCpp({ baseURL: "http://localhost:8080", fetch: f.fetch }); + + await assert.rejects(choose({ + ...choiceRequest, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }), /multimodal media marker/i); + assert.deepEqual(f.calls.map((call) => call.path), ["/props"]); +}); + +test("rejects an invalid llama.cpp detokenization response", async () => { + const f = fixture({ detokenizer: () => 42 }); + const choose = fromLlamaCpp({ baseURL: "http://localhost:8080", fetch: f.fetch }); + + await assert.rejects(choose({ + ...choiceRequest, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }), /invalid detokenization/i); + assert.equal(completions(f).length, 0); +}); + +test("rejects image inputs in minimal-prefix mode before contacting llama.cpp", async () => { + const f = fixture(); + const choose = fromLlamaCpp({ + baseURL: "http://localhost:8080", mode: "minimal-prefix", fetch: f.fetch, + }); + + await assert.rejects(choose({ + ...choiceRequest, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }), /require labels mode/i); + assert.equal(f.calls.length, 0); +}); + +test("rejects special-token insertion for llama.cpp image inputs", async () => { + const f = fixture(); + const choose = fromLlamaCpp({ + baseURL: "http://localhost:8080", addSpecialTokens: true, fetch: f.fetch, + }); + + await assert.rejects(choose({ + ...choiceRequest, + images: [{ mediaType: "image/png", base64: "aW1hZ2U=" }], + }), /require addSpecialTokens to be false/i); + assert.equal(f.calls.length, 0); +}); + test("scores candidates separately when top logprobs contain none of them", async () => { const scores = new Map([[encode("A")[0], -0.3], [encode("B")[0], -1.4]]); const f = fixture({ diff --git a/tests/ollama.test.mjs b/tests/ollama.test.mjs index 0fa4efa..6c0d665 100644 --- a/tests/ollama.test.mjs +++ b/tests/ollama.test.mjs @@ -97,6 +97,17 @@ test("normalizes an explicit API base URL", async () => { assert.equal(f.calls[0].url, "https://example.test/ollama/api/chat"); }); +test("sends raw base64 image data to Ollama", async () => { + const f = fixture(); + await chooser(f)({ + ...request, + images: [{ mediaType: "image/webp", base64: "d2VicA==" }], + }); + + assert.deepEqual(f.calls[0].body.messages[0].images, ["d2VicA=="]); + assert.ok(f.calls[0].body.messages[0].content.startsWith(request.context)); +}); + test("rejects invalid configuration before sending a request", () => { const f = fixture(); assert.throws(() => fromOllama({ model: "", fetch: f.fetch }), /model/i); diff --git a/tests/openrouter.test.mjs b/tests/openrouter.test.mjs index cdd95c0..b904cef 100644 --- a/tests/openrouter.test.mjs +++ b/tests/openrouter.test.mjs @@ -128,6 +128,25 @@ test("pins an optional provider without fallback", async () => { }); }); +test("sends image inputs as OpenRouter content parts", async () => { + const f = fixture(); + await chooser(f)({ + ...request, + images: [ + { mediaType: "image/png", base64: "cG5n" }, + { mediaType: "image/jpeg", base64: "anBlZw==" }, + ], + }); + + const content = f.calls[0].body.messages[0].content; + assert.equal(content[0].type, "text"); + assert.ok(content[0].text.startsWith(request.context)); + assert.deepEqual(content.slice(1), [ + { type: "image_url", image_url: { url: "data:image/png;base64,cG5n" } }, + { type: "image_url", image_url: { url: "data:image/jpeg;base64,anBlZw==" } }, + ]); +}); + test("assigns zero probability to labels omitted from top logprobs", async () => { const f = fixture(response(unorderedTopLogprobs.filter(({ token }) => token !== "C"))); const decision = await chooser(f)(request); diff --git a/tests/types/api.mts b/tests/types/api.mts index e697ffa..59996bf 100644 --- a/tests/types/api.mts +++ b/tests/types/api.mts @@ -1,4 +1,4 @@ -import { createChooser, type Scorer, type Decision } from "choosekit"; +import { createChooser, type Scorer, type Decision, type ImageInput } from "choosekit"; import { fromLlamaCpp } from "choosekit/llama-cpp"; import { fromOllama } from "choosekit/ollama"; import { fromOpenRouter } from "choosekit/openrouter"; @@ -35,7 +35,11 @@ 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" } }); +const image: ImageInput = { mediaType: "image/png", base64: "aW1hZ2U=" }; +void remote({ context: "", question: "?", choices: { yes: "Yes", no: "No" }, images: [image] }); +void ollama({ context: "", question: "?", choices: { yes: "Yes", no: "No" }, images: [image] }); +// @ts-expect-error Image MIME types are restricted to supported formats. +void remote({ context: "", question: "?", choices: { yes: "Yes", no: "No" }, images: [{ mediaType: "image/bmp", base64: "YQ==" }] }); // @ts-expect-error The old implicit agent-state input is not part of this API. choose({ state: "Changed", question: "Next?", choices: { yes: "Yes", no: "No" } });