diff --git a/.changeset/model-type-extrabody.md b/.changeset/model-type-extrabody.md new file mode 100644 index 0000000..2fcbb89 --- /dev/null +++ b/.changeset/model-type-extrabody.md @@ -0,0 +1,44 @@ +--- +'smart-decisions': minor +--- + +**Breaking (pre-1.0):** the provider settings on `Question` moved into a new `Model` type. + +- `apiBaseUrl`, `apiKey` and `model` no longer sit flat on `Question`; they are now + `question.model.{apiBaseUrl, apiKey, model}`. This groups everything about how the + model is reached and served into one reusable object to pass to every `choice()`. +- `Model` gains `extraBody`: extra fields forwarded verbatim into the chat completions + request body for engine- or model-specific settings (i.e. `chat_template_kwargs` + thinking toggles on llama.cpp/vLLM/SGLang, `reasoning_effort` on OpenAI/OpenRouter/ + Ollama, Ollama's native `think`). Reserved request keys the library's single-token + trick depends on (`model`, `messages`, `stream`, `logprobs`, `top_logprobs`, + `max_tokens`, `temperature`) cannot be overridden through it; `chat_template_kwargs` + merges one level deep with user keys winning per key. +- System 1's request now asks for `top_logprobs: 20` instead of 50: 20 is the highest + portable window (OpenAI and OpenRouter cap it at 20, vLLM's server default + `--max-logprobs` is 20). On llama.cpp (accepts up to 50) the smaller window is + enough for realistic option counts. +- New `Model` type exported from the package root; `generateText` accepts extra + top-level body fields via `ChatCompletionRequest`'s index signature. + +### Hardening + +- The transport no longer follows redirects (`redirect: 'error'`): the Bearer token + stays off unexpected paths, and a misconfigured base URL fails loudly. +- Response bodies are read under a 10 MB safety cap instead of being buffered + unconditionally; an over-cap body fails fast and is not retried. A declared + `Content-Length` over the cap is refused before reading anything. +- Non-finite `maxRetries` (i.e. `NaN`) no longer causes an infinite retry loop: it + falls back to the default, negatives clamp to 0 and fractions to whole attempts. +- `extraBody` ignores `__proto__` and `constructor` keys, so prototype-smuggled + config can never re-parent the request object. +- Malformed entries inside `top_logprobs` (missing/null token) are skipped instead + of crashing mid-read. +- Error messages strip query strings from the request URL, so providers that take + credentials as query parameters cannot leak them into logs. +- The endpoint path is joined through the URL API, preserving a query string in + `apiBaseUrl` instead of swallowing it, and an invalid base URL throws a clear + `Invalid model.apiBaseUrl` error. +- CI actions are pinned by commit SHA; the release script spawns every subprocess + as an argv array (no shell), so interpolated values can never be re-parsed as + shell syntax. diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a270b22..a4259f8 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -22,12 +22,12 @@ jobs: node: [22, 24] # engines: >=22 — test the floor and the current LTS steps: - - uses: actions/checkout@v4 + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 # Pin the PR head SHA so coverage line numbers map onto the diff view. with: ref: ${{ github.event.pull_request.head.sha || github.sha }} - - uses: actions/setup-node@v4 + - uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4 with: node-version: ${{ matrix.node }} cache: npm @@ -46,7 +46,7 @@ jobs: matrix.node == 22 && github.event_name == 'pull_request' && !github.event.pull_request.head.repo.fork - uses: actions/upload-code-coverage@v1 + uses: actions/upload-code-coverage@bfa741d815a28cb064a8e3a0837e577457a017d5 # v1 with: file: coverage/cobertura-coverage.xml language: JavaScript diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index f0ba86c..908a154 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -25,9 +25,9 @@ jobs: release: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v4 + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 - - uses: actions/setup-node@v4 + - uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4 with: node-version: 22 cache: npm @@ -37,7 +37,7 @@ jobs: - run: npm ci - - uses: changesets/action@v1 + - uses: changesets/action@a45c4d594aa4e2c509dc14a9f2b3b67ba3780d0d # v1.9.0 with: # Single command: the action whitespace-splits this input and spawns # it as args (no shell), so chaining belongs in the package script. diff --git a/README.md b/README.md index e310839..1aa174c 100644 --- a/README.md +++ b/README.md @@ -40,10 +40,14 @@ npm install smart-decisions ```typescript import { choice } from 'smart-decisions'; -const answer = await choice({ +const model = { apiBaseUrl: 'http://localhost:8000/v1', // your API url. i.e. your llama.cpp server apiKey: 'a-super-secret-api-key', // as required or not by your provider model: '/models/Qwen3.5-4B-Q4_K_M.gguf', +}; + +const answer = await choice({ + model, state: "It is raining and I am at home. I'm bored.", instructions: 'Give me a good plan to do now', criteria: { @@ -82,6 +86,12 @@ Classify incoming tickets. When the distribution is too flat to trust, the retur ```typescript import { choice } from 'smart-decisions'; +const model = { + apiBaseUrl: 'http://localhost:8000/v1', + apiKey: 'a-super-secret-api-key', + model: '/models/Qwen3.5-4B-Q4_K_M.gguf', +}; + const criteria = { billing: 'Customer asks about invoices, payments, refunds or charges', technical: 'Customer reports a bug, error or product malfunction', @@ -90,9 +100,7 @@ const criteria = { }; const answer = await choice({ - apiBaseUrl: 'http://localhost:8000/v1', - apiKey: 'a-super-secret-api-key', - model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + model, mode: 'system1', state: 'Customer support ticket:\n"Hi, I was charged twice this month. Can you refund the extra payment?"', @@ -114,10 +122,7 @@ routeTo(answer.choice); // 'billing' > ⚠️ **Work in progress.** > Only [llama.cpp](https://github.com/ggml-org/llama.cpp) has been -> tested for now. -> -> Other OpenAI-compatible servers (vLLM, LM Studio, Ollama, …) -> should keep working as long as they return logprobs, but are unverified yet. +> tested for now ## Use case diff --git a/examples/basic.ts b/examples/basic.ts index 9e8f220..7bbfe8f 100644 --- a/examples/basic.ts +++ b/examples/basic.ts @@ -1,9 +1,13 @@ import { choice } from '../src/index.js'; -const answer = await choice({ +const model = { apiBaseUrl: process.env.API_BASE_URL!, apiKey: process.env.API_KEY!, model: '/models/Qwen3.5-4B-Q4_K_M.gguf', +}; + +const answer = await choice({ + model: model, mode: 'system1', state: "It is raining and I am at home. I'm bored.", instructions: 'Give me a good plan to do now', diff --git a/examples/ticket-routing.ts b/examples/ticket-routing.ts index 5ead6a6..69f4dd6 100644 --- a/examples/ticket-routing.ts +++ b/examples/ticket-routing.ts @@ -9,11 +9,15 @@ const criteria = { const ticket = 'Hi, I was charged twice this month. Can you refund the extra payment?'; +const model = { + apiBaseUrl: process.env.API_BASE_URL!, + apiKey: process.env.API_KEY!, + model: '/models/Qwen3.5-4B-Q4_K_M.gguf', +}; + async function main() { const answer = await choice({ - apiBaseUrl: process.env.API_BASE_URL!, - apiKey: process.env.API_KEY!, - model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + model: model, mode: 'system1', state: `Customer support ticket:\n"${ticket}"`, instructions: 'Classify the ticket into the right department', diff --git a/scripts/release.mjs b/scripts/release.mjs index b3a5438..513da0e 100644 --- a/scripts/release.mjs +++ b/scripts/release.mjs @@ -1,4 +1,4 @@ -import { execSync } from 'node:child_process'; +import { execFileSync } from 'node:child_process'; import { mkdtempSync, readFileSync, writeFileSync } from 'node:fs'; import { tmpdir } from 'node:os'; import { join } from 'node:path'; @@ -20,6 +20,10 @@ import { join } from 'node:path'; * exists. Runs on every publish-path run, so a stranded announcement heals * on the next run — necessary because the registry's post-upload validation * window can keep a freshly accepted version invisible to registry checks. + * + * Every subprocess is spawned as an argv array (no shell): interpolated values + * can never be re-parsed as shell syntax, so a tampered package.json version + * cannot turn into command injection. */ const { name, version } = JSON.parse(readFileSync('package.json', 'utf8')); const tag = `v${version}`; @@ -27,11 +31,11 @@ const tag = `v${version}`; const env = { ...process.env }; delete env.GITHUB_TOKEN; -execSync('npm run build', { stdio: 'inherit' }); +execFileSync('npm', ['run', 'build'], { stdio: 'inherit' }); let registered = false; try { - execSync(`npm view ${name}@${version} version`, { env, stdio: 'ignore' }); + execFileSync('npm', ['view', `${name}@${version}`, 'version'], { env, stdio: 'ignore' }); registered = true; } catch { registered = false; @@ -40,7 +44,10 @@ try { if (registered) { console.log(`${name}@${version} is already registered — skipping publish.`); } else { - execSync('npm publish --provenance --access public', { env, stdio: 'inherit' }); + execFileSync('npm', ['publish', '--provenance', '--access', 'public'], { + env, + stdio: 'inherit', + }); console.log(`Published ${name}@${version}`); } @@ -59,17 +66,20 @@ if (!section) { throw new Error(`No CHANGELOG.md section found for ${version}`); } -const tagRemote = execSync(`git ls-remote origin refs/tags/${tag}`, { encoding: 'utf8' }); +const tagRemote = execFileSync('git', ['ls-remote', 'origin', `refs/tags/${tag}`], { + encoding: 'utf8', +}); if (tagRemote.includes(`refs/tags/${tag}`)) { console.log(`Tag ${tag} already exists — skipping.`); } else { - execSync(`git tag ${tag} && git push origin ${tag}`, { stdio: 'inherit' }); + execFileSync('git', ['tag', tag], { stdio: 'inherit' }); + execFileSync('git', ['push', 'origin', tag], { stdio: 'inherit' }); console.log(`Tagged ${tag}`); } let releaseExists = false; try { - execSync(`gh release view ${tag}`, { stdio: 'ignore' }); + execFileSync('gh', ['release', 'view', tag], { stdio: 'ignore' }); releaseExists = true; } catch { releaseExists = false; @@ -80,7 +90,7 @@ if (releaseExists) { } else { const notesFile = join(mkdtempSync(join(tmpdir(), 'release-notes-')), 'notes.md'); writeFileSync(notesFile, `${section.trimEnd()}\n`); - execSync(`gh release create ${tag} --title ${tag} --notes-file ${notesFile}`, { + execFileSync('gh', ['release', 'create', tag, '--title', tag, '--notes-file', notesFile], { stdio: 'inherit', }); console.log(`Created GitHub release ${tag}`); diff --git a/src/choice/choice.ts b/src/choice/choice.ts index ddacece..3c0376b 100644 --- a/src/choice/choice.ts +++ b/src/choice/choice.ts @@ -12,7 +12,7 @@ import { system1Choice } from './system1-choice.js'; * inference (a single forward pass generating a single token), while System 2 is slower * and more costly: the LLM may reason and has to generate a full structured response. * - * @param question - The decision to make: options, state, instructions and provider settings. + * @param question - The decision to make: options, state, instructions and the model to query. * @returns The winning option, the probability distribution over every option (sums to 1), * and a 0..1 confidence (flat distribution → low, single peak → high). * @throws If there are fewer than 2 options, or if a criteria value is not a string. @@ -22,9 +22,11 @@ import { system1Choice } from './system1-choice.js'; * @example * ```ts * const answer = await choice({ - * apiBaseUrl: 'http://localhost:8000/v1', - * apiKey: process.env.API_KEY!, - * model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + * model: { + * apiBaseUrl: 'http://localhost:8000/v1', + * apiKey: process.env.API_KEY!, + * model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + * }, * state: "It is raining and I am at home. I'm bored.", * instructions: 'Give me a good plan to do now', * criteria: { diff --git a/src/choice/system1-choice.ts b/src/choice/system1-choice.ts index d9ce03e..79f4acb 100644 --- a/src/choice/system1-choice.ts +++ b/src/choice/system1-choice.ts @@ -1,4 +1,5 @@ import { generateText } from '../utils/llms/generate-text.js'; +import { applyExtraBody } from '../utils/llms/apply-extra-body.js'; import { normalizeEntropy } from '../utils/math/normalize-entropy.js'; import type { ChoiceAnswer } from '../types/choice-answer.js'; import type { Question } from '../types/question.js'; @@ -17,7 +18,7 @@ const LETTERS = 'ABCDEFGHIJKLMNOPQRSTUVWXYZ'.split(''); * candidate letter tokens. No deliberation, no chain-of-thought, no structured * output — that's System 2's job. * - * @param question - The decision to make: options, state, instructions and provider settings. + * @param question - The decision to make: options, state, instructions and the model to query. * @returns The winning option (highest letter probability), the probability * distribution over every option (sums to 1), and a 0..1 confidence * based on the distribution's entropy (flat → low, single peak → high). @@ -25,9 +26,11 @@ const LETTERS = 'ABCDEFGHIJKLMNOPQRSTUVWXYZ'.split(''); * @example * ```ts * const answer = await system1Choice({ - * apiBaseUrl: 'http://localhost:8000/v1', - * apiKey: process.env.API_KEY!, - * model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + * model: { + * apiBaseUrl: 'http://localhost:8000/v1', + * apiKey: process.env.API_KEY!, + * model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + * }, * state: "It is raining and I am at home. I'm bored.", * instructions: 'Give me a good plan to do now', * criteria: { walk: 'Go for a walk', movie: 'Watch a movie' }, @@ -62,16 +65,24 @@ export async function system1Choice(question: Question): Promise { // nudges the model towards emitting just the chosen option's letter. // TODO: probably better to move instructions to system prompt for KV cache reuse const res = await generateText( - { apiBaseUrl: question.apiBaseUrl, apiKey: question.apiKey }, - { - model: question.model, - messages: [{ role: 'user', content: prompt }], - max_tokens: 1, // the answer is a single letter - temperature: 0, // greedy: always the most likely letter - logprobs: true, - top_logprobs: 50, // llama.cpp server max; margin so every declared letter (and its token variants) lands in the report - chat_template_kwargs: { enable_thinking: false }, - }, + { apiBaseUrl: question.model.apiBaseUrl, apiKey: question.model.apiKey }, + applyExtraBody( + { + model: question.model.model, + messages: [{ role: 'user', content: prompt }], + max_tokens: 1, // the answer is a single letter + temperature: 0, // greedy: always the most likely letter + logprobs: true, + // 20 is the highest portable window: OpenAI and OpenRouter cap top_logprobs + // at 20, and vLLM's server default --max-logprobs is also 20. llama.cpp + // accepts up to 50, so 20 is safe everywhere logprobs exist. The margin is + // enough for realistic option counts; letters falling outside the window + // just contribute 0 to their option's probability. + top_logprobs: 20, + chat_template_kwargs: { enable_thinking: false }, + }, + question.model.extraBody, + ), { maxRetries: question.maxRetries, timeoutMs: question.timeoutMs }, ); @@ -90,6 +101,12 @@ export async function system1Choice(question: Question): Promise { // TODO: check if keeping the highest one is the best option const letterProbability = new Map(); for (const t of tops) { + // A non-conforming server can send malformed entries (missing or null token); + // skip them instead of crashing mid-read — the remaining candidates still + // carry the decision. + if (typeof t?.token !== 'string') { + continue; + } const letter = t.token.trim().toUpperCase(); const p = Math.exp(t.logprob); if (p > (letterProbability.get(letter) ?? 0)) { diff --git a/src/index.ts b/src/index.ts index db38d9e..9c8b105 100644 --- a/src/index.ts +++ b/src/index.ts @@ -1,3 +1,4 @@ export type { ChoiceMode } from './types/choice-mode.js'; +export type { Model } from './types/model.js'; export type { Question } from './types/question.js'; export { choice } from './choice/choice.js'; diff --git a/src/types/chat-completion-request.ts b/src/types/chat-completion-request.ts index 1449c13..d3ca54d 100644 --- a/src/types/chat-completion-request.ts +++ b/src/types/chat-completion-request.ts @@ -34,4 +34,10 @@ export interface ChatCompletionRequest { top_logprobs?: number; /** Extra parameters forwarded to the model's chat template, i.e. `{ enable_thinking: false }` for llama.cpp */ chat_template_kwargs?: Record; + /** + * Engine- or model-specific fields beyond the standard keys above, i.e. merged in + * from `Model['extraBody']`. They ride on the body verbatim: OpenAI-compatible + * engines ignore unknown fields, so only keys the backend understands take effect. + */ + [key: string]: unknown; } diff --git a/src/types/model.ts b/src/types/model.ts new file mode 100644 index 0000000..b51f06c --- /dev/null +++ b/src/types/model.ts @@ -0,0 +1,42 @@ +/** + * The model to query and the OpenAI-compatible v1 API serving it. + * + * Carries everything about *how the model is reached and served*: endpoint, + * credentials, model identifier and any engine- or model-specific request settings + * the library does not model. One `Model` describing, i.e., a deployed llama.cpp + * server can back any number of `Question`s. + * + * @example + * ```ts + * const model: Model = { + * apiBaseUrl: 'http://localhost:8000/v1', + * apiKey: 'a-super-secret-api-key', + * model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + * }; + * ``` + */ +export interface Model { + /** Base url of an OpenAI compatible v1 API, i.e. `'http://localhost:8000/v1'` */ + apiBaseUrl: string; + /** The API key as required or not by your provider */ + apiKey: string; + /** The model identifier as required by your API provider. i.e. `'/models/Qwen3.5-4B-Q4_K_M.gguf'` for llama.cpp */ + model: string; + /** + * Extra fields forwarded verbatim into the chat completions request body, for + * engine- or model-specific request settings this library does not model. + * OpenAI-compatible engines ignore unknown body fields, so only the keys your + * backend understands take effect. i.e.: + * + * - `{ chat_template_kwargs: { enable_thinking: false } }` — reasoning/thinking + * toggle on llama.cpp, vLLM and SGLang (the exact key can differ per model: + * `enable_thinking` for Qwen3/GLM, `thinking` for DeepSeek-V3.1/Granite) + * - `{ reasoning_effort: 'none' }` — OpenAI, OpenRouter and Ollama's + * `/v1/chat/completions` + * - `{ think: false }` — Ollama native API + * + * Note: servers with strict body validation (the hosted OpenAI family) reject + * genuinely unknown fields with HTTP 400. + */ + extraBody?: Record; +} diff --git a/src/types/question.ts b/src/types/question.ts index 5c2ae74..e2acc0f 100644 --- a/src/types/question.ts +++ b/src/types/question.ts @@ -1,14 +1,17 @@ import type { ChoiceMode } from './choice-mode.js'; +import type { Model } from './model.js'; /** - * The decision to make: options, state, instructions and provider settings. + * The decision to make: options, state, instructions, mode and the model to query. * * @example * ```ts * const question: Question = { - * apiBaseUrl: 'http://localhost:8000/v1', - * apiKey: 'a-super-secret-api-key', - * model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + * model: { + * apiBaseUrl: 'http://localhost:8000/v1', + * apiKey: 'a-super-secret-api-key', + * model: '/models/Qwen3.5-4B-Q4_K_M.gguf', + * }, * state: "It is raining and I am at home. I'm bored.", * instructions: 'Give me a good plan to do now', * criteria: { walk: 'Go for a walk', movie: 'Watch a movie' }, @@ -16,12 +19,8 @@ import type { ChoiceMode } from './choice-mode.js'; * ``` */ export interface Question { - /** Base url of an OpenAI compatible v1 API, i.e. `'http://localhost:8000/v1'` */ - apiBaseUrl: string; - /** The API key as required or not by your provider */ - apiKey: string; - /** The model identifier as required by your API provider. i.e. `'/models/Qwen3.5-4B-Q4_K_M.gguf'` for llama.cpp */ - model: string; + /** The model to use and the OpenAI-compatible API serving it (endpoint, credentials, model id and extra request settings) */ + model: Model; /** Maximum number of retries after a failed attempt (network errors, 408/409/429/5xx). Defaults to 2 */ maxRetries?: number; /** Per-attempt timeout in milliseconds. Defaults to 600000 */ diff --git a/src/utils/llms/apply-extra-body.ts b/src/utils/llms/apply-extra-body.ts new file mode 100644 index 0000000..abdc6af --- /dev/null +++ b/src/utils/llms/apply-extra-body.ts @@ -0,0 +1,102 @@ +import type { ChatCompletionRequest } from '../../types/chat-completion-request.js'; + +/** + * Keys of the chat completions request body that `applyExtraBody` never lets an + * `extraBody` override. + * + * They are exactly the request fields System 1's single-token logprobs trick + * depends on for correctness — `model` and `messages` define what is asked, + * `stream` keeps the response parsable, and `logprobs`, `top_logprobs`, + * `max_tokens` and `temperature` make the first generated token the answer + * distribution itself. Letting a user override any of them through `extraBody` + * could silently corrupt the probability readout. + */ +const RESERVED_KEYS: ReadonlySet = new Set([ + 'model', + 'messages', + 'stream', + 'logprobs', + 'top_logprobs', + 'max_tokens', + 'temperature', +]); + +/** The request key that carries chat-template arguments on OpenAI-compatible engines. */ +const CHAT_TEMPLATE_KWARGS = 'chat_template_kwargs'; + +const isPlainObject = (value: unknown): value is Record => + typeof value === 'object' && value !== null && !Array.isArray(value); + +/** + * Applies a `Model`'s `extraBody` to a chat completions request — the mechanism that + * lets users reach engine- or model-specific settings this library does not model. + * + * Rules: + * - Keys the library relies on for its own correctness (see `RESERVED_KEYS`) are + * ignored. + * - `__proto__` and `constructor` are ignored: `extraBody` is an arbitrary-key + * passthrough, and an own enumerable `__proto__` key (i.e. smuggled through + * `JSON.parse`) must never be able to replace the request object's prototype. + * - `chat_template_kwargs` (the llama.cpp / vLLM / SGLang thinking toggle) is merged + * one level deep, so engine defaults the library sets survive, user keys win per + * key, and sibling model-specific keys coexist (i.e. Qwen3's `enable_thinking` + * alongside DeepSeek's `thinking`). A non-object value for it is ignored, keeping + * the engine defaults intact. + * - Every other key is forwarded verbatim; OpenAI-compatible engines ignore unknown + * body fields, so only the keys the backend understands take effect. + * + * @param request - Chat completions request body in the API's own wire format. + * @param extraBody - Extra fields to forward, typically `Model['extraBody']`. + * @returns A new request body with the extras applied; the input is never mutated. + * @example + * ```ts + * const request = applyExtraBody( + * { model: 'm', messages: [], chat_template_kwargs: { enable_thinking: false } }, + * { chat_template_kwargs: { thinking: false }, think: false }, + * ); + * // → { model: 'm', messages: [], chat_template_kwargs: { enable_thinking: false, thinking: false }, think: false } + * ``` + */ +export function applyExtraBody( + request: ChatCompletionRequest, + extraBody: Record | undefined, +): ChatCompletionRequest { + // No extras: return the request untouched, so the wire body stays byte-identical + // to the pre-extraBody behavior. + if (extraBody === undefined) { + return request; + } + + const merged: ChatCompletionRequest = { ...request }; + + for (const [key, value] of Object.entries(extraBody)) { + // Reserved keys are the library's own (see RESERVED_KEYS); skip silently. + if (RESERVED_KEYS.has(key)) { + continue; + } + + // A computed __proto__ key is an own enumerable property Object.entries happily + // yields; assigning it with [[Set]] would replace the request object's + // prototype. Prototype keys never belong in a request body. + if (key === '__proto__' || key === 'constructor') { + continue; + } + + // Chat template kwargs are special in both directions: a plain object merges one + // level deep (user keys win per key, engine defaults survive, sibling model keys + // coexist), and any non-object value is ignored so the defaults stay intact. + if (key === CHAT_TEMPLATE_KWARGS) { + if (isPlainObject(value)) { + merged[CHAT_TEMPLATE_KWARGS] = { + ...(merged[CHAT_TEMPLATE_KWARGS] as Record | undefined), + ...value, + }; + } + continue; + } + + merged[key] = value; + } + + return merged; +} diff --git a/src/utils/llms/generate-text.ts b/src/utils/llms/generate-text.ts index 4654f20..35ace76 100644 --- a/src/utils/llms/generate-text.ts +++ b/src/utils/llms/generate-text.ts @@ -1,6 +1,7 @@ import { messageOf } from '../error/message-of.js'; import { truncate } from '../text/truncate.js'; import { sleep } from '../time/sleep.js'; +import { readBodyCapped, isBodyTooLargeError } from '../network/read-body-capped.js'; import { backoffDelay } from '../network/backoff-delay.js'; import { isTimeoutError } from '../network/is-timeout-error.js'; import { retryDelayFromHeaders } from '../network/retry-delay-from-headers.js'; @@ -18,6 +19,12 @@ const DEFAULT_MAX_RETRIES = 2; const DEFAULT_TIMEOUT_MS = 600_000; /** How many characters of an error body end up in thrown error messages. */ const MAX_ERROR_BODY_CHARS = 500; +/** + * Safety cap on the buffered response body, in bytes. A one-token logprobs answer + * is a few kilobytes at most, so this only ever fires on a misbehaving server — + * and prevents it from exhausting the host process's memory. + */ +const MAX_RESPONSE_BYTES = 10 * 1024 * 1024; /** * Generates a chat completion from an OpenAI-compatible v1 API and returns its parsed @@ -27,15 +34,18 @@ const MAX_ERROR_BODY_CHARS = 500; * The request is always a single non-streaming POST to `{apiBaseUrl}/chat/completions` * with Bearer authentication. Failed attempts — network errors, plus 408/409/429/5xx * responses or any `x-should-retry` answer — back off exponentially and are retried; - * timeouts are not. + * timeouts and oversized response bodies are not. Redirects are refused rather than + * followed. * - * @param connection - API endpoint and credentials, mirroring `Question`'s provider settings. - * @param request - Chat completions request body in the API's own wire format. + * @param connection - API endpoint and credentials, mirroring `Model`'s transport fields. + * @param request - Chat completions request body in the API's own wire format. Extra + * engine- or model-specific fields (i.e. from `Model['extraBody']`) may be + * present alongside the standard keys; they are serialized into the body as is. * @param options - Retry and timeout settings; defaults to 2 retries with a 10-minute - * timeout per attempt. + * timeout per attempt. Non-finite `maxRetries` falls back to the default. * @returns The parsed response body. Response fields we don't consume are passed through untouched. - * @throws If the API answers with a failing status, an unparsable body, or keeps failing - * after every attempt. + * @throws If `apiBaseUrl` is not a valid URL, the API answers with a failing status, + * an unparsable or over-cap body, or keeps failing after every attempt. * @example * ```ts * const completion = await generateText( @@ -58,17 +68,28 @@ export async function generateText( request: ChatCompletionRequest, options: GenerateTextOptions = {}, ): Promise { - // Not validated on purpose: 0 retries means exactly one attempt, a tiny timeout aborts - // immediately, and negative values degrade monotonically in the same direction. - const maxRetries = options.maxRetries ?? DEFAULT_MAX_RETRIES; + // Sanitized, not just defaulted: a NaN from JS callers would make every + // comparison against the attempt counter false and retry forever; negatives + // degrade to 0 (one attempt) and fractions to whole attempts. + const maxRetries = + typeof options.maxRetries === 'number' && Number.isFinite(options.maxRetries) + ? Math.max(0, Math.floor(options.maxRetries)) + : DEFAULT_MAX_RETRIES; const timeoutMs = options.timeoutMs ?? DEFAULT_TIMEOUT_MS; - // A trailing slash on the base URL must not double when appending the path. - const url = new URL( - connection.apiBaseUrl.endsWith('/') - ? connection.apiBaseUrl + PATH.slice(1) - : connection.apiBaseUrl + PATH, - ); + // The endpoint path is joined through the URL API so a query string on the base + // URL (i.e. `https://host/v1?x=1`) survives the join instead of swallowing the + // appended path. + let url: URL; + try { + url = new URL(connection.apiBaseUrl); + } catch (err) { + throw new Error(`Invalid model.apiBaseUrl "${connection.apiBaseUrl}": ${messageOf(err)}`); + } + url.pathname = `${url.pathname.replace(/\/+$/, '')}${PATH}`; + // Errors are user-facing: some providers take credentials as query parameters, + // and those must never leak into exception messages or logs. + const displayUrl = `${url.origin}${url.pathname}`; // One full request per attempt, so maxRetries counts *retries* // (total attempts = maxRetries + 1). @@ -86,16 +107,26 @@ export async function generateText( }, body: JSON.stringify({ ...request, stream: false }), signal: AbortSignal.timeout(timeoutMs), + // An API should never answer this call with a redirect; refusing to + // follow one keeps the Bearer token off unexpected paths and avoids a + // pointless round trip. + redirect: 'error', }); // Read before the status check: error responses carry their explanation in the - // body, and it can only be consumed once. - text = await response.text(); + // body, and it can only be consumed once. Capped: an oversized body is refused + // instead of buffered (see MAX_RESPONSE_BYTES). + text = await readBodyCapped(response, MAX_RESPONSE_BYTES); } catch (err) { + // An over-cap body is deterministic: the same request cannot shrink on a + // retry, so fail without burning the retry budget. + if (isBodyTooLargeError(err)) { + throw err; + } // Timeouts are not retried: the full per-attempt budget was already spent, so // another wait just multiplies it. Any other fetch failure (connection refused, // reset, ...) is worth another attempt. if (isTimeoutError(err) || attempt >= maxRetries) { - throw new Error(`LLM request to ${url} failed: ${messageOf(err)}`, { cause: err }); + throw new Error(`LLM request to ${displayUrl} failed: ${messageOf(err)}`, { cause: err }); } await sleep(backoffDelay(attempt)); continue; diff --git a/src/utils/network/read-body-capped.ts b/src/utils/network/read-body-capped.ts new file mode 100644 index 0000000..2fb1b17 --- /dev/null +++ b/src/utils/network/read-body-capped.ts @@ -0,0 +1,97 @@ +/** + * Reads an HTTP response body into a string, refusing to buffer more than a + * safety cap of bytes. + * + * The naive `response.text()` buffers the whole body in memory before the caller + * ever looks at the status code — so a misbehaving server (broken, compromised, + * or reached through a bad base URL) can exhaust the host process with a single + * multi-gigabyte answer. This reads the body in chunks instead and aborts the + * moment the cap is exceeded, in either of two ways: + * + * - A declared `Content-Length` above the cap fails immediately, without reading + * a single byte of the body. + * - A streamed (or undeclared) body fails as soon as the accumulated bytes cross + * the cap, tearing the connection down instead of buffering the rest. + * + * The thrown error is tagged (see `isBodyTooLargeError`) so callers can tell this + * deterministic failure apart from transient network errors and skip retrying it. + * + * @param response - The (not yet consumed) HTTP response to read. + * @param maxBytes - Safety cap, in bytes, on the buffered body. + * @returns The full body text, decoded as UTF-8. + * @throws A `BodyTooLargeError`-tagged error when the body is over the cap — + * recognized by `isBodyTooLargeError`. + * @example + * ```ts + * const text = await readBodyCapped(response, 10 * 1024 * 1024); + * try { + * await readBodyCapped(hugeResponse, 1024); + * } catch (err) { + * isBodyTooLargeError(err); // true + * } + * ``` + */ +export async function readBodyCapped(response: Response, maxBytes: number): Promise { + const tooLarge = (): Error => { + const mb = maxBytes / (1024 * 1024); + // The default cap reads as "10 MB"; test caps below a megabyte as exact bytes. + const cap = mb >= 1 ? `${Math.floor(mb)} MB` : `${maxBytes} bytes`; + const err = new Error(`response body exceeds the ${cap} safety cap — refusing to read it`); + err.name = BODY_TOO_LARGE; + return err; + }; + + // A declared length over the cap is settled before the first byte is read. + const declared = Number.parseInt(response.headers.get('content-length') ?? '', 10); + if (Number.isInteger(declared) && declared > maxBytes) { + throw tooLarge(); + } + + // No stream to read from (i.e. `new Response(null)`): nothing to cap, fall back. + const reader = response.body?.getReader(); + if (!reader) { + return response.text(); + } + + const decoder = new TextDecoder(); + let bytes = 0; + let text = ''; + for (;;) { + const { done, value } = await reader.read(); + if (done) break; + bytes += value.byteLength; + if (bytes > maxBytes) { + // Tear the connection down instead of buffering the rest of the body. + // A rejection here surfaces through the caller's error handling like any + // other read failure. + await reader.cancel(); + throw tooLarge(); + } + // stream: true keeps multi-byte characters that span chunk boundaries intact. + text += decoder.decode(value, { stream: true }); + } + return text + decoder.decode(); +} + +/** Error name tagging the deterministic "body over the cap" failure. */ +const BODY_TOO_LARGE = 'BodyTooLargeError'; + +/** + * Whether a thrown value is the tagged error `readBodyCapped` raises for a body + * over its cap. + * + * Unlike timeouts or resets, an oversized body cannot shrink on a retry — the + * same request will answer oversized again — so callers use this to fail fast + * instead of burning their retry budget. + * + * @param err - The thrown value, of any type. + * @returns Whether the value is the tagged too-large error. + * @example + * ```ts + * isBodyTooLargeError(await readBodyCapped(huge, 1024).catch((e) => e)); // true + * isBodyTooLargeError(new DOMException('timed out', 'TimeoutError')); // false + * isBodyTooLargeError('boom'); // false + * ``` + */ +export const isBodyTooLargeError = (err: unknown): boolean => + err instanceof Error && err.name === BODY_TOO_LARGE; diff --git a/test/choice/choice.test.ts b/test/choice/choice.test.ts index 2133002..e246695 100644 --- a/test/choice/choice.test.ts +++ b/test/choice/choice.test.ts @@ -14,9 +14,11 @@ const question = ( mode?: 'system1' | 'system2', criteria: Record = { a: 'A', b: 'B' }, ): Question => ({ - apiBaseUrl: 'https://example.com/v1', - apiKey: 'key', - model: 'model', + model: { + apiBaseUrl: 'https://example.com/v1', + apiKey: 'key', + model: 'model', + }, // Omitted when undefined: exactOptionalPropertyTypes forbids an explicit `mode: undefined`. ...(mode !== undefined && { mode }), criteria, diff --git a/test/choice/system1-choice.test.ts b/test/choice/system1-choice.test.ts index e4f5c9c..37053ba 100644 --- a/test/choice/system1-choice.test.ts +++ b/test/choice/system1-choice.test.ts @@ -19,9 +19,11 @@ afterEach(() => { // for right now (walking or a movie fit a rainy day at home, the beach does not); // options map to letters A..C. const question = () => ({ - apiBaseUrl: 'https://example.com/v1', - apiKey: 'key', - model: 'model', + model: { + apiBaseUrl: 'https://example.com/v1', + apiKey: 'key', + model: 'model', + }, criteria: { walk: 'Go for a walk', movie: 'Watch a movie', @@ -78,7 +80,7 @@ describe('system1Choice', () => { expect(params.max_tokens).toBe(1); // one forward pass, one generated token expect(params.temperature).toBe(0); expect(params.logprobs).toBe(true); - expect(params.top_logprobs).toBe(50); + expect(params.top_logprobs).toBe(20); expect(params.chat_template_kwargs).toEqual({ enable_thinking: false }); expect(params.stream).toBe(false); // Prompt includes the lettered options and the single-letter instruction. @@ -86,6 +88,91 @@ describe('system1Choice', () => { expect(params.messages[0].content).toContain('exactly one letter'); }); + it('ignores malformed logprob entries instead of crashing on them', async () => { + // A non-conforming server can send nulls or token-less entries inside + // top_logprobs; they are skipped, and the well-formed candidates decide. + fetchMock.mockResolvedValueOnce( + new Response( + JSON.stringify({ + choices: [ + { + logprobs: { + content: [ + { + top_logprobs: [ + null, + { token: null, logprob: Math.log(0.9) }, + { token: ' A', logprob: Math.log(0.3) }, + { token: ' B', logprob: Math.log(0.7) }, + ], + }, + ], + }, + }, + ], + }), + { status: 200 }, + ), + ); + const answer = await system1Choice(question()); + expect(answer.choice).toBe('movie'); + expect(answer.probabilities.walk).toBeCloseTo(0.3, 10); + expect(answer.probabilities.movie).toBeCloseTo(0.7, 10); + }); + + it('forwards extraBody extra keys and lets the user override the thinking default', async () => { + fetchMock.mockResolvedValueOnce(response([logprobToken('A', Math.log(0.6))])); + await system1Choice({ + ...question(), + model: { + ...question().model, + extraBody: { + think: false, // Ollama-native key; engines ignore what they don't know + chat_template_kwargs: { enable_thinking: true }, // user override wins per key + }, + }, + }); + + const [, init] = fetchMock.mock.calls[0] as [URL, RequestInit]; + const params = JSON.parse(init.body as string); + // Exotic key passes through verbatim... + expect(params.think).toBe(false); + // ...while the user's chat_template_kwargs override replaced the default, + // and the reserved sampling knobs were not touched. + expect(params.chat_template_kwargs).toEqual({ enable_thinking: true }); + expect(params.max_tokens).toBe(1); + expect(params.temperature).toBe(0); + expect(params.logprobs).toBe(true); + expect(params.top_logprobs).toBe(20); + }); + + it('keeps the default thinking toggle and ignores reserved keys passed through extraBody', async () => { + fetchMock.mockResolvedValueOnce(response([logprobToken('A', Math.log(0.6))])); + await system1Choice({ + ...question(), + model: { + ...question().model, + extraBody: { + model: 'other', + max_tokens: 100, + stream: true, + messages: [], + logprobs: false, + }, + }, + }); + + const [, init] = fetchMock.mock.calls[0] as [URL, RequestInit]; + const params = JSON.parse(init.body as string); + // Reserved keys never leak into the body as overrides; the request stays intact. + expect(params.model).toBe('model'); + expect(params.max_tokens).toBe(1); + expect(params.stream).toBe(false); + expect(params.messages).toHaveLength(1); + expect(params.logprobs).toBe(true); + expect(params.chat_template_kwargs).toEqual({ enable_thinking: false }); + }); + it('forwards retry and timeout settings to the transport', async () => { const timeoutSpy = vi.spyOn(AbortSignal, 'timeout'); try { diff --git a/test/types/chat-completion-request.test.ts b/test/types/chat-completion-request.test.ts index ec4aa74..f60ef9e 100644 --- a/test/types/chat-completion-request.test.ts +++ b/test/types/chat-completion-request.test.ts @@ -18,4 +18,16 @@ describe('chatCompletionRequest', () => { Record | undefined >(); }); + + it('carries an index signature so engine-specific extras type-check on the body', () => { + // Extra fields (i.e. merged in from Model['extraBody']) must be representable + // on the wire type itself, and a plain request must remain assignable to it. + expectTypeOf().toExtend>(); + const withExtras: ChatCompletionRequest = { + model: 'm', + messages: [], + reasoning_effort: 'none', // not a declared key; accepted via the index signature + }; + expectTypeOf(withExtras).not.toBeNever(); + }); }); diff --git a/test/types/model.test.ts b/test/types/model.test.ts new file mode 100644 index 0000000..9ac3060 --- /dev/null +++ b/test/types/model.test.ts @@ -0,0 +1,14 @@ +import { describe, expectTypeOf, it } from 'vitest'; +import type { Model } from '../../src/types/model.js'; + +// src/types/model.ts contains only types, so these are compile-time +// assertions: vitest type-checks them with expectTypeOf and the tests pass +// trivially at runtime. +describe('model', () => { + it('has the expected fields and optionality', () => { + expectTypeOf().toEqualTypeOf(); + expectTypeOf().toEqualTypeOf(); + expectTypeOf().toEqualTypeOf(); + expectTypeOf().toEqualTypeOf | undefined>(); + }); +}); diff --git a/test/types/question.test.ts b/test/types/question.test.ts index da81bf9..8d4c98a 100644 --- a/test/types/question.test.ts +++ b/test/types/question.test.ts @@ -1,5 +1,6 @@ import { describe, expectTypeOf, it } from 'vitest'; import type { ChoiceMode } from '../../src/types/choice-mode.js'; +import type { Model } from '../../src/types/model.js'; import type { Question } from '../../src/types/question.js'; // src/types/question.ts contains only types, so these are compile-time @@ -7,9 +8,7 @@ import type { Question } from '../../src/types/question.js'; // trivially at runtime. describe('question', () => { it('has the expected fields and optionality', () => { - expectTypeOf().toEqualTypeOf(); - expectTypeOf().toEqualTypeOf(); - expectTypeOf().toEqualTypeOf(); + expectTypeOf().toEqualTypeOf(); expectTypeOf().toEqualTypeOf(); expectTypeOf().toEqualTypeOf(); expectTypeOf().toEqualTypeOf(); diff --git a/test/utils/llms/apply-extra-body.test.ts b/test/utils/llms/apply-extra-body.test.ts new file mode 100644 index 0000000..2bdf904 --- /dev/null +++ b/test/utils/llms/apply-extra-body.test.ts @@ -0,0 +1,119 @@ +import { describe, expect, it } from 'vitest'; +import type { ChatCompletionRequest } from '../../../src/types/chat-completion-request.js'; +import { applyExtraBody } from '../../../src/utils/llms/apply-extra-body.js'; + +// Baseline body mirroring what system1Choice builds: the two required fields plus +// every sampling knob the single-token trick depends on. +const baseline = (): ChatCompletionRequest => ({ + model: 'model', + messages: [{ role: 'user', content: 'hi' }], + max_tokens: 1, + temperature: 0, + logprobs: true, + top_logprobs: 20, + chat_template_kwargs: { enable_thinking: false }, +}); + +describe('applyExtraBody', () => { + it('returns the request untouched when extraBody is undefined', () => { + const request = baseline(); + expect(applyExtraBody(request, undefined)).toBe(request); + expect(applyExtraBody(request, undefined)).toEqual(baseline()); + }); + + it('adds unknown keys verbatim', () => { + expect(applyExtraBody(baseline(), { reasoning_effort: 'none', think: false })).toEqual({ + ...baseline(), + reasoning_effort: 'none', + think: false, + }); + }); + + it('ignores every reserved key', () => { + // Each reserved key tries to override the baseline with a poisoned value; the + // baseline must come out untouched for all of them. + const sabotage = { + model: 'other', + messages: [], + stream: true, + logprobs: false, + top_logprobs: 1, + max_tokens: 100, + temperature: 2, + } as Record; + expect(applyExtraBody(baseline(), sabotage)).toEqual(baseline()); + }); + + it('merges chat_template_kwargs one level deep, keeping the default and adding sibling keys', () => { + expect( + applyExtraBody(baseline(), { + chat_template_kwargs: { thinking: false }, // DeepSeek/Granite key, alongside Qwen3's + }), + ).toEqual({ + ...baseline(), + chat_template_kwargs: { enable_thinking: false, thinking: false }, + }); + }); + + it('lets user keys win per key inside chat_template_kwargs', () => { + expect( + applyExtraBody(baseline(), { + chat_template_kwargs: { enable_thinking: true }, // user flips the default + }), + ).toEqual({ + ...baseline(), + chat_template_kwargs: { enable_thinking: true }, + }); + }); + + it('ignores non-object chat_template_kwargs, keeping the engine defaults', () => { + // null, a string and an array are all non-plain-object values for the merge; the + // baseline default must survive each of them. + expect(applyExtraBody(baseline(), { chat_template_kwargs: null })).toEqual(baseline()); + expect(applyExtraBody(baseline(), { chat_template_kwargs: 'enable_thinking=false' })).toEqual( + baseline(), + ); + expect(applyExtraBody(baseline(), { chat_template_kwargs: ['enable_thinking'] })).toEqual( + baseline(), + ); + }); + + it('adds chat_template_kwargs when the baseline request has none', () => { + const request: ChatCompletionRequest = { model: 'm', messages: [] }; + expect(applyExtraBody(request, { chat_template_kwargs: { thinking: false } })).toEqual({ + model: 'm', + messages: [], + chat_template_kwargs: { thinking: false }, + }); + }); + + it('adds no chat_template_kwargs key when neither side has one', () => { + const request: ChatCompletionRequest = { model: 'm', messages: [] }; + expect(applyExtraBody(request, { seed: 42 })).toEqual({ model: 'm', messages: [], seed: 42 }); + }); + + it('ignores prototype-smuggling keys so the request object cannot be reparented', () => { + // JSON.parse is the classic way an own enumerable __proto__ property appears + // in a config object; assigning it with [[Set]] would replace the request's + // prototype instead of adding a harmless body key. + const poisoned: Record = JSON.parse( + '{"__proto__": {"polluted": true}, "constructor": 42, "seed": 7}', + ); + const merged = applyExtraBody(baseline(), poisoned); + expect(Object.getPrototypeOf(merged)).toBe(Object.prototype); + expect(Object.hasOwn(merged, '__proto__')).toBe(false); + expect(merged).toEqual({ ...baseline(), seed: 7 }); // normal keys still flow + }); + + it('keeps extra top-level keys already present on the request and returns a new object, never mutating the input', () => { + const request = { ...baseline(), reasoning_effort: 'low' }; + const merged = applyExtraBody(request, { seed: 7 }); + expect(merged).not.toBe(request); + expect(merged).toEqual({ ...baseline(), reasoning_effort: 'low', seed: 7 }); + // Input untouched: still the value it was constructed with (baseline + the + // reasoning_effort key the caller set). + expect(request).toEqual({ ...baseline(), reasoning_effort: 'low' }); + // The reasoning_effort key the caller already set survives the merge. + expect(merged).toHaveProperty('reasoning_effort', 'low'); + }); +}); diff --git a/test/utils/llms/generate-text.test.ts b/test/utils/llms/generate-text.test.ts index c78f3e7..8512566 100644 --- a/test/utils/llms/generate-text.test.ts +++ b/test/utils/llms/generate-text.test.ts @@ -22,6 +22,7 @@ interface FetchInit { headers: Record; body: string; signal: AbortSignal; + redirect: string; } beforeEach(() => { @@ -50,6 +51,7 @@ describe('generateText', () => { expect(init.headers.Accept).toBe('application/json'); expect(init.headers.Authorization).toBe('Bearer secret'); expect(init.headers['User-Agent']).toMatch(/^smart-decisions\//); + expect(init.redirect).toBe('error'); // never follow redirects expect(JSON.parse(init.body)).toEqual({ model: 'model', messages: [{ role: 'user', content: 'hi' }], @@ -69,6 +71,106 @@ describe('generateText', () => { expect(String(fetchMock.mock.calls[1]![0])).toBe('https://example.com/v1/chat/completions'); }); + it('joins the endpoint path through the URL API, preserving base URL queries', async () => { + fetchMock.mockResolvedValueOnce(ok({ choices: [] })); + await generateText({ apiBaseUrl: 'https://example.com/v1?team=1', apiKey: 'k' }, request()); + + // String concatenation would put the path inside the query; the URL join + // keeps the query and appends the path where it belongs. + expect(String(fetchMock.mock.calls[0]![0])).toBe( + 'https://example.com/v1/chat/completions?team=1', + ); + }); + + it('rejects an invalid base URL without attempting a request', async () => { + await expect(generateText({ apiBaseUrl: 'not a url', apiKey: 'k' }, request())).rejects.toThrow( + 'Invalid model.apiBaseUrl "not a url"', + ); + expect(fetchMock).not.toHaveBeenCalled(); + }); + + it('strips query strings from URLs in error messages so query-borne keys cannot leak', async () => { + fetchMock.mockRejectedValue(new TypeError('fetch failed')); + + const err = (await generateText( + { apiBaseUrl: 'https://example.com/v1?key=secret', apiKey: 'k' }, + request(), + { maxRetries: 0 }, + ).catch((e: unknown) => e)) as Error; + expect(err.message).toBe( + 'LLM request to https://example.com/v1/chat/completions failed: fetch failed', + ); + expect(err.message).not.toContain('key=secret'); + }); + + it('refuses response bodies over the safety cap without retrying', async () => { + // Declared over cap: fails before a single body byte is read. + fetchMock.mockResolvedValueOnce( + new Response('tiny', { + status: 200, + headers: { 'content-length': String(10 * 1024 * 1024 + 1) }, + }), + ); + await expect(generateText(connection(), request())).rejects.toThrow( + 'response body exceeds the 10 MB safety cap — refusing to read it', + ); + expect(fetchMock).toHaveBeenCalledTimes(1); + + // Undeclared (streamed) over cap: fails mid-read, also without a retry. + const chunks = ['x'.repeat(5 * 1024 * 1024), 'x'.repeat(5 * 1024 * 1024), 'x'.repeat(1024)]; + let i = 0; + const stream = new ReadableStream({ + pull(controller) { + if (i < chunks.length) controller.enqueue(new TextEncoder().encode(chunks[i++]!)); + else controller.close(); + }, + }); + fetchMock.mockResolvedValueOnce(new Response(stream, { status: 200 })); + await expect(generateText(connection(), request())).rejects.toThrow( + 'response body exceeds the 10 MB safety cap — refusing to read it', + ); + expect(fetchMock).toHaveBeenCalledTimes(2); + }); + + it('sanitizes non-finite maxRetries to the default, never retrying forever', async () => { + // NaN would make every `attempt >= maxRetries` comparison false → infinite + // retries against a failing server. Infinity falls back to the default too. + for (const maxRetries of [Number.NaN, Number.POSITIVE_INFINITY]) { + // A factory, not mockResolvedValue: a Response body can only be read once, + // so every attempt needs a fresh Response. + fetchMock.mockImplementation(async () => + fail(429, 'rate limited', { 'retry-after-ms': '1' }), + ); + const pending = expect(generateText(connection(), request(), { maxRetries })).rejects.toThrow( + 'LLM API returned HTTP 429: rate limited', + ); + await vi.advanceTimersByTimeAsync(2); // two 1ms header-requested waits + await pending; + // Default maxRetries (2) → exactly 3 attempts, then give up. + expect(fetchMock).toHaveBeenCalledTimes(3); + fetchMock.mockReset(); + } + }); + + it('degrades negative and fractional maxRetries to whole, non-negative attempt counts', async () => { + // -5 → 0 retries → a single attempt, no wait involved. + fetchMock.mockResolvedValue(fail(429, 'rate limited', { 'retry-after-ms': '1' })); + await expect(generateText(connection(), request(), { maxRetries: -5 })).rejects.toThrow( + 'LLM API returned HTTP 429: rate limited', + ); + expect(fetchMock).toHaveBeenCalledTimes(1); + + // 1.5 → floor to 1 retry → two attempts. + fetchMock.mockReset(); + fetchMock.mockImplementation(async () => fail(429, 'rate limited', { 'retry-after-ms': '1' })); + const pending = expect( + generateText(connection(), request(), { maxRetries: 1.5 }), + ).rejects.toThrow('LLM API returned HTTP 429: rate limited'); + await vi.advanceTimersByTimeAsync(1); + await pending; + expect(fetchMock).toHaveBeenCalledTimes(2); + }); + it('applies the default timeout and honors an explicit one, per attempt', async () => { fetchMock.mockResolvedValueOnce(ok({ choices: [] })); await generateText(connection(), request()); @@ -79,6 +181,27 @@ describe('generateText', () => { expect(timeoutSpy).toHaveBeenCalledWith(42); }); + it('adds extra top-level request keys into the body untouched', async () => { + fetchMock.mockResolvedValueOnce(ok({ choices: [] })); + + // Engine/model-specific fields (i.e. from Model['extraBody']) ride on the + // wire type's index signature; the transport must serialize them as is. + await generateText(connection(), { + ...request(), + reasoning_effort: 'none', + chat_template_kwargs: { enable_thinking: false }, + }); + + const [, init] = fetchMock.mock.calls[0] as [URL, FetchInit]; + expect(JSON.parse(init.body)).toEqual({ + model: 'model', + messages: [{ role: 'user', content: 'hi' }], + reasoning_effort: 'none', + chat_template_kwargs: { enable_thinking: false }, + stream: false, + }); + }); + it('throws immediately on a non-retryable status, without retrying', async () => { fetchMock.mockResolvedValueOnce(fail(400, '{"error":{"message":"bad"}}')); diff --git a/test/utils/network/read-body-capped.test.ts b/test/utils/network/read-body-capped.test.ts new file mode 100644 index 0000000..ca58efd --- /dev/null +++ b/test/utils/network/read-body-capped.test.ts @@ -0,0 +1,87 @@ +import { describe, expect, it } from 'vitest'; +import { + isBodyTooLargeError, + readBodyCapped, +} from '../../../src/utils/network/read-body-capped.js'; + +// Builds a body stream with no Content-Length, handing out the given chunks +// one read at a time — the shape an undeclared/chunked server response has. +const streamBody = (...chunks: string[]): ReadableStream => { + const encoded = chunks.map((c) => new TextEncoder().encode(c)); + return new ReadableStream({ + pull(controller) { + const chunk = encoded.shift(); + if (chunk === undefined) { + controller.close(); + return; + } + controller.enqueue(chunk); + }, + }); +}; + +describe('readBodyCapped', () => { + it('returns the text of a body within the cap', async () => { + const response = new Response('hello logprobs', { headers: { 'content-length': '14' } }); + await expect(readBodyCapped(response, 1024)).resolves.toBe('hello logprobs'); + }); + + it('rejects a declared Content-Length over the cap without reading the body', async () => { + // If the body were ever read, its pull would throw and the rejection would be + // THAT error instead of the cap error — so the assertion below doubles as the + // proof that the cap is settled from the header alone. (The stream's eager + // controller pull erroring on its own is unobserved: nobody reads it.) + const body = new ReadableStream({ + pull() { + throw new Error('the body must never be read'); + }, + }); + const response = new Response(body, { headers: { 'content-length': '3000000' } }); + + await expect(readBodyCapped(response, 2 * 1024 * 1024)).rejects.toThrow( + 'response body exceeds the 2 MB safety cap — refusing to read it', + ); + }); + + it('rejects a streamed body the moment it crosses the cap', async () => { + // No Content-Length: the cap fires mid-stream, on the chunk that overflows. + const response = new Response(streamBody('aaaa', 'bbbb', 'cccc')); + await expect(readBodyCapped(response, 6)).rejects.toThrow( + 'response body exceeds the 6 bytes safety cap — refusing to read it', + ); + }); + + it('assembles a streamed body that stays under the cap', async () => { + const response = new Response(streamBody('aaaa', 'bb')); + await expect(readBodyCapped(response, 1024)).resolves.toBe('aaaabb'); + }); + + it('keeps multi-byte characters intact across chunk boundaries', async () => { + // 'é' is two UTF-8 bytes; split them across two chunks and cap above the total. + const bytes = new TextEncoder().encode('é'); + const stream = new ReadableStream({ + start(controller) { + controller.enqueue(bytes.slice(0, 1)); + controller.enqueue(bytes.slice(1)); + controller.close(); + }, + }); + const response = new Response(stream); + await expect(readBodyCapped(response, 1024)).resolves.toBe('é'); + }); + + it('falls back to text() when the response has no stream at all', async () => { + // new Response(null) exposes body === null; nothing to cap, nothing to read. + const response = new Response(null, { status: 200 }); + await expect(readBodyCapped(response, 1024)).resolves.toBe(''); + }); + + it('tags its rejections so callers can avoid retrying them', async () => { + const response = new Response('x'.repeat(17), { headers: { 'content-length': '17' } }); + const err = await readBodyCapped(response, 16).catch((e: unknown) => e); + expect(isBodyTooLargeError(err)).toBe(true); + expect(isBodyTooLargeError(new Error('network reset'))).toBe(false); + expect(isBodyTooLargeError(new DOMException('timed out', 'TimeoutError'))).toBe(false); + expect(isBodyTooLargeError('boom')).toBe(false); + }); +});