#!/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 || ''}`); 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; }); }