217 lines
10 KiB
JavaScript
217 lines
10 KiB
JavaScript
#!/usr/bin/env node
|
|
|
|
import fs from 'node:fs/promises';
|
|
import { fileURLToPath } from 'node:url';
|
|
|
|
function hitId(hit) {
|
|
return String(typeof hit === 'object' && hit !== null ? hit.fragmentId : hit);
|
|
}
|
|
|
|
function hitSource(hit) {
|
|
return typeof hit === 'object' && hit !== null ? String(hit.sourceName || hit.title || '') : '';
|
|
}
|
|
|
|
export function reciprocalRank(retrieved, relevant) {
|
|
const relevantSet = new Set(relevant.map(String));
|
|
const rank = retrieved.findIndex((hit) => relevantSet.has(hitId(hit)));
|
|
return rank < 0 ? 0 : 1 / (rank + 1);
|
|
}
|
|
|
|
export function recallAtK(retrieved, relevant) {
|
|
const relevantSet = new Set(relevant.map(String));
|
|
if (relevantSet.size === 0) return 0;
|
|
const hits = retrieved.filter((hit) => relevantSet.has(hitId(hit)));
|
|
return new Set(hits.map(hitId)).size / relevantSet.size;
|
|
}
|
|
|
|
export function ndcgAtK(retrieved, relevant) {
|
|
const relevantSet = new Set(relevant.map(String));
|
|
if (relevantSet.size === 0) return 0;
|
|
const dcg = retrieved.reduce((sum, hit, index) =>
|
|
sum + (relevantSet.has(hitId(hit)) ? 1 / Math.log2(index + 2) : 0), 0);
|
|
const idealLength = Math.min(relevantSet.size, retrieved.length);
|
|
const ideal = Array.from({ length: idealLength }, (_, index) => 1 / Math.log2(index + 2))
|
|
.reduce((sum, value) => sum + value, 0);
|
|
return ideal === 0 ? 0 : dcg / ideal;
|
|
}
|
|
|
|
export function evaluateCase(retrieved, testCase, k) {
|
|
const ranked = retrieved.slice(0, k);
|
|
const relevant = testCase.relevantFragmentIds || [];
|
|
const forbiddenIds = new Set((testCase.forbiddenFragmentIds || []).map(String));
|
|
const forbiddenSources = new Set((testCase.forbiddenSourceNames || []).map((value) => value.toLowerCase()));
|
|
const requiredSources = new Set((testCase.requiredSourceNames || []).map((value) => value.toLowerCase()));
|
|
const returnedSources = new Set(ranked.map(hitSource).filter(Boolean).map((value) => value.toLowerCase()));
|
|
const forbiddenLeak = ranked.some((hit) => forbiddenIds.has(hitId(hit))
|
|
|| forbiddenSources.has(hitSource(hit).toLowerCase()));
|
|
const missingRequiredSource = requiredSources.size > 0
|
|
&& ![...requiredSources].some((source) => returnedSources.has(source));
|
|
const expectedNoEvidence = Boolean(testCase.expectedNoEvidence);
|
|
return {
|
|
id: testCase.id,
|
|
caseType: testCase.caseType || (expectedNoEvidence ? 'no-answer' : 'answerable'),
|
|
answerable: !expectedNoEvidence,
|
|
recallAtK: expectedNoEvidence ? null : recallAtK(ranked, relevant),
|
|
reciprocalRank: expectedNoEvidence ? null : reciprocalRank(ranked, relevant),
|
|
ndcgAtK: expectedNoEvidence ? null : ndcgAtK(ranked, relevant),
|
|
forbiddenLeak,
|
|
wrongSource: forbiddenLeak || missingRequiredSource,
|
|
returned: ranked.length,
|
|
noEvidence: ranked.length === 0,
|
|
noAnswerFalsePositive: expectedNoEvidence && ranked.length > 0
|
|
};
|
|
}
|
|
|
|
function rate(rows, predicate) {
|
|
return rows.length ? rows.filter(predicate).length / rows.length : 0;
|
|
}
|
|
|
|
function average(rows, key) {
|
|
return rows.length ? rows.reduce((sum, row) => sum + row[key], 0) / rows.length : 0;
|
|
}
|
|
|
|
export function summarize(results) {
|
|
const answerable = results.filter((row) => row.answerable);
|
|
const noAnswer = results.filter((row) => !row.answerable);
|
|
const byCaseType = Object.fromEntries([...new Set(results.map((row) => row.caseType))].sort().map((caseType) => {
|
|
const rows = results.filter((row) => row.caseType === caseType);
|
|
return [caseType, { cases: rows.length, forbiddenLeakageRate: rate(rows, (row) => row.forbiddenLeak) }];
|
|
}));
|
|
return {
|
|
cases: results.length,
|
|
answerableCases: answerable.length,
|
|
noAnswerCases: noAnswer.length,
|
|
recallAtK: average(answerable, 'recallAtK'),
|
|
mrr: average(answerable, 'reciprocalRank'),
|
|
ndcgAtK: average(answerable, 'ndcgAtK'),
|
|
forbiddenLeakageRate: rate(results, (row) => row.forbiddenLeak),
|
|
wrongSourceRate: rate(results, (row) => row.wrongSource),
|
|
noAnswerFalsePositiveRate: rate(noAnswer, (row) => row.noAnswerFalsePositive),
|
|
noAnswerPrecision: 1 - rate(noAnswer, (row) => row.noAnswerFalsePositive),
|
|
byCaseType
|
|
};
|
|
}
|
|
|
|
export function thresholdFailures(summary, thresholds) {
|
|
const failures = [];
|
|
if (thresholds.minRecall !== undefined && summary.recallAtK < thresholds.minRecall) {
|
|
failures.push(`recallAtK ${summary.recallAtK} < ${thresholds.minRecall}`);
|
|
}
|
|
if (thresholds.minNdcg !== undefined && summary.ndcgAtK < thresholds.minNdcg) {
|
|
failures.push(`ndcgAtK ${summary.ndcgAtK} < ${thresholds.minNdcg}`);
|
|
}
|
|
if (thresholds.minMrr !== undefined && summary.mrr < thresholds.minMrr) {
|
|
failures.push(`mrr ${summary.mrr} < ${thresholds.minMrr}`);
|
|
}
|
|
if (thresholds.maxLeakage !== undefined && summary.forbiddenLeakageRate > thresholds.maxLeakage) {
|
|
failures.push(`forbiddenLeakageRate ${summary.forbiddenLeakageRate} > ${thresholds.maxLeakage}`);
|
|
}
|
|
if (thresholds.maxNoAnswerFalsePositive !== undefined
|
|
&& summary.noAnswerFalsePositiveRate > thresholds.maxNoAnswerFalsePositive) {
|
|
failures.push(`noAnswerFalsePositiveRate ${summary.noAnswerFalsePositiveRate} > ${thresholds.maxNoAnswerFalsePositive}`);
|
|
}
|
|
if (thresholds.maxWrongSource !== undefined && summary.wrongSourceRate > thresholds.maxWrongSource) {
|
|
failures.push(`wrongSourceRate ${summary.wrongSourceRate} > ${thresholds.maxWrongSource}`);
|
|
}
|
|
return failures;
|
|
}
|
|
|
|
export function parseArgs(argv) {
|
|
const options = { baseUrl: process.env.AIHR_EVAL_BASE_URL || 'http://127.0.0.1:8080', k: 5 };
|
|
for (let index = 0; index < argv.length; index += 1) {
|
|
const value = argv[index];
|
|
if (value === '--dataset') options.dataset = argv[++index];
|
|
else if (value === '--base-url') options.baseUrl = argv[++index];
|
|
else if (value === '--token') options.token = argv[++index];
|
|
else if (value === '--clientid') options.clientid = argv[++index];
|
|
else if (value === '--k') options.k = Number(argv[++index]);
|
|
else if (value === '--dry-run') options.dryRun = true;
|
|
else if (value === '--min-recall') options.minRecall = Number(argv[++index]);
|
|
else if (value === '--min-ndcg') options.minNdcg = Number(argv[++index]);
|
|
else if (value === '--min-mrr') options.minMrr = Number(argv[++index]);
|
|
else if (value === '--max-leakage') options.maxLeakage = Number(argv[++index]);
|
|
else if (value === '--max-wrong-source') options.maxWrongSource = Number(argv[++index]);
|
|
else if (value === '--max-no-answer-false-positive') options.maxNoAnswerFalsePositive = Number(argv[++index]);
|
|
else throw new Error(`Unknown option: ${value}`);
|
|
}
|
|
if (!options.dataset) throw new Error('--dataset is required');
|
|
if (!Number.isInteger(options.k) || options.k < 1 || options.k > 50) {
|
|
throw new Error('--k must be an integer between 1 and 50');
|
|
}
|
|
for (const key of ['minRecall', 'minNdcg', 'minMrr', 'maxLeakage', 'maxWrongSource', 'maxNoAnswerFalsePositive']) {
|
|
if (options[key] !== undefined && (!Number.isFinite(options[key]) || options[key] < 0 || options[key] > 1)) {
|
|
throw new Error(`--${key.replace(/[A-Z]/g, (letter) => `-${letter.toLowerCase()}`)} must be between 0 and 1`);
|
|
}
|
|
}
|
|
return options;
|
|
}
|
|
|
|
export function validateDataset(dataset) {
|
|
if (dataset?.schemaVersion !== 1) throw new Error('dataset.schemaVersion must be 1');
|
|
if (!dataset || !Array.isArray(dataset.cases) || dataset.cases.length === 0) {
|
|
throw new Error('dataset.cases must be a non-empty array');
|
|
}
|
|
const ids = new Set();
|
|
for (const item of dataset.cases) {
|
|
if (!item.id || ids.has(item.id) || !item.query) throw new Error(`invalid or duplicate case: ${item.id || '<missing>'}`);
|
|
ids.add(item.id);
|
|
if (!item.expectedNoEvidence && (!Array.isArray(item.relevantFragmentIds) || item.relevantFragmentIds.length === 0)) {
|
|
throw new Error(`answerable case must define relevantFragmentIds: ${item.id}`);
|
|
}
|
|
if (item.expectedNoEvidence && (item.relevantFragmentIds || []).length > 0) {
|
|
throw new Error(`no-answer case cannot define relevantFragmentIds: ${item.id}`);
|
|
}
|
|
}
|
|
}
|
|
|
|
async function requestSearch(options, testCase, fetchImpl) {
|
|
const headers = { 'content-type': 'application/json' };
|
|
if (options.token) headers.Authorization = options.token.startsWith('Bearer ') ? options.token : `Bearer ${options.token}`;
|
|
if (options.clientid) headers.clientid = options.clientid;
|
|
const response = await fetchImpl(`${options.baseUrl.replace(/\/$/, '')}/api/knowledge/search`, {
|
|
method: 'POST', headers,
|
|
body: JSON.stringify({
|
|
queryText: testCase.query,
|
|
category: testCase.category || undefined,
|
|
position: testCase.position || undefined,
|
|
source: 'knowledge_search',
|
|
limit: options.k
|
|
})
|
|
});
|
|
const body = await response.json();
|
|
if (!response.ok || body.code !== 200) {
|
|
throw new Error(`case ${testCase.id} search failed with HTTP ${response.status} code ${body.code}`);
|
|
}
|
|
const snippets = body.data?.snippets || [];
|
|
return snippets.filter((snippet) => snippet.fragmentId !== null && snippet.fragmentId !== undefined)
|
|
.map((snippet) => ({ fragmentId: snippet.fragmentId, sourceName: snippet.title || '' }));
|
|
}
|
|
|
|
export async function run(options, dependencies = {}) {
|
|
const dataset = JSON.parse(await fs.readFile(options.dataset, 'utf8'));
|
|
validateDataset(dataset);
|
|
if (options.dryRun) {
|
|
return { datasetCode: dataset.datasetCode, schemaVersion: dataset.schemaVersion,
|
|
cases: dataset.cases.length, k: options.k, dryRun: true };
|
|
}
|
|
const fetchImpl = dependencies.fetchImpl || globalThis.fetch;
|
|
const results = [];
|
|
for (const testCase of dataset.cases) {
|
|
const retrieved = await requestSearch(options, testCase, fetchImpl);
|
|
results.push(evaluateCase(retrieved, testCase, options.k));
|
|
}
|
|
const summary = summarize(results);
|
|
const failures = thresholdFailures(summary, options);
|
|
return { datasetCode: dataset.datasetCode, schemaVersion: dataset.schemaVersion, k: options.k,
|
|
summary, thresholdPassed: failures.length === 0, thresholdFailures: failures, cases: results };
|
|
}
|
|
|
|
if (fileURLToPath(import.meta.url) === process.argv[1]) {
|
|
run(parseArgs(process.argv.slice(2)))
|
|
.then((result) => {
|
|
console.log(JSON.stringify(result, null, 2));
|
|
if (result.thresholdPassed === false) process.exitCode = 2;
|
|
})
|
|
.catch((error) => { console.error(error.message); process.exitCode = 1; });
|
|
}
|