Configure clinical assistant reasoning profile
All checks were successful
Forgejo Android APK / Build signed APK (push) Successful in 1m58s
All checks were successful
Forgejo Android APK / Build signed APK (push) Successful in 1m58s
This commit is contained in:
parent
f556d50a09
commit
e710b1c7bd
5 changed files with 73 additions and 20 deletions
|
|
@ -36,6 +36,7 @@ var {
|
||||||
var {
|
var {
|
||||||
buildSystemPrompt,
|
buildSystemPrompt,
|
||||||
buildUserPrompt,
|
buildUserPrompt,
|
||||||
|
assistantGenerationOptions,
|
||||||
finalizeAssistantAnswer
|
finalizeAssistantAnswer
|
||||||
} = require('../utils/clinicalAnswer');
|
} = require('../utils/clinicalAnswer');
|
||||||
|
|
||||||
|
|
@ -182,12 +183,17 @@ router.post('/clinical-assistant/chat', async function(req, res) {
|
||||||
var prepared = await prepareAssistantChat(req.body);
|
var prepared = await prepareAssistantChat(req.body);
|
||||||
if (prepared.direct) return res.json(prepared.direct);
|
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,
|
model: prepared.chatModel || undefined,
|
||||||
temperature: 0.15,
|
temperature: 0.15,
|
||||||
maxTokens: 2600
|
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;
|
var answer = finalized.answer;
|
||||||
ai = finalized.ai;
|
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('sources', { sources: safeSources, search: prepared.search });
|
||||||
sendEvent('status', { message: 'Generating answer...' });
|
sendEvent('status', { message: 'Generating answer...' });
|
||||||
|
|
||||||
var ai = await callAIStream(prepared.messages, {
|
var ai = await callAIStream(prepared.messages, assistantGenerationOptions({
|
||||||
model: prepared.chatModel || undefined,
|
model: prepared.chatModel || undefined,
|
||||||
temperature: 0.15,
|
temperature: 0.15,
|
||||||
maxTokens: 2600
|
maxTokens: 2600
|
||||||
}, function(delta) {
|
}), function(delta) {
|
||||||
sendEvent('token', { token: delta });
|
sendEvent('token', { token: delta });
|
||||||
});
|
});
|
||||||
|
|
||||||
|
|
@ -249,6 +255,7 @@ router.post('/clinical-assistant/chat/stream', async function(req, res) {
|
||||||
messages: prepared.messages,
|
messages: prepared.messages,
|
||||||
chatModel: prepared.chatModel,
|
chatModel: prepared.chatModel,
|
||||||
callAI: callAI,
|
callAI: callAI,
|
||||||
|
generationOptions: assistantGenerationOptions({ temperature: 0.15 }),
|
||||||
streamed: true,
|
streamed: true,
|
||||||
onRegenerating: function() { sendEvent('status', { message: 'Completing answer...' }); }
|
onRegenerating: function() { sendEvent('status', { message: 'Completing answer...' }); }
|
||||||
});
|
});
|
||||||
|
|
@ -467,11 +474,11 @@ async function rewriteSearchQuery(message, history, chatModel) {
|
||||||
role: 'user',
|
role: 'user',
|
||||||
content: 'Conversation:\n' + hist + '\n\nLatest user question:\n' + message + '\n\nStandalone search query:'
|
content: 'Conversation:\n' + hist + '\n\nLatest user question:\n' + message + '\n\nStandalone search query:'
|
||||||
}
|
}
|
||||||
], {
|
], assistantGenerationOptions({
|
||||||
model: chatModel || undefined,
|
model: chatModel || undefined,
|
||||||
temperature: 0,
|
temperature: 0,
|
||||||
maxTokens: 80
|
maxTokens: 80
|
||||||
});
|
}));
|
||||||
var rewritten = String(ai.content || '').replace(/^['"]|['"]$/g, '').replace(/\s+/g, ' ').trim();
|
var rewritten = String(ai.content || '').replace(/^['"]|['"]$/g, '').replace(/\s+/g, ' ').trim();
|
||||||
if (!rewritten || rewritten.length < 6 || rewritten.length > 300) return message;
|
if (!rewritten || rewritten.length < 6 || rewritten.length > 300) return message;
|
||||||
if (/^(yes|no|maybe|i don'?t know)$/i.test(rewritten)) return message;
|
if (/^(yes|no|maybe|i don'?t know)$/i.test(rewritten)) return message;
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@
|
||||||
const { OpenAI } = require('openai');
|
const { OpenAI } = require('openai');
|
||||||
const { DEFAULT_MODEL, FALLBACK_MODEL, getBedrockModelId, getBedrockMaxOut } = require('./models');
|
const { DEFAULT_MODEL, FALLBACK_MODEL, getBedrockModelId, getBedrockMaxOut } = require('./models');
|
||||||
const logger = require('./logger');
|
const logger = require('./logger');
|
||||||
|
const { resolveGenerationOptions } = require('./generationOptions');
|
||||||
|
|
||||||
var activeProvider = process.env.AI_PROVIDER || (process.env.LITELLM_API_BASE ? 'litellm' : 'openrouter');
|
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)
|
// 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');
|
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,
|
model: model,
|
||||||
messages: messages,
|
messages: messages,
|
||||||
temperature: temperature,
|
temperature: temperature,
|
||||||
max_tokens: maxTokens
|
max_tokens: maxTokens
|
||||||
});
|
}, generation || {}));
|
||||||
|
|
||||||
return {
|
return {
|
||||||
success: true,
|
success: true,
|
||||||
|
|
@ -421,8 +428,9 @@ async function callAIStream(messages, options, onToken) {
|
||||||
options = options || {};
|
options = options || {};
|
||||||
var requestedModel = options.model;
|
var requestedModel = options.model;
|
||||||
var model = await resolveModel(requestedModel);
|
var model = await resolveModel(requestedModel);
|
||||||
var temperature = options.temperature || 0.3;
|
var generation = resolveGenerationOptions(options);
|
||||||
var maxTokens = options.maxTokens || 4000;
|
var temperature = generation.temperature;
|
||||||
|
var maxTokens = generation.maxTokens;
|
||||||
var startTime = Date.now();
|
var startTime = Date.now();
|
||||||
await assertModelAllowed(model, options);
|
await assertModelAllowed(model, options);
|
||||||
|
|
||||||
|
|
@ -443,13 +451,13 @@ async function callAIStream(messages, options, onToken) {
|
||||||
|
|
||||||
var content = '';
|
var content = '';
|
||||||
var finishReason = null;
|
var finishReason = null;
|
||||||
var stream = await client.chat.completions.create({
|
var stream = await client.chat.completions.create(addReasoningOptions({
|
||||||
model: model,
|
model: model,
|
||||||
messages: messages,
|
messages: messages,
|
||||||
temperature: temperature,
|
temperature: temperature,
|
||||||
max_tokens: maxTokens,
|
max_tokens: maxTokens,
|
||||||
stream: true
|
stream: true
|
||||||
});
|
}, generation));
|
||||||
for await (var part of stream) {
|
for await (var part of stream) {
|
||||||
var choice = part && part.choices && part.choices[0] ? part.choices[0] : null;
|
var choice = part && part.choices && part.choices[0] ? part.choices[0] : null;
|
||||||
if (choice && choice.finish_reason) finishReason = choice.finish_reason;
|
if (choice && choice.finish_reason) finishReason = choice.finish_reason;
|
||||||
|
|
@ -470,8 +478,9 @@ async function callAI(messages, options) {
|
||||||
options = options || {};
|
options = options || {};
|
||||||
var requestedModel = options.model;
|
var requestedModel = options.model;
|
||||||
var model = await resolveModel(requestedModel);
|
var model = await resolveModel(requestedModel);
|
||||||
var temperature = options.temperature || 0.3;
|
var generation = resolveGenerationOptions(options);
|
||||||
var maxTokens = options.maxTokens || 4000;
|
var temperature = generation.temperature;
|
||||||
|
var maxTokens = generation.maxTokens;
|
||||||
var startTime = Date.now();
|
var startTime = Date.now();
|
||||||
|
|
||||||
// Server-side whitelist: reject any model the operator hasn't enabled.
|
// 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) {
|
} else if (activeProvider === 'vertex' && vertexClient) {
|
||||||
result = await callVertex(messages, model, temperature, maxTokens);
|
result = await callVertex(messages, model, temperature, maxTokens);
|
||||||
} else if (activeProvider === 'litellm' && litellmClient) {
|
} else if (activeProvider === 'litellm' && litellmClient) {
|
||||||
result = await callLiteLLM(messages, model, temperature, maxTokens);
|
result = await callLiteLLM(messages, model, temperature, maxTokens, generation);
|
||||||
} else if (openrouter) {
|
} else if (openrouter) {
|
||||||
result = await callOpenRouter(messages, model, temperature, maxTokens);
|
result = await callOpenRouter(messages, model, temperature, maxTokens);
|
||||||
} else {
|
} else {
|
||||||
|
|
@ -552,7 +561,7 @@ async function callAI(messages, options) {
|
||||||
if (activeProvider === 'litellm' && model !== FALLBACK_MODEL && litellmClient) {
|
if (activeProvider === 'litellm' && model !== FALLBACK_MODEL && litellmClient) {
|
||||||
logger.warn('Trying fallback model on LiteLLM: ' + FALLBACK_MODEL);
|
logger.warn('Trying fallback model on LiteLLM: ' + FALLBACK_MODEL);
|
||||||
try {
|
try {
|
||||||
var litellmFallback = await callLiteLLM(messages, FALLBACK_MODEL, temperature, maxTokens);
|
var litellmFallback = await callLiteLLM(messages, FALLBACK_MODEL, temperature, maxTokens, generation);
|
||||||
litellmFallback.fallback = true;
|
litellmFallback.fallback = true;
|
||||||
litellmFallback.duration = Date.now() - startTime;
|
litellmFallback.duration = Date.now() - startTime;
|
||||||
logger.info('LiteLLM fallback success', { model: FALLBACK_MODEL });
|
logger.info('LiteLLM fallback success', { model: FALLBACK_MODEL });
|
||||||
|
|
|
||||||
|
|
@ -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.';
|
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) {
|
async function finalizeAssistantAnswer(ai, options) {
|
||||||
options = options || {};
|
options = options || {};
|
||||||
var answer = stripModelSourcesSection(String(ai && ai.content || '').trim());
|
var answer = stripModelSourcesSection(String(ai && ai.content || '').trim());
|
||||||
if (shouldRegenerateTruncatedAnswer(answer, ai && ai.finishReason) && typeof options.callAI === 'function') {
|
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) });
|
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();
|
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,
|
model: options.chatModel || undefined,
|
||||||
temperature: 0.15,
|
|
||||||
maxTokens: 5000
|
maxTokens: 5000
|
||||||
});
|
}));
|
||||||
answer = stripModelSourcesSection(String(completed.content || '').trim()) || answer;
|
answer = stripModelSourcesSection(String(completed.content || '').trim()) || answer;
|
||||||
ai.model = completed.model || ai.model;
|
ai.model = completed.model || ai.model;
|
||||||
ai.provider = completed.provider || ai.provider;
|
ai.provider = completed.provider || ai.provider;
|
||||||
|
|
@ -60,6 +66,7 @@ function shouldRegenerateTruncatedAnswer(answer, finishReason) {
|
||||||
module.exports = {
|
module.exports = {
|
||||||
buildSystemPrompt: buildSystemPrompt,
|
buildSystemPrompt: buildSystemPrompt,
|
||||||
buildUserPrompt: buildUserPrompt,
|
buildUserPrompt: buildUserPrompt,
|
||||||
|
assistantGenerationOptions: assistantGenerationOptions,
|
||||||
finalizeAssistantAnswer: finalizeAssistantAnswer,
|
finalizeAssistantAnswer: finalizeAssistantAnswer,
|
||||||
stripModelSourcesSection: stripModelSourcesSection,
|
stripModelSourcesSection: stripModelSourcesSection,
|
||||||
shouldRegenerateTruncatedAnswer: shouldRegenerateTruncatedAnswer
|
shouldRegenerateTruncatedAnswer: shouldRegenerateTruncatedAnswer
|
||||||
|
|
|
||||||
11
src/utils/generationOptions.js
Normal file
11
src/utils/generationOptions.js
Normal file
|
|
@ -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 };
|
||||||
19
test/clinical-generation-options.test.js
Normal file
19
test/clinical-generation-options.test.js
Normal file
|
|
@ -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);
|
||||||
|
});
|
||||||
Loading…
Reference in a new issue