#!/usr/bin/env python3
"""Atelier Nymphar, 2026-10-05. Corpus fictif. Aucun appel réseau par défaut.
API optionnelle : Voyage historique, api.voyageai.com (pas endpoint Atlas UE).
"""
import argparse
import json
import math
import os
import time
from datetime import datetime, timezone
from pathlib import Path
from urllib.error import HTTPError, URLError
from urllib.request import Request, urlopen

PRICES = {'rerank-2.5': 0.05, 'rerank-3': 0.05, 'rerank-3-lite': 0.02}
ENDPOINT = 'https://api.voyageai.com/v1/rerank'


def metrics(order, relevant, k=3):
    """Rappel avec tous les positifs annotés, même absents des candidats."""
    if len(order) != len(set(order)):
        raise ValueError('Identifiants dupliqués dans le classement')
    positives = {doc for doc, grade in relevant.items() if grade > 0}
    if not positives:
        return {'recall_at_3': None, 'mrr_at_3': None, 'ndcg_at_3': None}
    top = order[:k]
    recall = len(set(top) & positives) / len(positives)
    mrr = next((1 / (i + 1) for i, doc in enumerate(top) if doc in positives), 0)
    dcg = sum((2 ** relevant.get(doc, 0) - 1) / math.log2(i + 2)
              for i, doc in enumerate(top))
    ideal = sum((2 ** grade - 1) / math.log2(i + 2)
                for i, grade in enumerate(sorted(relevant.values(), reverse=True)[:k]))
    return {'recall_at_3': recall, 'mrr_at_3': mrr, 'ndcg_at_3': dcg / ideal}


def map_results(data, candidates, expected):
    indexes = [row['index'] for row in data]
    if len(indexes) != expected or len(indexes) != len(set(indexes)):
        raise ValueError('Résultat incomplet ou indices dupliqués')
    if any(type(i) is not int or i < 0 or i >= len(candidates) for i in indexes):
        raise ValueError('Indice de document invalide')
    return [candidates[i] for i in indexes]


def validate(corpus):
    docs = corpus['documents']
    seen = set()
    for case in corpus['cases']:
        if case['id'] in seen:
            raise ValueError('Identifiant de question dupliqué')
        seen.add(case['id'])
        ids = case['candidates']
        if not ids or len(ids) > 1000 or len(ids) != len(set(ids)):
            raise ValueError('Liste de candidats invalide')
        if any(i not in docs for i in ids):
            raise ValueError('Document candidat absent du corpus')
        if any(i not in docs or g not in (0, 1, 2) for i, g in case['relevance'].items()):
            raise ValueError('Annotation invalide')


def evaluate(corpus, model=None):
    validate(corpus)
    key = os.environ.get('VOYAGE_API_KEY') if model else None
    if model and not key:
        raise ValueError('VOYAGE_API_KEY doit être configurée pour le mode live')
    rows = []
    for case in corpus['cases']:
        ids = case['candidates']
        row = {'case_id': case['id'], 'model': model or 'baseline_fictive',
               'candidates': ids, 'status': 'ok', 'duration_ms': None,
               'total_tokens': None, 'list_price_usd': None,
               'answerable_in_corpus': any(case['relevance'].values())}
        order = ids[:3]
        if model:
            body = {'model': model, 'query': case['query'],
                    'documents': [corpus['documents'][i] for i in ids],
                    'top_k': min(3, len(ids)), 'truncation': False,
                    'return_documents': False}
            request = Request(ENDPOINT, data=json.dumps(body).encode(), headers={
                'Authorization': 'Bearer ' + key, 'Content-Type': 'application/json'})
            started = time.perf_counter()
            try:
                with urlopen(request, timeout=30) as response:
                    result = json.load(response)
                order = map_results(result['data'], ids, min(3, len(ids)))
                row['duration_ms'] = round((time.perf_counter() - started) * 1000, 2)
                row['scores'] = [r['relevance_score'] for r in result['data']]
                tokens = result.get('usage', {}).get('total_tokens')
                row['total_tokens'] = tokens
                row['list_price_usd'] = tokens * PRICES[model] / 1_000_000 if tokens is not None else None
            except (HTTPError, URLError, TimeoutError, ValueError, KeyError, TypeError) as exc:
                row.update(status='error', error_type=type(exc).__name__,
                           http_status=getattr(exc, 'code', None), order=None,
                           recall_at_3=None, mrr_at_3=None, ndcg_at_3=None)
                rows.append(row)
                continue
        row.update(order=order, **metrics(order, case['relevance']))
        rows.append(row)
    return {'created_at': datetime.now(timezone.utc).isoformat(),
            'corpus_kind': corpus['kind'], 'mode': 'live' if model else 'baseline',
            'endpoint': ENDPOINT if model else None, 'k': 3,
            'note': 'Exercice fictif, pas un benchmark représentatif. Prix catalogue hors crédits.',
            'rows': rows, 'status': 'partial' if any(r['status'] == 'error' for r in rows) else 'complete'}


def self_test():
    assert metrics(['a', 'b', 'c'], {'a': 2}) == {
        'recall_at_3': 1, 'mrr_at_3': 1, 'ndcg_at_3': 1}
    delayed = metrics(['x', 'a', 'b'], {'a': 2})
    assert delayed['mrr_at_3'] == 0.5
    assert math.isclose(delayed['ndcg_at_3'], 1 / math.log2(3))
    assert metrics(['x'], {'absent': 2})['recall_at_3'] == 0
    assert metrics(['x'], {})['recall_at_3'] is None
    assert map_results([{'index': 2}, {'index': 0}], ['a', 'b', 'c'], 2) == ['c', 'a']
    for bad in [[{'index': 0}, {'index': 0}], [{'index': -1}], [{'index': 8}]]:
        try:
            map_results(bad, ['a', 'b', 'c'], len(bad))
        except ValueError:
            pass
        else:
            raise AssertionError('Une réponse invalide a été acceptée')
    print('OK : rangs, rappel global, NDCG, absence de réponse et indices invalides')


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--mode', choices=['baseline', 'live'], default='baseline')
    parser.add_argument('--model', choices=PRICES, default='rerank-3')
    parser.add_argument('--cases', type=Path, default=Path(__file__).with_name('cas_fr.json'))
    parser.add_argument('--output', type=Path)
    parser.add_argument('--self-test', action='store_true')
    args = parser.parse_args()
    if args.self_test:
        self_test()
        return
    report = evaluate(json.loads(args.cases.read_text()), args.model if args.mode == 'live' else None)
    rendered = json.dumps(report, ensure_ascii=False, indent=2) + '\n'
    if args.output:
        args.output.write_text(rendered)
    else:
        print(rendered, end='')
    if report['status'] == 'partial':
        raise SystemExit(1)


if __name__ == '__main__':
    main()
