"""A small lexical retrieval exercise, not an embedding or LLM benchmark.

Run: python compare.py
Writes results.csv and summary.json next to this script. No network requests.
"""
import csv
import json
from pathlib import Path


def rank(documents, question, include_title=False):
    terms = set(question.casefold().split())
    candidates = []
    for document in documents:
        text = document['body']
        if include_title:
            text = document['title'] + '\n' + text
        score = sum(term in text.casefold() for term in terms)
        if score:
            candidates.append({'id': document['id'], 'score': score, 'text': text})
    # Stable tie-breaking is explicit; ties must not be presented as confidence.
    return sorted(candidates, key=lambda item: (-item['score'], item['id']))


def compare(fixture):
    rows = []
    for query in fixture['queries']:
        for mode in ('body_only', 'title_and_body'):
            ranked = rank(fixture['documents'], query['text'], mode == 'title_and_body')
            top = ranked[0] if ranked else None
            expected = query['expected']
            predicted = top['id'] if top else None
            rows.append({
                'query_id': query['id'], 'question': query['text'], 'mode': mode,
                'expected': expected or '', 'predicted': predicted or '',
                'score': top['score'] if top else 0,
                'top_tie_count': sum(item['score'] == top['score'] for item in ranked) if top else 0,
                'correct': predicted == expected,
                'answerable': expected is not None,
                'retrieved_text': top['text'] if top else '',
            })
    summary = {}
    for mode in ('body_only', 'title_and_body'):
        answerable = [r for r in rows if r['mode'] == mode and r['answerable']]
        unanswerable = [r for r in rows if r['mode'] == mode and not r['answerable']]
        summary[mode] = {
            'answerable_top1_hits': sum(r['correct'] for r in answerable),
            'answerable_questions': len(answerable),
            'unanswerable_false_candidates': sum(bool(r['predicted']) for r in unanswerable),
            'unanswerable_questions': len(unanswerable),
        }
    return rows, summary


def main():
    root = Path(__file__).resolve().parent
    fixture = json.loads((root / 'fixture.json').read_text(encoding='utf-8'))
    rows, summary = compare(fixture)
    with (root / 'results.csv').open('w', encoding='utf-8-sig', newline='') as stream:
        writer = csv.DictWriter(stream, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)
    (root / 'summary.json').write_text(json.dumps(summary, indent=2) + '\n', encoding='utf-8')
    print(json.dumps(summary, indent=2))


if __name__ == '__main__':
    main()
