diff --git a/src/routes/clinicalAssistant.js b/src/routes/clinicalAssistant.js index f72530c..9e45e65 100644 --- a/src/routes/clinicalAssistant.js +++ b/src/routes/clinicalAssistant.js @@ -36,6 +36,7 @@ var { var { buildSystemPrompt, buildUserPrompt, + assistantGenerationOptions, finalizeAssistantAnswer } = require('../utils/clinicalAnswer'); @@ -182,12 +183,17 @@ router.post('/clinical-assistant/chat', async function(req, res) { var prepared = await prepareAssistantChat(req.body); if (prepared.direct) return res.json(prepared.direct); - var ai = await callAI(prepared.messages, { + var ai = await callAI(prepared.messages, assistantGenerationOptions({ model: prepared.chatModel || undefined, temperature: 0.15, maxTokens: 2600 + })); + var finalized = await finalizeAssistantAnswer(ai, { + messages: prepared.messages, + chatModel: prepared.chatModel, + callAI: callAI, + generationOptions: assistantGenerationOptions({ temperature: 0.15 }) }); - var finalized = await finalizeAssistantAnswer(ai, { messages: prepared.messages, chatModel: prepared.chatModel, callAI: callAI }); var answer = finalized.answer; ai = finalized.ai; @@ -237,11 +243,11 @@ router.post('/clinical-assistant/chat/stream', async function(req, res) { sendEvent('sources', { sources: safeSources, search: prepared.search }); sendEvent('status', { message: 'Generating answer...' }); - var ai = await callAIStream(prepared.messages, { + var ai = await callAIStream(prepared.messages, assistantGenerationOptions({ model: prepared.chatModel || undefined, temperature: 0.15, maxTokens: 2600 - }, function(delta) { + }), function(delta) { sendEvent('token', { token: delta }); }); @@ -249,6 +255,7 @@ router.post('/clinical-assistant/chat/stream', async function(req, res) { messages: prepared.messages, chatModel: prepared.chatModel, callAI: callAI, + generationOptions: assistantGenerationOptions({ temperature: 0.15 }), streamed: true, onRegenerating: function() { sendEvent('status', { message: 'Completing answer...' }); } }); @@ -467,11 +474,11 @@ async function rewriteSearchQuery(message, history, chatModel) { role: 'user', content: 'Conversation:\n' + hist + '\n\nLatest user question:\n' + message + '\n\nStandalone search query:' } - ], { + ], assistantGenerationOptions({ model: chatModel || undefined, temperature: 0, maxTokens: 80 - }); + })); var rewritten = String(ai.content || '').replace(/^['"]|['"]$/g, '').replace(/\s+/g, ' ').trim(); if (!rewritten || rewritten.length < 6 || rewritten.length > 300) return message; if (/^(yes|no|maybe|i don'?t know)$/i.test(rewritten)) return message; diff --git a/src/utils/ai.js b/src/utils/ai.js index 7557c89..cafd1de 100644 --- a/src/utils/ai.js +++ b/src/utils/ai.js @@ -7,6 +7,7 @@ const { OpenAI } = require('openai'); const { DEFAULT_MODEL, FALLBACK_MODEL, getBedrockModelId, getBedrockMaxOut } = require('./models'); const logger = require('./logger'); +const { resolveGenerationOptions } = require('./generationOptions'); var activeProvider = process.env.AI_PROVIDER || (process.env.LITELLM_API_BASE ? 'litellm' : 'openrouter'); @@ -367,15 +368,21 @@ async function callVertex(messages, model, temperature, maxTokens) { // ============================================================ // CALL LITELLM (OpenAI-compatible proxy) // ============================================================ -async function callLiteLLM(messages, model, temperature, maxTokens) { +function addReasoningOptions(request, generation) { + if (generation.reasoningEffort != null) request.reasoning_effort = generation.reasoningEffort; + if (generation.reasoningFormat != null) request.reasoning_format = generation.reasoningFormat; + return request; +} + +async function callLiteLLM(messages, model, temperature, maxTokens, generation) { if (!litellmClient) throw new Error('LiteLLM not configured. Set LITELLM_API_BASE in .env'); - var completion = await litellmClient.chat.completions.create({ + var completion = await litellmClient.chat.completions.create(addReasoningOptions({ model: model, messages: messages, temperature: temperature, max_tokens: maxTokens - }); + }, generation || {})); return { success: true, @@ -421,8 +428,9 @@ async function callAIStream(messages, options, onToken) { options = options || {}; var requestedModel = options.model; var model = await resolveModel(requestedModel); - var temperature = options.temperature || 0.3; - var maxTokens = options.maxTokens || 4000; + var generation = resolveGenerationOptions(options); + var temperature = generation.temperature; + var maxTokens = generation.maxTokens; var startTime = Date.now(); await assertModelAllowed(model, options); @@ -443,13 +451,13 @@ async function callAIStream(messages, options, onToken) { var content = ''; var finishReason = null; - var stream = await client.chat.completions.create({ + var stream = await client.chat.completions.create(addReasoningOptions({ model: model, messages: messages, temperature: temperature, max_tokens: maxTokens, stream: true - }); + }, generation)); for await (var part of stream) { var choice = part && part.choices && part.choices[0] ? part.choices[0] : null; if (choice && choice.finish_reason) finishReason = choice.finish_reason; @@ -470,8 +478,9 @@ async function callAI(messages, options) { options = options || {}; var requestedModel = options.model; var model = await resolveModel(requestedModel); - var temperature = options.temperature || 0.3; - var maxTokens = options.maxTokens || 4000; + var generation = resolveGenerationOptions(options); + var temperature = generation.temperature; + var maxTokens = generation.maxTokens; var startTime = Date.now(); // Server-side whitelist: reject any model the operator hasn't enabled. @@ -492,7 +501,7 @@ async function callAI(messages, options) { } else if (activeProvider === 'vertex' && vertexClient) { result = await callVertex(messages, model, temperature, maxTokens); } else if (activeProvider === 'litellm' && litellmClient) { - result = await callLiteLLM(messages, model, temperature, maxTokens); + result = await callLiteLLM(messages, model, temperature, maxTokens, generation); } else if (openrouter) { result = await callOpenRouter(messages, model, temperature, maxTokens); } else { @@ -552,7 +561,7 @@ async function callAI(messages, options) { if (activeProvider === 'litellm' && model !== FALLBACK_MODEL && litellmClient) { logger.warn('Trying fallback model on LiteLLM: ' + FALLBACK_MODEL); try { - var litellmFallback = await callLiteLLM(messages, FALLBACK_MODEL, temperature, maxTokens); + var litellmFallback = await callLiteLLM(messages, FALLBACK_MODEL, temperature, maxTokens, generation); litellmFallback.fallback = true; litellmFallback.duration = Date.now() - startTime; logger.info('LiteLLM fallback success', { model: FALLBACK_MODEL }); diff --git a/src/utils/clinicalAnswer.js b/src/utils/clinicalAnswer.js index 9d72676..fcfa695 100644 --- a/src/utils/clinicalAnswer.js +++ b/src/utils/clinicalAnswer.js @@ -9,17 +9,23 @@ function buildUserPrompt(question, context, history, searchQuery) { return 'Question:\n' + question + searchNote + '\n\nRecent conversation, if relevant:\n' + (hist || 'None') + '\n\nRetrieved sources:\n' + context + '\n\nWrite the answer now. If the question is a short misspelled or partial term and the sources point to a likely concept, answer the likely concept rather than asking for clarification.'; } +function assistantGenerationOptions(overrides) { + return Object.assign({ + reasoningEffort: 'low', + reasoningFormat: 'hidden' + }, overrides || {}); +} + async function finalizeAssistantAnswer(ai, options) { options = options || {}; var answer = stripModelSourcesSection(String(ai && ai.content || '').trim()); if (shouldRegenerateTruncatedAnswer(answer, ai && ai.finishReason) && typeof options.callAI === 'function') { console.warn('[clinical-assistant] answer looked truncated; regenerating final answer', { finishReason: ai && ai.finishReason, chars: answer.length, streamed: Boolean(options.streamed) }); if (typeof options.onRegenerating === 'function') options.onRegenerating(); - var completed = await options.callAI(options.messages, { + var completed = await options.callAI(options.messages, Object.assign({}, options.generationOptions || {}, { model: options.chatModel || undefined, - temperature: 0.15, maxTokens: 5000 - }); + })); answer = stripModelSourcesSection(String(completed.content || '').trim()) || answer; ai.model = completed.model || ai.model; ai.provider = completed.provider || ai.provider; @@ -60,6 +66,7 @@ function shouldRegenerateTruncatedAnswer(answer, finishReason) { module.exports = { buildSystemPrompt: buildSystemPrompt, buildUserPrompt: buildUserPrompt, + assistantGenerationOptions: assistantGenerationOptions, finalizeAssistantAnswer: finalizeAssistantAnswer, stripModelSourcesSection: stripModelSourcesSection, shouldRegenerateTruncatedAnswer: shouldRegenerateTruncatedAnswer diff --git a/src/utils/generationOptions.js b/src/utils/generationOptions.js new file mode 100644 index 0000000..21705a1 --- /dev/null +++ b/src/utils/generationOptions.js @@ -0,0 +1,11 @@ +function resolveGenerationOptions(options) { + options = options || {}; + return { + temperature: options.temperature ?? 0.3, + maxTokens: options.maxTokens ?? 4000, + reasoningEffort: options.reasoningEffort, + reasoningFormat: options.reasoningFormat + }; +} + +module.exports = { resolveGenerationOptions }; diff --git a/test/clinical-generation-options.test.js b/test/clinical-generation-options.test.js new file mode 100644 index 0000000..2f85a7e --- /dev/null +++ b/test/clinical-generation-options.test.js @@ -0,0 +1,19 @@ +const test = require('node:test'); +const assert = require('node:assert/strict'); + +const { assistantGenerationOptions } = require('../src/utils/clinicalAnswer'); +const { resolveGenerationOptions } = require('../src/utils/generationOptions'); + +test('clinical assistant profile requests low reasoning without exposing it', () => { + assert.deepEqual(assistantGenerationOptions({ temperature: 0, maxTokens: 80 }), { + reasoningEffort: 'low', + reasoningFormat: 'hidden', + temperature: 0, + maxTokens: 80 + }); +}); + +test('generation defaults preserve an explicit zero temperature', () => { + assert.equal(resolveGenerationOptions({ temperature: 0 }).temperature, 0); + assert.equal(resolveGenerationOptions({}).temperature, 0.3); +});