Files
prop-ai-hr/scripts/evaluate-knowledge-quality.mjs

294 lines
14 KiB
JavaScript
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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 normalizeSourceName(value) {
return String(value || '')
.replace(/^第\s*\d+\s*段[::]\s*/u, '')
.replace(/\.(?:pdf|docx?|pptx?|xlsx?|txt|md)$/iu, '')
.replace(/\s+/gu, '')
.toLowerCase();
}
function sourceMatches(actual, expected) {
const left = normalizeSourceName(actual);
const right = normalizeSourceName(expected);
return Boolean(left && right && (left.includes(right) || right.includes(left)));
}
function requiredSourceMatches(hit, source, authorities, kinds) {
if (!sourceMatches(hitSource(hit), source)) return false;
if (typeof hit !== 'object' || hit === null) return authorities.length === 0 && kinds.length === 0;
const authority = String(hit.sourceAuthority || '').toUpperCase();
const kind = String(hit.sourceKind || '').toUpperCase();
return (authorities.length === 0 || authorities.includes(authority))
&& (kinds.length === 0 || kinds.includes(kind));
}
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 credited = new Set();
const dcg = retrieved.reduce((sum, hit, index) => {
const id = hitId(hit);
if (!relevantSet.has(id) || credited.has(id)) return sum;
credited.add(id);
return sum + 1 / Math.log2(index + 2);
}, 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, actualNoEvidence = retrieved.length === 0) {
const ranked = retrieved.slice(0, k);
const relevant = testCase.relevantFragmentIds || [];
const forbiddenIds = new Set((testCase.forbiddenFragmentIds || []).map(String));
const forbiddenSources = testCase.forbiddenSourceNames || [];
const requiredSources = testCase.requiredSourceNames || [];
const requiredAuthorities = (testCase.requiredSourceAuthorities || []).map((value) => String(value).toUpperCase());
const requiredKinds = (testCase.requiredSourceKinds || []).map((value) => String(value).toUpperCase());
const visibleEvidence = ranked.filter((hit) => typeof hit !== 'object' || hit === null
|| hit.selectedForEvidence === undefined || hit.selectedForEvidence === true);
const forbiddenLeak = visibleEvidence.some((hit) => forbiddenIds.has(hitId(hit))
|| forbiddenSources.some((source) => sourceMatches(hitSource(hit), source)));
const missingRequiredSource = requiredSources.length > 0
&& !requiredSources.some((source) => ranked.some((hit) =>
requiredSourceMatches(hit, source, requiredAuthorities, requiredKinds)));
const sourceRanking = requiredSources.length > 0 && relevant.length === 0
? ranked.map((hit, index) => requiredSources.some((source) =>
requiredSourceMatches(hit, source, requiredAuthorities, requiredKinds))
? normalizeSourceName(hitSource(hit)) : `__nonmatching_source_${index}`) : null;
const normalizedRelevantSources = requiredSources.map(normalizeSourceName);
const metricHits = sourceRanking || ranked;
const metricRelevant = sourceRanking ? normalizedRelevantSources : relevant;
const expectedNoEvidence = Boolean(testCase.expectedNoEvidence);
const retrievalEvaluated = metricRelevant.length > 0;
return {
id: testCase.id,
caseType: testCase.caseType || (expectedNoEvidence ? 'no-answer' : 'answerable'),
answerable: retrievalEvaluated,
retrievalEvaluated,
expectedNoEvidence,
recallAtK: retrievalEvaluated ? recallAtK(metricHits, metricRelevant) : null,
recallAt5: retrievalEvaluated ? recallAtK(metricHits.slice(0, 5), metricRelevant) : null,
recallAt10: retrievalEvaluated ? recallAtK(metricHits.slice(0, 10), metricRelevant) : null,
recallAt20: retrievalEvaluated ? recallAtK(metricHits.slice(0, 20), metricRelevant) : null,
reciprocalRank: retrievalEvaluated ? reciprocalRank(metricHits, metricRelevant) : null,
ndcgAtK: retrievalEvaluated ? ndcgAtK(metricHits, metricRelevant) : null,
forbiddenLeak,
wrongSource: forbiddenLeak || missingRequiredSource,
returned: ranked.length,
noEvidence: actualNoEvidence,
noAnswerFalsePositive: expectedNoEvidence && !actualNoEvidence,
topCandidates: ranked.slice(0, 5).map((hit, index) => ({
rank: index + 1,
fragmentId: typeof hit === 'object' && hit !== null ? hit.fragmentId : hit,
sourceName: hitSource(hit),
sourceAuthority: typeof hit === 'object' && hit !== null ? hit.sourceAuthority : null,
sourceKind: typeof hit === 'object' && hit !== null ? hit.sourceKind : null,
selectedForEvidence: typeof hit === 'object' && hit !== null ? hit.selectedForEvidence : undefined
}))
};
}
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.retrievalEvaluated ?? row.answerable);
const noAnswer = results.filter((row) => row.expectedNoEvidence ?? !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'),
recallAt5: average(answerable, 'recallAt5'),
recallAt10: average(answerable, 'recallAt10'),
recallAt20: average(answerable, 'recallAt20'),
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)
&& (!Array.isArray(item.requiredSourceNames) || item.requiredSourceNames.length === 0)) {
throw new Error(`answerable case must define relevantFragmentIds or requiredSourceNames: ${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/query`, {
method: 'POST', headers,
body: JSON.stringify({
queryText: testCase.query,
category: testCase.category || undefined,
position: testCase.position || undefined,
source: testCase.source || 'mobile_uni_agent',
// Response display remains capped at 10; retrievalCandidates is the decoupled pre-generation pool.
limit: Math.min(options.k, 10)
})
});
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}: ${body.msg || ''}`);
}
const candidates = body.data?.retrievalCandidates;
const fallback = body.data?.citations || body.data?.legacy?.snippets || [];
const rows = Array.isArray(candidates) && candidates.length > 0
? [...candidates].sort((left, right) => {
const leftRank = Number.isInteger(left.candidateRank) ? left.candidateRank : Number.MAX_SAFE_INTEGER;
const rightRank = Number.isInteger(right.candidateRank) ? right.candidateRank : Number.MAX_SAFE_INTEGER;
return leftRank - rightRank;
})
: fallback;
return {
noEvidence: Boolean(body.data?.noEvidence),
retrieved: rows.filter((item) => item.fragmentId !== null && item.fragmentId !== undefined)
.map((item) => ({
fragmentId: item.fragmentId,
sourceName: item.title || '',
sourceAuthority: item.sourceAuthority || null,
sourceKind: item.sourceKind || null,
candidateRank: item.candidateRank,
selectedForEvidence: item.selectedForEvidence
}))
};
}
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 result = await requestSearch(options, testCase, fetchImpl);
results.push(evaluateCase(result.retrieved, testCase, options.k, result.noEvidence));
}
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; });
}