File size: 7,142 Bytes
0390c03
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
"""التحقق من هلوسة القرآن والحديث وتصحيحها: Gradio interface.

    python app.py                      # http://127.0.0.1:7860

Mode A verifies pasted text. Mode B asks a language model first and verifies its answer. API keys are taken from the
form or, preferably, from environment variables (GEMINI_API_KEY, OPENAI_API_KEY, HF_TOKEN) so that a deployment can
keep them as secrets.
"""
from __future__ import annotations

import json
import logging
import os
import threading
from pathlib import Path
from typing import List, Optional

import ui
from llm_client import PROVIDERS, LLMError, LLMSettings, generate
from verifier import MAX_INPUT_CHARS, IslamicContentVerifier

logger = logging.getLogger(__name__)

EXAMPLES_PATH = Path(__file__).resolve().parent / "demo" / "examples.json"
PROVIDER_LABELS = {"gemini": "جوجل جيميناي", "openai": "أوبن إيه آي", "huggingface": "هاغينغ فيس"}
KEY_ENV = {"gemini": "GEMINI_API_KEY", "openai": "OPENAI_API_KEY", "huggingface": "HF_TOKEN"}

_pipeline: Optional[IslamicContentVerifier] = None
_lock = threading.Lock()


def get_pipeline() -> IslamicContentVerifier:
    """Created once. The Quran index loads immediately; the Hadith index loads lazily (see ``warm_in_background``)."""
    global _pipeline
    with _lock:
        if _pipeline is None:
            _pipeline = IslamicContentVerifier()
        return _pipeline


def warm_in_background() -> None:
    threading.Thread(target=lambda: get_pipeline().retriever.warm(), daemon=True).start()


def load_examples(path: Path = EXAMPLES_PATH) -> List[dict]:
    try:
        with open(path, encoding="utf-8") as handle:
            return json.load(handle)
    except (OSError, json.JSONDecodeError):
        logger.exception("Could not load demo examples from %s", path)
        return []


def verify_text(text: str) -> str:
    """Mode A. Never raises: problems become Arabic notices."""
    if not text or not text.strip():
        return ui.render_message("الرجاء إدخال نص للتحقق منه.", "warn")
    try:
        return ui.render_results(get_pipeline().analyze(text))
    except ValueError:
        return ui.render_message(f"النص طويل جدًا (الحد الأقصى {MAX_INPUT_CHARS} حرف).", "warn")
    except Exception:
        logger.exception("Verification failed")
        return ui.render_message("حدث خطأ غير متوقع أثناء التحقق.", "bad")


def verify_generated_answer(answer: str) -> str:
    """Verify a model answer and show it above the report (also used by the in-browser page)."""
    try:
        return ui.render_results(get_pipeline().analyze(answer), generated_answer=answer)
    except Exception:
        logger.exception("Verification of the generated answer failed")
        return ui.render_message("تعذّر التحقق من إجابة النموذج.", "bad")


def analyze_benchmark(text: str, response_id: str = "R001") -> str:
    """IslamicEval-style JSON (1A/1B/1C rows + TSV) for a text. Used by the browser page's export button."""
    from benchmark import benchmark_json
    return benchmark_json(get_pipeline().analyze(text), response_id)


def detect_spans(text: str) -> str:
    """Detection only (Subtask 1A), as JSON ``[{label, start, end, text}]`` for the browser's model-selection step."""
    spans = get_pipeline().detect(text)
    return json.dumps([{"label": s.label, "start": s.start, "end": s.end, "text": s.text} for s in spans], ensure_ascii=False)


def verify_given_spans(text: str, spans_json: str) -> str:
    """Verify and correct spans supplied by an external detector (e.g. fine-tuned CAMeLBERT-MSA); returns the HTML report."""
    try:
        spans = [s for s in json.loads(spans_json) if s["label"] in ("Ayah", "Hadith") and 0 <= s["start"] < s["end"] <= len(text)]
        return ui.render_results(get_pipeline().analyze_spans(text, spans))
    except Exception:
        logger.exception("Verification of given spans failed")
        return ui.render_message("تعذّر التحقق من المقاطع المحدَّدة.", "bad")


def ask_then_verify(prompt: str, provider: str = "openai", model: str = "", api_key: str = "") -> str:
    """Mode B: answer with the pre-configured ChatGPT client (key from OPENAI_API_KEY), then verify every quotation."""
    key = (api_key or "").strip() or os.environ.get(KEY_ENV.get(provider, ""), "")
    try:
        answer = generate(LLMSettings(provider=provider, api_key=key, model=model or ""), prompt)
    except LLMError as exc:
        return ui.render_message(str(exc), "warn")
    return verify_generated_answer(answer)


def build_interface():
    import gradio as gr

    examples = load_examples()

    def next_example(index: int):
        if not examples:
            return "", 0
        return examples[index % len(examples)]["text"], (index + 1) % len(examples)

    def show_default_model(provider: str):
        return PROVIDERS[provider]["model"]

    with gr.Blocks(title="التحقق من هلوسة القرآن والحديث وتصحيحها", css=ui.CSS, theme=gr.themes.Base(primary_hue="emerald", neutral_hue="stone")) as demo:
        gr.HTML(ui.HERO)
        with gr.Tabs():
            with gr.Tab("تحقق مباشر"):
                example_index = gr.State(0)
                text_input = gr.Textbox(label="النص المراد التحقق منه", lines=9, max_lines=24, placeholder=ui.PLACEHOLDER,
                                        rtl=True, elem_classes="input-area")
                with gr.Row():
                    verify_button = gr.Button("تحقّق من النص", variant="primary", scale=3)
                    example_button = gr.Button("جرّب مثالًا", variant="secondary", scale=2)
                results = gr.HTML(elem_classes="results")
                verify_button.click(verify_text, inputs=text_input, outputs=results)
                example_button.click(next_example, inputs=example_index, outputs=[text_input, example_index]).then(
                    verify_text, inputs=text_input, outputs=results)
            with gr.Tab("اسأل ثم تحقّق"):
                gr.HTML('<div class="icv"><div class="notice">اكتب سؤالًا، وسيجيب ChatGPT، ثم يفحص النظام كل آية '
                        'وحديث ورد في إجابته ويعرض الأخطاء والتصحيحات.</div></div>')
                prompt = gr.Textbox(label="سؤالك", lines=3, placeholder=ui.PROMPT_PLACEHOLDER, rtl=True, elem_classes="input-area")
                ask_button = gr.Button("اسأل ثم تحقّق", variant="primary")
                answer_results = gr.HTML(elem_classes="results")
                ask_button.click(ask_then_verify, inputs=prompt, outputs=answer_results)
        gr.HTML(ui.DISCLAIMER)
    return demo


def main() -> None:
    logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")
    get_pipeline()
    warm_in_background()
    build_interface().queue().launch(share=os.environ.get("ICV_SHARE") == "1")


if __name__ == "__main__":
    main()