diff --git a/README.md b/README.md index e00a21f..7e53a04 100644 --- a/README.md +++ b/README.md @@ -98,6 +98,13 @@ const result = await ctx.text2run(` models({ // Set default model when no @model specified default: "gpt-4.1-mini", + + // Or set the default model and its request defaults together. + // All fields except `model` are added to its outgoing request payload. + default: { + model: "openai/gpt-5.6-terra", + reasoning: { max_tokens: 8000 } + }, // Create aliases for models alias: { "gpt": "gpt-4.1-mini" }, @@ -117,8 +124,23 @@ models({ openai: process.env.OPENAI_KEY, anthropic: process.env.ANTHROPIC_KEY, ollama: process.env.OLLAMA_URL + }, + + // Provider request fields used as defaults for individual models. + // Keys are real model IDs, so aliases and `default` inherit them. + params: { + "openai/gpt-5.6-luna": { + reasoning: { max_tokens: 8000 } + }, + "claude-sonnet-4-20250514": { + thinking: { type: "enabled", budget_tokens: 8000 }, + max_tokens: 12000 + } } }) + +// `params` are deep-merged into the outgoing request. Fields set later in a +// chat (for example by `| prop`) override these configured defaults. ``` ### Provider-Specific diff --git a/src/index.js b/src/index.js index 1e98f84..da79904 100644 --- a/src/index.js +++ b/src/index.js @@ -11,16 +11,47 @@ const tune = require("tune-sdk") const man = require("tune-sdk/man"); man.addPackage(__dirname) +// Merge request defaults without replacing nested provider options. Values in +// `override` win, so inline model processors can override configured defaults. +function mergeParams(defaults = {}, override = {}) { + const result = { ...defaults } + for (const [key, value] of Object.entries(override)) { + result[key] = value && typeof value === "object" && !Array.isArray(value) + ? mergeParams(defaults[key] || {}, value) + : value + } + return result +} + +function withParams(node, defaults) { + if (!node || !defaults || node.type !== "llm") return node + return { + ...node, + exec: (payload, ctx) => node.exec(mergeParams(defaults, payload), ctx) + } +} + +// `default` may be a model name or `{ model, ...requestDefaults }`. +function defaultConfig(value) { + if (typeof value === "string") return { model: value, params: {} } + if (!value || typeof value !== "object") return { model: undefined, params: {} } + const { model, ...params } = value + return { model, params } +} + function createModelsMiddleware(options = {}) { const { cache = true, // disk cache by default for text editors cacheTtl = 3600000, // 1 hour default - default: defaultModel, + default: defaultOption, apiKeys = {}, expose = undefined, - alias = {} + alias = {}, + params = {} } = options; + const { model: defaultModel, params: defaultParams } = defaultConfig(defaultOption) + // Create configured providers const providers = [ ollama({ cache: false, apiKey: apiKeys.ollama }), @@ -32,14 +63,19 @@ function createModelsMiddleware(options = {}) { openrouter({ cache, cacheTtl, apiKey: apiKeys.openrouter }), ]; + async function resolveModel(context, modelName, args, requestDefaults = {}) { + const node = await tune.resolve(context, modelName, args, providers) + return withParams(node, mergeParams(params[modelName], requestDefaults)) + } + return async function models(name, args) { // Handle default model resolution if (name === "default" && args.type === "llm" && defaultModel) { - return tune.resolve(this, defaultModel, args, providers); + return resolveModel(this, defaultModel, args, defaultParams) } - - // Handle aliases + // Resolve aliases before looking up request defaults. This makes aliases + // and @default share the configuration of their real model ID. const resolvedName = alias[name] || name; // TODO: what if name is regex? @@ -47,7 +83,7 @@ function createModelsMiddleware(options = {}) { return } - return tune.resolve(this, resolvedName, args, providers); + return resolveModel(this, resolvedName, args) } } @@ -60,4 +96,7 @@ module.exports.gemini = gemini; module.exports.mistral = mistral; module.exports.groq = groq; module.exports.ollama = ollama; +module.exports.mergeParams = mergeParams; +module.exports.withParams = withParams; +module.exports.defaultConfig = defaultConfig; diff --git a/test/index.js b/test/index.js index 1adfb9d..fafd8e0 100644 --- a/test/index.js +++ b/test/index.js @@ -5,7 +5,7 @@ const fs = require('fs'); const path = require('path'); const tune = require('tune-sdk'); const models = require('../src/index'); -const { openai, anthropic, groq, mistral, gemini, openrouter } = require('../src/index'); +const { openai, anthropic, groq, mistral, gemini, openrouter, mergeParams, withParams, defaultConfig } = require('../src/index'); const llmUtils = require('../src/llm-utils.js') require('dotenv').config() @@ -22,6 +22,47 @@ const env = { const tests = {}; + +tests.params_merge = async function() { + const defaults = { + reasoning: { max_tokens: 8000, enabled: true }, + max_tokens: 12000, + messages: [{ role: "system", content: "default" }] + } + const payload = { + reasoning: { max_tokens: 2000 }, + max_tokens: 16000, + messages: [{ role: "user", content: "call" }] + } + const merged = mergeParams(defaults, payload) + assert.deepEqual(merged.reasoning, { max_tokens: 2000, enabled: true }) + assert.equal(merged.max_tokens, 16000) + assert.deepEqual(merged.messages, payload.messages) + assert.deepEqual(defaults.reasoning, { max_tokens: 8000, enabled: true }) + + let received + const node = withParams({ + type: "llm", + exec: async (request) => { + received = request + return "request created" + } + }, defaults) + assert.equal(await node.exec(payload), "request created") + assert.deepEqual(received, merged) + + assert.deepEqual(defaultConfig("gpt-4.1-mini"), { + model: "gpt-4.1-mini", params: {} + }) + assert.deepEqual(defaultConfig({ + model: "openai/gpt-5.6-terra", + reasoning: { max_tokens: 8000 } + }), { + model: "openai/gpt-5.6-terra", + params: { reasoning: { max_tokens: 8000 } } + }) +} + tests.api_keys = async function(){ assert.ok(process.env.OPENAI_KEY, "OPENAI_KEY has to be set for testing") assert.ok(process.env.ANTHROPIC_KEY, "ANTHROPIC_KEY has to be set for testing")