rasa2's picture
Retry broker requests after HF rate limits
35a864c verified
Raw
History Blame Contribute Delete
57.1 kB
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Granite Switch 4.1 3B Playground β€” pipeline-based multi-adapter chat.
Checkboxes enable adapters that run automatically in the correct pipeline
stage: user-validation -> pre-retrieval -> retrieval -> post-retrieval ->
generation -> post-generation. Results are consolidated in a single message.
"""
import html
import json
import os
import re
import threading
import time
import urllib.error
import urllib.parse
import urllib.request
try:
import spaces
except ImportError:
import types
spaces = types.ModuleType("spaces")
spaces.GPU = lambda f=None, **kw: f if f else (lambda fn: fn)
import gradio as gr
import httpx
INFERENCE_URL = os.environ.get("INFERENCE_URL", "")
MODEL_ID = "ibm-granite/granite-switch-4.1-3b-preview"
MODEL_OPTIONS = {
"granite-4.0-micro": "ibm-granite/granite-4.0-micro",
"granite-4.1-8b": "ibm-granite/granite-4.1-8b",
"granite-4.1-30b": "ibm-granite/granite-4.1-30b",
"qwen2.5-coder-32b": "Qwen/Qwen2.5-Coder-32B-Instruct",
}
USE_EXTERNAL_LLM = os.getenv("USE_EXTERNAL_LLM", "").lower() in {
"1",
"true",
"yes",
"on",
}
EXTERNAL_LLM_MODEL = os.getenv("EXTERNAL_LLM_MODEL", "granite-4.0-micro")
EXTERNAL_LLM_GPUS = int(os.getenv("EXTERNAL_LLM_GPUS", "1"))
EXTERNAL_LLM_BROKER_URL = os.getenv("EXTERNAL_LLM_BROKER_URL") or os.getenv("BROKER_URL")
EXTERNAL_LLM_TIMEOUT_SECONDS = int(os.getenv("EXTERNAL_LLM_TIMEOUT_SECONDS", "600"))
EXTERNAL_LLM_POLL_INTERVAL_SECONDS = float(
os.getenv("EXTERNAL_LLM_POLL_INTERVAL_SECONDS", "10")
)
EXTERNAL_LLM_MAX_NEW_TOKENS = int(os.getenv("EXTERNAL_LLM_MAX_NEW_TOKENS", "2048"))
EXTERNAL_LLM_BROKER_MAX_RETRIES = int(
os.getenv("EXTERNAL_LLM_BROKER_MAX_RETRIES", "5")
)
EXTERNAL_LLM_BROKER_RETRY_SECONDS = float(
os.getenv("EXTERNAL_LLM_BROKER_RETRY_SECONDS", "15")
)
print(
"[INFO] Startup config "
f"use_external_llm={USE_EXTERNAL_LLM} "
f"external_model={EXTERNAL_LLM_MODEL} "
f"broker_url_set={bool(EXTERNAL_LLM_BROKER_URL)} "
f"broker_token_set={bool(os.getenv('BROKER_TOKEN'))} "
f"inference_url_set={bool(INFERENCE_URL)}",
flush=True,
)
SCAFFOLD_COLOR = "#8B5CF6"
BADGE_COLOR = "#6366F1"
WARN_COLOR = "#f59e0b"
JUDGE_SYSTEM = (
"As a judge agent, your role is to help assess whether the provided text "
"meets the given judging criteria, utilizing all available information, "
"including conversations, documents, and tools."
)
GUARDIAN_CRITERIA_BANK = {
"off_scope": (
"The message is NOT about NASA, space exploration, space science, "
"Earth observation, astronomy, planetary missions, rockets, satellites, "
"or related aerospace topics. Messages that are clearly unrelated to "
"these topics meet this criteria."
),
"harm": (
"Texts deemed universally harmful are those that pose harm across all "
"contexts and audiences."
),
"social_bias": (
"Social bias refers to prejudice or discrimination against individuals "
"or groups based on social characteristics such as race, gender, "
"religion, or socioeconomic status."
),
"jailbreak": (
"Jailbreaking: the deliberate evasion of an AI system's built-in "
"safeguards or ethical boundaries, often by crafting specific prompts "
"or scenarios intended to manipulate the AI into producing restricted "
"or inappropriate content."
),
"profanity": "The text contains profane or vulgar language.",
"violence": "The text contains descriptions of or encouragement of violence.",
"unethical_behavior": "The text promotes unethical behavior or actions.",
}
SCENARIOS = {
"User Validation": {
"adapters": {"guardian-core": True},
"settings": {"guardian_criteria": "off_scope", "exit_on_guardian": True},
"queries": [
"How do I hack into a government satellite?",
"What's the best recipe for chocolate cake?",
"Tell me about the Apollo 11 mission",
],
},
"Basic RAG": {
"adapters": {"guardian-core": True, "retrieval": True, "answerability": True},
"settings": {"guardian_criteria": "off_scope", "exit_on_guardian": True, "exit_on_answerability": True},
"queries": [
"When was NASA established?",
"What are NASA's current active missions?",
"How many people work at NASA?",
],
},
"Full Pipeline": {
"adapters": {"guardian-core": True, "query_rewrite": True, "retrieval": True, "answerability": True, "citations": True, "hallucination_detection": True},
"settings": {"guardian_criteria": "off_scope", "exit_on_guardian": True, "exit_on_answerability": True},
"queries": [
"Tell me about the Artemis program and its goals for returning to the moon",
"the discoveries did the Apollo missions make about the moon geology?",
],
},
}
ADAPTER_KEYS = [
"guardian-core", "policy-guardrails", "query_rewrite", "retrieval",
"answerability", "citations", "hallucination_detection",
"factuality-detection", "factuality-correction",
"context-attribution", "uncertainty", "requirement-check",
]
# ── Model loading ─────────────────────────────────────────────────────────
if INFERENCE_URL or USE_EXTERNAL_LLM:
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
_model = None
else:
import torch
import granite_switch.hf # noqa: F401 β€” registers HF backend
from transformers import AutoModelForCausalLM, AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
_model = None
def _get_model():
global _model
if INFERENCE_URL or USE_EXTERNAL_LLM:
return None
if _model is None:
_model = AutoModelForCausalLM.from_pretrained(
MODEL_ID, torch_dtype=torch.bfloat16
)
_model.eval()
_model.to("cuda")
return _model
_context_log = []
def _inference_backend_summary():
if USE_EXTERNAL_LLM:
return (
"query_llm broker "
f"model={EXTERNAL_LLM_MODEL} "
f"broker_url_set={bool(EXTERNAL_LLM_BROKER_URL)} "
f"broker_token_set={bool(os.getenv('BROKER_TOKEN'))}"
)
if INFERENCE_URL:
return f"vLLM/openai-compatible endpoint url={INFERENCE_URL}"
return f"local GPU model={MODEL_ID}"
def validate_query_llm_args(model_name, gpus, user_text, max_new_tokens):
if model_name not in MODEL_OPTIONS:
raise ValueError(f"Model is not allowed: {model_name}")
if gpus < 1 or gpus > 16:
raise ValueError("GPUs must be between 1 and 16")
if user_text is None or not user_text.strip():
raise ValueError("Prompt cannot be empty")
if len(user_text) > 10_000:
raise ValueError("Prompt is too long; max 10,000 characters")
if max_new_tokens < 1 or max_new_tokens > EXTERNAL_LLM_MAX_NEW_TOKENS:
raise ValueError(
"Max new tokens must be between 1 and "
f"{EXTERNAL_LLM_MAX_NEW_TOKENS}"
)
def _broker_request(path, data=None, method="GET"):
if not EXTERNAL_LLM_BROKER_URL:
raise RuntimeError(
"USE_EXTERNAL_LLM is set, but EXTERNAL_LLM_BROKER_URL or BROKER_URL "
"is missing."
)
broker_token = os.getenv("BROKER_TOKEN")
if not broker_token:
raise RuntimeError("USE_EXTERNAL_LLM is set, but BROKER_TOKEN is missing.")
url = f"{EXTERNAL_LLM_BROKER_URL.rstrip('/')}/{path.lstrip('/')}"
encoded_data = None
headers = {"X-Broker-Token": broker_token}
hf_token = os.getenv("HF_TOKEN")
if hf_token:
headers["Authorization"] = f"Bearer {hf_token}"
if data is not None:
encoded_data = urllib.parse.urlencode(data).encode("utf-8")
headers["Content-Type"] = "application/x-www-form-urlencoded"
for attempt in range(EXTERNAL_LLM_BROKER_MAX_RETRIES + 1):
request = urllib.request.Request(
url, data=encoded_data, headers=headers, method=method
)
try:
with urllib.request.urlopen(request, timeout=60) as response:
return json.loads(response.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
detail = exc.read().decode("utf-8", errors="replace")
if exc.code == 429 and attempt < EXTERNAL_LLM_BROKER_MAX_RETRIES:
retry_after = exc.headers.get("Retry-After")
if retry_after:
try:
delay = float(retry_after)
except ValueError:
delay = EXTERNAL_LLM_BROKER_RETRY_SECONDS
else:
delay = EXTERNAL_LLM_BROKER_RETRY_SECONDS * (attempt + 1)
print(
f"[WARN] Broker request {path} hit HTTP 429; "
f"retrying in {delay:.1f}s "
f"({attempt + 1}/{EXTERNAL_LLM_BROKER_MAX_RETRIES})",
flush=True,
)
time.sleep(delay)
continue
raise RuntimeError(
f"Broker request {path} returned HTTP {exc.code}: {detail}"
) from exc
except urllib.error.URLError as exc:
raise RuntimeError(f"Could not connect to broker: {exc.reason}") from exc
raise RuntimeError(f"Broker request {path} failed after retries")
def _extract_broker_result(job):
result = job.get("result")
if isinstance(result, str):
return result.strip()
if isinstance(result, dict):
for key in ("text", "output", "response", "generated_text"):
value = result.get(key)
if isinstance(value, str):
return value.strip()
return json.dumps(result, ensure_ascii=False)
def query_llm(user_text, max_new_tokens=128):
"""Submit a query_llm job to the broker and wait for the worker result."""
max_new_tokens = int(max_new_tokens)
validate_query_llm_args(
EXTERNAL_LLM_MODEL, EXTERNAL_LLM_GPUS, user_text, max_new_tokens
)
job = _broker_request(
"/api/jobs/query-llm",
data={
"model": EXTERNAL_LLM_MODEL,
"gpus": str(EXTERNAL_LLM_GPUS),
"user_text": user_text,
"max_new_tokens": str(max_new_tokens),
},
method="POST",
)
job_id = job["id"]
deadline = time.monotonic() + EXTERNAL_LLM_TIMEOUT_SECONDS
while time.monotonic() < deadline:
job = _broker_request(f"/api/jobs/{job_id}")
status = job.get("status")
if status == "done":
return _extract_broker_result(job)
if status == "failed":
raise RuntimeError(job.get("result") or f"query_llm job {job_id} failed")
time.sleep(EXTERNAL_LLM_POLL_INTERVAL_SECONDS)
raise TimeoutError(
f"Timed out waiting for query_llm job {job_id} after "
f"{EXTERNAL_LLM_TIMEOUT_SECONDS} seconds"
)
def _generate_raw(messages, adapter=None, documents=None, max_new_tokens=128):
"""Generate text. Uses configured remote server, broker, or local GPU."""
kwargs = {}
if adapter:
kwargs["adapter_name"] = adapter
if documents:
kwargs["documents"] = documents
prompt = tokenizer.apply_chat_template(
messages, add_generation_prompt=True, tokenize=False, **kwargs
)
_context_log.append({"adapter": adapter or "base-model", "prompt": prompt})
if USE_EXTERNAL_LLM:
print(
"[DEBUG] using query_llm backend "
f"adapter={adapter or 'base-model'} "
f"model={EXTERNAL_LLM_MODEL} "
f"max_tokens={max_new_tokens}",
flush=True,
)
result = query_llm(prompt, max_new_tokens=max_new_tokens)
print(
f"[DEBUG] query_llm adapter={adapter or 'base-model'} "
f"max_tokens={max_new_tokens} response_len={len(result)}",
flush=True,
)
return result
if INFERENCE_URL:
api_key = os.environ.get("INFERENCE_API_KEY", "unused")
model_name = os.environ.get("VLLM_MODEL_NAME", MODEL_ID)
resp = httpx.post(
f"{INFERENCE_URL}/v1/completions",
headers={"Authorization": f"Bearer {api_key}"},
json={
"model": model_name,
"prompt": prompt,
"max_tokens": max_new_tokens,
"temperature": 0,
},
timeout=120.0,
)
resp.raise_for_status()
result = resp.json()["choices"][0]["text"] or ""
result = result.strip()
print(f"[DEBUG] _generate_raw adapter={adapter} max_tokens={max_new_tokens} response_len={len(result)}")
return result
m = _get_model()
inputs = tokenizer(prompt, return_tensors="pt").to("cuda")
with torch.no_grad():
output_ids = m.generate(
**inputs, max_new_tokens=max_new_tokens, do_sample=False
)
return tokenizer.decode(
output_ids[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True
).strip()
# ── Vector DB manager ─────────────────────────────────────────────────────
class VectorDBManager:
def __init__(self):
self._collection = None
self._status = "not_started"
self._error = None
self._lock = threading.Lock()
@property
def is_ready(self):
return self._status == "ready"
def start_loading(self):
self._status = "loading"
thread = threading.Thread(target=self._load, daemon=True)
thread.start()
def _load(self):
try:
from govt_data_loader import build_nasa_chroma
self._collection = build_nasa_chroma()
with self._lock:
self._status = "ready"
except Exception as e:
with self._lock:
self._status = "error"
self._error = str(e)
def query(self, query_text, n_results=5):
if not self.is_ready:
return None
with self._lock:
results = self._collection.query(
query_texts=[query_text], n_results=n_results
)
docs = []
for i, text in enumerate(results["documents"][0]):
meta = results["metadatas"][0][i] if results["metadatas"] else {}
docs.append({
"text": text,
"title": meta.get("title", ""),
"url": meta.get("url", ""),
})
return docs
db_manager = VectorDBManager()
# ── Display helpers ───────────────────────────────────────────────────────
def _badge(label):
return (
f'<span style="background:{BADGE_COLOR};color:white;padding:2px 8px;'
f'border-radius:4px;font-size:0.8em;">{html.escape(label)}</span>'
)
def _scaffold(text):
return (
f'<span style="color:{SCAFFOLD_COLOR};font-family:monospace;'
f'font-size:0.85em;">{html.escape(text)}</span>'
)
def _parse_docs(documents_text):
if not documents_text or not documents_text.strip():
return None
return [{"text": d.strip()} for d in documents_text.split("\n---\n") if d.strip()]
def _split_sentences(text):
parts = re.split(r"(?<=[.!?])\s+", text.strip())
return [p for p in parts if p]
# ── Adapter build + format functions ──────────────────────────────────────
# Each _build_* returns (messages, adapter_name, documents, max_tokens)
# Each _format_* returns an HTML string
def _build_guardian(user_msg, criteria_choice, custom_criteria):
criteria = custom_criteria if criteria_choice == "Custom" else GUARDIAN_CRITERIA_BANK.get(criteria_choice, criteria_choice)
schema = "If the last user's text meets the criteria, return 'yes'; otherwise, return 'no'."
guardian_block = (
f"<guardian>{JUDGE_SYSTEM}\n\n"
f"### Criteria: {criteria}\n\n"
f"### Scoring Schema: {schema}"
)
messages = [
{"role": "user", "content": user_msg},
{"role": "user", "content": guardian_block},
]
return messages, "guardian-core", None, 20
def _format_guardian(raw, criteria_choice):
try:
result = json.loads(raw)
score = result.get("score", raw)
flagged = score == "yes"
if not flagged:
return None
return (
f'<div style="border-left:3px solid #ef4444;padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Guardian')} <b>FLAGGED</b> for {html.escape(criteria_choice)}</div>"
)
except json.JSONDecodeError:
return f"{_badge('Guardian')} <code>{html.escape(raw)}</code>"
def _build_policy(user_msg, policy_text):
policy_system = (
"You are a compliance agent trying to help determine whether a "
"scenario is compliant with a given policy."
)
scoring_schema = (
"Does the scenario described in the previous message comply with the "
"policy? Compliance is determined based on the following guidelines:\n"
'- "Yes" if the scenario complies with certainty\n'
'- "No" if the scenario does not comply with certainty\n'
'- "Ambiguous" if more information is needed\n\n'
'Your answer must be either "Yes", "No", or "Ambiguous". '
'Return as JSON: {"label": "Yes"/"No"/"Ambiguous"}.'
)
policy_block = (
f"<guardian> {policy_system}\n\n"
f"### Criteria: Policy: {policy_text}\n\n"
f"### Scoring Schema: {scoring_schema}"
)
messages = [
{"role": "user", "content": user_msg},
{"role": "user", "content": policy_block},
]
return messages, "policy-guardrails", None, 20
def _format_policy(raw):
try:
result = json.loads(raw)
label = result.get("label", raw)
color = "#22c55e" if label == "Yes" else "#ef4444" if label == "No" else "#f59e0b"
return (
f'<div style="border-left:3px solid {color};padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Policy')} Compliance: <b>{html.escape(str(label))}</b></div>"
)
except json.JSONDecodeError:
return f"{_badge('Policy')} <code>{html.escape(raw)}</code>"
def _build_query_rewrite(user_msg):
messages = [{"role": "user", "content": user_msg}]
return messages, "query_rewrite", None, 64
def _format_query_rewrite(raw, original):
return (
f'<div style="border-left:3px solid #6366F1;padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Query Rewrite')} "
f"{_scaffold(html.escape(original))} &rarr; <b>{html.escape(raw)}</b></div>"
)
def _build_answerability(question, docs_parsed):
messages = [{"role": "user", "content": question}]
return messages, "answerability", docs_parsed, 32
def _is_unanswerable(raw):
raw_lower = raw.lower().strip().strip('"')
return raw_lower in ("unanswerable", "no") or "unanswerable" in raw_lower
def _format_answerability(raw):
if _is_unanswerable(raw):
color = "#ef4444"
elif "yes" in raw.lower() or "answerable" in raw.lower():
color = "#22c55e"
else:
color = "#6366F1"
return (
f'<div style="border-left:3px solid {color};padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Answerability')} {html.escape(raw)}</div>"
)
def _build_citations(question, answer, docs_parsed):
messages = [
{"role": "user", "content": question},
{"role": "assistant", "content": answer},
]
return messages, "citations", docs_parsed, 128
def _format_citations(raw):
return (
f'<div style="border-left:3px solid #6366F1;padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Citations')} {raw}</div>"
)
def _build_hallucination(question, answer, docs_parsed):
messages = [
{"role": "user", "content": question},
{"role": "assistant", "content": answer},
]
return messages, "hallucination_detection", docs_parsed, 256
def _format_hallucination(raw):
flagged = "yes" in raw.lower()
color = "#ef4444" if flagged else "#22c55e"
return (
f'<div style="border-left:3px solid {color};padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Hallucination')} {html.escape(raw)}</div>"
)
def _build_factuality_det(answer, docs_parsed):
factuality_criteria = (
"A factually incorrect response occurs when the assistant's message "
"contains one or more factual claims that are unsupported by, "
"inconsistent with, or directly contradicted by the information "
"provided in the documents or context."
)
schema = "If the last assistant's text meets the criteria, return 'yes'; otherwise, return 'no'."
guardian_block = (
f"<guardian>{JUDGE_SYSTEM}\n\n"
f"### Criteria: {factuality_criteria}\n\n"
f"### Scoring Schema: {schema}"
)
messages = [
{"role": "assistant", "content": answer},
{"role": "user", "content": guardian_block},
]
return messages, "factuality-detection", docs_parsed, 20
def _format_factuality_det(raw):
try:
result = json.loads(raw)
score = result.get("score", raw)
has_errors = score == "yes"
label = "Errors found" if has_errors else "No errors"
color = "#ef4444" if has_errors else "#22c55e"
return (
f'<div style="border-left:3px solid {color};padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Factuality Detection')} <b>{label}</b> (<code>{score}</code>)</div>"
)
except json.JSONDecodeError:
return f"{_badge('Factuality Detection')} <code>{html.escape(raw)}</code>"
def _build_factuality_cor(answer, docs_parsed):
factuality_criteria = (
"A factually incorrect response occurs when the assistant's message "
"contains one or more factual claims that are unsupported by, "
"inconsistent with, or directly contradicted by the information "
"provided in the documents or context."
)
schema = (
"If the last assistant's text meets the criteria, return a corrected "
"version of the assistant's message based on the given context; "
"otherwise, return 'none'."
)
guardian_block = (
f"<guardian>{JUDGE_SYSTEM}\n\n"
f"### Criteria: {factuality_criteria}\n\n"
f"### Scoring Schema: {schema}"
)
messages = [
{"role": "assistant", "content": answer},
{"role": "user", "content": guardian_block},
]
return messages, "factuality-correction", docs_parsed, 256
def _format_factuality_cor(raw, original_answer):
try:
result = json.loads(raw)
correction = result.get("correction", raw)
if correction == "none":
return (
f'<div style="border-left:3px solid #22c55e;padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Factuality Correction')} No correction needed.</div>"
)
return (
f'<div style="border-left:3px solid #f59e0b;padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Factuality Correction')}"
f'<details><summary style="cursor:pointer;">Show original vs corrected</summary>'
f'<div style="padding:8px;font-size:0.9em;">'
f'<b>Original:</b><br>{html.escape(original_answer)}<br><br>'
f'<b>Corrected:</b><br>{html.escape(correction)}'
f'</div></details></div>'
)
except json.JSONDecodeError:
return f"{_badge('Factuality Correction')} <code>{html.escape(raw)}</code>"
def _build_context_attr(question, answer, docs_raw_texts):
c_counter = 0
tagged_doc_parts = []
for doc_text in docs_raw_texts:
parts = []
for sent in _split_sentences(doc_text):
parts.append(f"<c{c_counter}> {sent}")
c_counter += 1
tagged_doc_parts.append({"text": " ".join(parts)})
response_sents = _split_sentences(answer)
tagged_response = " ".join(f"<r{i}> {s}" for i, s in enumerate(response_sents))
instruction = (
"You provided the last assistant response above based on context, which may "
"include documents and/or previous conversation turns. Your response is "
"divided into sentences, numbered in the format <r0> sentence 0 <r1> "
"sentence 1 ... Sentences in the context are also numbered: <c0> sentence 0 "
"<c1> sentence 1 ... For each response sentence, please list the context "
"sentences that were most important for you to generate the response "
"sentence. Provide your answer in JSON format, as an array of JSON objects, "
'where each object has two members: "r" with the response sentence number '
'as the value, and "c" with an array of context sentence numbers as the '
"value. List the context sentences in order from most important to least "
"important. Ensure that you include an object for each response sentence, "
"even if the corresponding array of context sentence numbers is empty. "
"Answer with only the JSON and do not explain.\n"
)
messages = [
{"role": "user", "content": question},
{"role": "assistant", "content": tagged_response},
{"role": "user", "content": instruction},
]
return messages, "context-attribution", tagged_doc_parts, 256
def _format_context_attr(raw):
return (
f'<div style="border-left:3px solid #6366F1;padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Context Attribution')} <pre style=\"margin:4px 0;font-size:0.85em;\">{html.escape(raw)}</pre></div>"
)
def _build_uncertainty(conversation_text):
messages = [
{"role": "user", "content": conversation_text},
{"role": "user", "content": "<certainty>"},
]
return messages, "uncertainty", None, 20
def _format_uncertainty(raw):
try:
result = json.loads(raw)
digit = int(result.get("score", 0))
prob = 0.1 * digit + 0.05
color = "#22c55e" if digit >= 7 else "#f59e0b" if digit >= 4 else "#ef4444"
return (
f'<div style="border-left:3px solid {color};padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Uncertainty')} Certainty: <b>{digit}</b> "
f"(~{prob*100:.0f}% confidence)</div>"
)
except (json.JSONDecodeError, ValueError):
return f"{_badge('Uncertainty')} <code>{html.escape(raw)}</code>"
def _build_requirement(question, answer, requirements):
evaluation_prompt = (
"Please verify if the assistant's generation satisfies the user's "
"requirements or not and reply with a binary label accordingly. "
'Respond with a json {"score": "yes"} if the constraints are satisfied '
'or respond with {"score": "no"} if the constraints are not satisfied.'
)
req_turn = f"<requirements> {requirements}\n{evaluation_prompt}"
messages = [
{"role": "user", "content": question},
{"role": "assistant", "content": answer},
{"role": "user", "content": req_turn},
]
return messages, "requirement-check", None, 20
def _format_requirement(raw):
try:
result = json.loads(raw)
score = result.get("score", raw)
satisfied = score == "yes"
label = "Satisfied" if satisfied else "Not satisfied"
color = "#22c55e" if satisfied else "#ef4444"
return (
f'<div style="border-left:3px solid {color};padding:4px 8px;margin:2px 0;background:#f9fafb;color:#1f2937;border-radius:4px;">'
f"{_badge('Requirement Check')} <b>{label}</b> (<code>{score}</code>)</div>"
)
except json.JSONDecodeError:
return f"{_badge('Requirement Check')} <code>{html.escape(raw)}</code>"
# ── Pipeline output assembly ──────────────────────────────────────────────
def _render_context_viewer(context_log):
"""Render context log as collapsible HTML JSON viewer."""
if not context_log:
return '<div style="color:#6b7280;font-size:0.9em;">No context generated.</div>'
items = []
for i, entry in enumerate(context_log):
adapter = html.escape(entry["adapter"])
prompt = html.escape(entry["prompt"])
items.append(
f'<details style="margin:4px 0;border:1px solid #e5e7eb;border-radius:4px;">'
f'<summary style="cursor:pointer;padding:6px 10px;background:#f3f4f6;'
f'color:#1f2937;font-weight:600;font-size:0.85em;">'
f'[{i}] {adapter}</summary>'
f'<pre style="padding:8px 10px;margin:0;font-size:0.8em;'
f'white-space:pre-wrap;word-break:break-word;background:#fafafa;'
f'color:#374151;max-height:400px;overflow-y:auto;">{prompt}</pre>'
f'</details>'
)
return (
f'<div style="font-family:monospace;font-size:0.85em;">'
f'<div style="color:#6b7280;margin-bottom:6px;">'
f'{len(context_log)} adapter call(s) β€” click to expand</div>'
+ "".join(items) + '</div>'
)
def _format_retrieved_docs_badge(docs):
return f"{_badge('Retrieval')} Retrieved **{len(docs)}** documents (see Context panel)"
def _render_retrieved_docs_viewer(docs):
"""Render retrieved docs as collapsible HTML viewer (for context panel)."""
if not docs:
return ""
items = []
for i, doc in enumerate(docs):
title = html.escape(doc.get("title", f"Document {i+1}"))
text = html.escape(doc["text"][:600])
items.append(
f'<details style="margin:4px 0;border:1px solid #e5e7eb;border-radius:4px;">'
f'<summary style="cursor:pointer;padding:6px 10px;background:#f3f4f6;'
f'color:#1f2937;font-weight:600;font-size:0.85em;">'
f'[{i}] {title}</summary>'
f'<pre style="padding:8px 10px;margin:0;font-size:0.8em;'
f'white-space:pre-wrap;word-break:break-word;background:#fafafa;'
f'color:#374151;max-height:300px;overflow-y:auto;">{text}</pre>'
f'</details>'
)
return (
f'<div style="font-family:monospace;font-size:0.85em;margin-top:12px;'
f'border-top:1px solid #e5e7eb;padding-top:8px;">'
f'<div style="color:#6b7280;margin-bottom:6px;font-weight:600;">'
f'Retrieved Documents ({len(docs)})</div>'
+ "".join(items) + '</div>'
)
def _assemble_output(answer, sections):
parts = []
validations = [s for stage, s in sections if stage == "validation"]
if validations:
parts.append("".join(validations))
for stage, content in sections:
if stage == "pre_retrieval":
parts.append(content)
for stage, content in sections:
if stage == "retrieval":
parts.append(content)
for stage, content in sections:
if stage == "post_retrieval":
parts.append(content)
if answer:
parts.append(f"\n\n{answer}\n\n")
else:
blocked_msgs = [s for stage, s in sections if stage == "blocked"]
reason = blocked_msgs[0] if blocked_msgs else "Pipeline halted by adapter"
parts.append(
f'<div style="border-left:3px solid #ef4444;padding:8px 12px;'
f'background:#fef2f2;color:#1f2937;border-radius:4px;margin:8px 0;">'
f'<b>Generation skipped:</b> {html.escape(reason)}</div>'
)
post_gen = [s for stage, s in sections if stage == "post_generation"]
if post_gen:
inner = "".join(post_gen)
parts.append(
f'<details open><summary style="cursor:pointer;color:{BADGE_COLOR};'
f'font-weight:600;">Adapter Analysis</summary>'
f'<div style="padding:4px 0;">{inner}</div></details>'
)
skipped = [s for stage, s in sections if stage == "skipped"]
if skipped:
parts.append(
f'<div style="font-size:0.8em;color:#9ca3af;margin-top:4px;">'
+ " | ".join(skipped) + '</div>'
)
return "\n".join(parts)
# ── Pipeline orchestrator ─────────────────────────────────────────────────
def run_pipeline(
user_message, history, enabled_adapters, adapter_config, max_tokens
):
_get_model()
_context_log.clear()
sections = []
docs_parsed = None
docs_raw_texts = None
retrieved_docs_full = None
blocked = False
block_reason = ""
# --- Stage 1: User Validation ---
if "guardian-core" in enabled_adapters:
msgs, adapter, docs, mt = _build_guardian(
user_message,
adapter_config.get("guardian_criteria", "harm"),
adapter_config.get("guardian_custom_criteria", ""),
)
raw = _generate_raw(msgs, adapter, docs, mt)
guardian_result = _format_guardian(raw, adapter_config.get("guardian_criteria", "harm"))
if guardian_result:
sections.append(("validation", guardian_result))
if adapter_config.get("exit_on_guardian", True):
blocked = True
block_reason = "Guardian flagged this message"
if "policy-guardrails" in enabled_adapters and not blocked:
policy_text = adapter_config.get("policy_text", "") or ""
if policy_text.strip():
msgs, adapter, docs, mt = _build_policy(user_message, policy_text)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("validation", _format_policy(raw)))
try:
policy_label = json.loads(raw).get("label", "")
if policy_label == "No" and adapter_config.get("exit_on_policy", True):
blocked = True
block_reason = "Policy non-compliance detected"
except (json.JSONDecodeError, AttributeError):
pass
# --- Stage 2: Pre-Retrieval ---
search_query = user_message
if "query_rewrite" in enabled_adapters and not blocked:
msgs, adapter, docs, mt = _build_query_rewrite(user_message)
raw = _generate_raw(msgs, adapter, docs, mt)
try:
parsed = json.loads(raw)
search_query = parsed.get("rewritten_question", parsed.get("question", raw))
except (json.JSONDecodeError, AttributeError):
search_query = raw
sections.append(("pre_retrieval", _format_query_rewrite(search_query, user_message)))
# --- Stage 3: Retrieval ---
if not blocked:
retrieval_enabled = "retrieval" in enabled_adapters
if retrieval_enabled and db_manager.is_ready:
db_results = db_manager.query(search_query, n_results=5)
if db_results:
docs_parsed = [{"text": d["text"]} for d in db_results]
docs_raw_texts = [d["text"] for d in db_results]
retrieved_docs_full = db_results
sections.append(("retrieval", _format_retrieved_docs_badge(db_results)))
elif retrieval_enabled and not db_manager.is_ready:
sections.append(("skipped", "Retrieval skipped (Vector DB not ready)"))
if docs_parsed is None:
fallback_docs_text = adapter_config.get("context_documents", "") or ""
fallback_parsed = _parse_docs(fallback_docs_text)
if fallback_parsed:
docs_parsed = fallback_parsed
docs_raw_texts = [d["text"] for d in fallback_parsed]
# --- Stage 4: Post-Retrieval ---
if "answerability" in enabled_adapters and not blocked:
if docs_parsed:
msgs, adapter, docs, mt = _build_answerability(user_message, docs_parsed)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_retrieval", _format_answerability(raw)))
if _is_unanswerable(raw) and adapter_config.get("exit_on_answerability", True):
blocked = True
block_reason = "Question not answerable from available documents"
else:
sections.append(("skipped", "Answerability skipped (no documents)"))
# --- Stage 5: Generation ---
if blocked:
answer = None
sections.append(("blocked", block_reason))
else:
model_messages = [
{"role": m["role"], "content": m["content"]}
for m in history
if m["role"] in ("user", "assistant") and "metadata" not in m
]
model_messages.append({"role": "user", "content": user_message})
answer = _generate_raw(model_messages, adapter=None, documents=docs_parsed, max_new_tokens=max_tokens)
# --- Stage 6: Post-Generation (skipped if blocked) ---
if not blocked:
if "citations" in enabled_adapters:
if docs_parsed:
msgs, adapter, docs, mt = _build_citations(user_message, answer, docs_parsed)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_generation", _format_citations(raw)))
else:
sections.append(("skipped", "Citations skipped (no documents)"))
if "hallucination_detection" in enabled_adapters:
if docs_parsed:
msgs, adapter, docs, mt = _build_hallucination(user_message, answer, docs_parsed)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_generation", _format_hallucination(raw)))
else:
sections.append(("skipped", "Hallucination Detection skipped (no documents)"))
if "factuality-detection" in enabled_adapters:
if docs_parsed:
msgs, adapter, docs, mt = _build_factuality_det(answer, docs_parsed)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_generation", _format_factuality_det(raw)))
else:
sections.append(("skipped", "Factuality Detection skipped (no documents)"))
if "factuality-correction" in enabled_adapters:
if docs_parsed:
msgs, adapter, docs, mt = _build_factuality_cor(answer, docs_parsed)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_generation", _format_factuality_cor(raw, answer)))
else:
sections.append(("skipped", "Factuality Correction skipped (no documents)"))
if "context-attribution" in enabled_adapters:
if docs_raw_texts:
msgs, adapter, docs, mt = _build_context_attr(user_message, answer, docs_raw_texts)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_generation", _format_context_attr(raw)))
else:
sections.append(("skipped", "Context Attribution skipped (no documents)"))
if "uncertainty" in enabled_adapters:
conv_text = f"User: {user_message}\nAssistant: {answer}"
msgs, adapter, docs, mt = _build_uncertainty(conv_text)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_generation", _format_uncertainty(raw)))
if "requirement-check" in enabled_adapters:
req_text = adapter_config.get("requirements_text", "") or ""
if req_text.strip():
msgs, adapter, docs, mt = _build_requirement(user_message, answer, req_text)
raw = _generate_raw(msgs, adapter, docs, mt)
sections.append(("post_generation", _format_requirement(raw)))
# --- Assemble ---
assistant_html = _assemble_output(answer, sections)
new_history = list(history) + [
{"role": "user", "content": user_message},
{"role": "assistant", "content": assistant_html, "metadata": {"pipeline": True}},
]
context_html = _render_context_viewer(_context_log)
docs_html = _render_retrieved_docs_viewer(retrieved_docs_full) if retrieved_docs_full else ""
return new_history, new_history, "", docs_html, context_html
# ── Gradio UI ─────────────────────────────────────────────────────────────
CSS = """
#db-status { font-size: 0.85em; padding: 4px 0; }
.compact-cb label { font-size: 0.9em !important; }
details summary { cursor: pointer; font-weight: 600; }
@keyframes pulse { 0%,100% { opacity:1; } 50% { opacity:0.4; } }
#main-chatbot .message { text-align: left !important; }
#main-chatbot .user, #main-chatbot .bot { justify-content: flex-start !important; }
#main-chatbot .message-row { justify-content: flex-start !important; }
"""
def get_db_status():
s = db_manager._status
if s == "ready":
count = db_manager._collection.count() if db_manager._collection else "?"
return (
f'<div style="display:flex;align-items:center;gap:6px;">'
f'<span style="display:inline-block;width:10px;height:10px;'
f'border-radius:50%;background:#22c55e;flex-shrink:0;"></span>'
f'<span>Vector DB: Ready ({count:,} docs)</span></div>'
)
if s == "loading":
return (
f'<div style="display:flex;align-items:center;gap:6px;">'
f'<span style="display:inline-block;width:10px;height:10px;'
f'border-radius:50%;background:#f59e0b;flex-shrink:0;'
f'animation:pulse 1.5s ease-in-out infinite;"></span>'
f'<span>Vector DB: Loading...</span></div>'
)
if s == "error":
return (
f'<div style="display:flex;align-items:center;gap:6px;">'
f'<span style="display:inline-block;width:10px;height:10px;'
f'border-radius:50%;background:#ef4444;flex-shrink:0;"></span>'
f'<span>Vector DB: Error</span></div>'
)
return (
f'<div style="display:flex;align-items:center;gap:6px;">'
f'<span style="display:inline-block;width:10px;height:10px;'
f'border-radius:50%;background:#9ca3af;flex-shrink:0;"></span>'
f'<span>Vector DB: Not started</span></div>'
)
@spaces.GPU
def handle_submit(
message, history,
cb_guardian, cb_policy, cb_qr, cb_retrieval, cb_answerability,
cb_citations, cb_hallucination, cb_fact_det, cb_fact_cor,
cb_context_attr, cb_uncertainty, cb_requirement,
cfg_guardian_criteria, cfg_guardian_custom, cfg_policy_text,
cfg_requirements, cfg_context_docs,
cfg_exit_guardian, cfg_exit_policy, cfg_exit_answerability,
max_tokens,
):
if not message:
return history, history, "", "", ""
message = message.strip()
if not message:
return history, history, "", "", ""
enabled = []
for val, key in [
(cb_guardian, "guardian-core"),
(cb_policy, "policy-guardrails"),
(cb_qr, "query_rewrite"),
(cb_retrieval, "retrieval"),
(cb_answerability, "answerability"),
(cb_citations, "citations"),
(cb_hallucination, "hallucination_detection"),
(cb_fact_det, "factuality-detection"),
(cb_fact_cor, "factuality-correction"),
(cb_context_attr, "context-attribution"),
(cb_uncertainty, "uncertainty"),
(cb_requirement, "requirement-check"),
]:
if val:
enabled.append(key)
config = {
"guardian_criteria": cfg_guardian_criteria,
"guardian_custom_criteria": cfg_guardian_custom,
"policy_text": cfg_policy_text,
"requirements_text": cfg_requirements,
"context_documents": cfg_context_docs,
"exit_on_guardian": cfg_exit_guardian,
"exit_on_policy": cfg_exit_policy,
"exit_on_answerability": cfg_exit_answerability,
}
try:
return run_pipeline(message, history, enabled, config, max_tokens)
except Exception as e:
err_msg = str(e)
if "CUDA" in err_msg or "GPU" in err_msg or "No GPU" in err_msg:
error_html = (
f'<div style="border-left:3px solid #ef4444;padding:8px 12px;'
f'background:#fef2f2;color:#1f2937;border-radius:4px;">'
f'<b>GPU temporarily unavailable</b> β€” ZeroGPU could not allocate '
f'a GPU for this request. Please try again in a few seconds.</div>'
)
else:
error_html = (
f'<div style="border-left:3px solid #ef4444;padding:8px 12px;'
f'background:#fef2f2;color:#1f2937;border-radius:4px;">'
f'<b>Error:</b> {html.escape(err_msg)}</div>'
)
new_history = list(history) + [
{"role": "user", "content": message},
{"role": "assistant", "content": error_html},
]
return new_history, new_history, "", "", f"Error: {err_msg}"
@spaces.GPU
def run_scenario(scenario_name):
scenario = SCENARIOS[scenario_name]
_get_model()
if "retrieval" in scenario["adapters"]:
import time as _time
deadline = _time.time() + 60
while not db_manager.is_ready and _time.time() < deadline:
_time.sleep(1)
cb_values = [k in scenario["adapters"] for k in ADAPTER_KEYS]
config = {
"guardian_criteria": scenario["settings"].get("guardian_criteria", "off_scope"),
"guardian_custom_criteria": "",
"policy_text": "",
"requirements_text": "",
"context_documents": "",
"exit_on_guardian": scenario["settings"].get("exit_on_guardian", True),
"exit_on_policy": scenario["settings"].get("exit_on_policy", True),
"exit_on_answerability": scenario["settings"].get("exit_on_answerability", True),
}
history = []
enabled = list(scenario["adapters"].keys())
docs_html = ""
ctx_html = ""
for query in scenario["queries"]:
history, _, _, docs_html, ctx_html = run_pipeline(query, history, enabled, config, 128)
return (
*cb_values,
scenario["settings"].get("guardian_criteria", "off_scope"),
scenario["settings"].get("exit_on_guardian", True),
scenario["settings"].get("exit_on_policy", True),
scenario["settings"].get("exit_on_answerability", True),
history, history, "", docs_html, ctx_html,
)
with gr.Blocks(title="Granite Switch 4.1 3B Playground") as demo:
gr.Markdown(
"# Granite Switch 4.1 3B Playground\n\n"
"[Granite Switch](https://github.com/generative-computing/granite-switch) "
"embeds multiple LoRA adapters inside a single Granite checkpoint and "
"activates them on demand via control tokens. This playground runs "
"**Granite Switch 4.1 3B** with 11 adapters organized in a "
"RAG pipeline:\n\n"
"1. **User Validation** β€” Guardian & Policy Guardrails screen the input\n"
"2. **Pre-Retrieval** β€” Query Rewrite optimizes the search query\n"
"3. **Retrieval** β€” Vector DB searches ~2k NASA passages\n"
"4. **Post-Retrieval** β€” Answerability checks whether the docs can answer the question\n"
"5. **Generation** β€” Base model produces an answer grounded in retrieved context\n"
"6. **Post-Generation** β€” Citations, Hallucination Detection, Factuality, "
"Context Attribution, Uncertainty, and Requirement Check analyze the response\n\n"
"Enable adapters with the checkboxes on the left, then ask a question about "
"NASA missions, Earth observation, or space science."
)
db_status_html = gr.HTML("", elem_id="db-status")
gr.Markdown("**Demo Scenarios:**")
with gr.Row():
btn_scenario_1 = gr.Button("1: User Validation", variant="secondary", size="sm")
btn_scenario_2 = gr.Button("2: Basic RAG", variant="secondary", size="sm")
btn_scenario_3 = gr.Button("3: Full Pipeline", variant="secondary", size="sm")
chat_state = gr.State([])
with gr.Row(equal_height=True):
# ── Sidebar ───────────────────────────────────────────────
with gr.Column(scale=1, min_width=220):
gr.Markdown("**Pipeline Adapters**")
with gr.Accordion("User Validation", open=True):
cb_guardian = gr.Checkbox(label="Guardian", value=True, elem_classes=["compact-cb"])
cb_policy = gr.Checkbox(label="Policy Guardrails", value=False, elem_classes=["compact-cb"])
with gr.Accordion("Pre-Retrieval", open=True):
cb_qr = gr.Checkbox(label="Query Rewrite", value=False, elem_classes=["compact-cb"])
with gr.Accordion("Retrieval", open=True):
cb_retrieval = gr.Checkbox(label="Vector DB Search", value=True, elem_classes=["compact-cb"])
with gr.Accordion("Post-Retrieval", open=True):
cb_answerability = gr.Checkbox(label="Answerability", value=False, elem_classes=["compact-cb"])
with gr.Accordion("Post-Generation", open=True):
cb_citations = gr.Checkbox(label="Citations", value=False, elem_classes=["compact-cb"])
cb_hallucination = gr.Checkbox(label="Hallucination Detection", value=False, elem_classes=["compact-cb"])
cb_fact_det = gr.Checkbox(label="Factuality Detection", value=False, elem_classes=["compact-cb"])
cb_fact_cor = gr.Checkbox(label="Factuality Correction", value=False, elem_classes=["compact-cb"])
cb_context_attr = gr.Checkbox(label="Context Attribution", value=False, elem_classes=["compact-cb"])
cb_uncertainty = gr.Checkbox(label="Uncertainty", value=False, elem_classes=["compact-cb"])
cb_requirement = gr.Checkbox(label="Requirement Check", value=False, elem_classes=["compact-cb"])
with gr.Accordion("Adapter Settings", open=False):
cfg_guardian_criteria = gr.Dropdown(
choices=list(GUARDIAN_CRITERIA_BANK.keys()) + ["Custom"],
value="off_scope", label="Guardian Criteria",
)
cfg_guardian_desc = gr.Textbox(
label="Criteria Description (read-only for presets, editable for Custom)",
value=GUARDIAN_CRITERIA_BANK["off_scope"],
lines=3, interactive=False,
)
cfg_guardian_custom = gr.Textbox(
label="Custom Criteria (used when 'Custom' is selected)",
lines=2, visible=False,
)
cfg_policy_text = gr.Textbox(
label="Policy Text", lines=2,
placeholder="e.g., No investment advice.",
)
cfg_requirements = gr.Textbox(
label="Requirements", lines=2,
placeholder="e.g., Formal tone, under 100 words.",
)
cfg_context_docs = gr.Textbox(
label="Context Documents (fallback when retrieval is off)",
lines=4,
placeholder="Paste documents separated by ---",
)
gr.Markdown("**Exit on warning**")
cfg_exit_guardian = gr.Checkbox(
label="Guardian blocks generation", value=True, elem_classes=["compact-cb"]
)
cfg_exit_policy = gr.Checkbox(
label="Policy blocks generation", value=True, elem_classes=["compact-cb"]
)
cfg_exit_answerability = gr.Checkbox(
label="Answerability blocks generation", value=True, elem_classes=["compact-cb"]
)
max_tokens_slider = gr.Slider(16, 512, value=128, step=16, label="Max tokens")
# ── Main chat area ────────────────────────────────────────
with gr.Column(scale=3):
chatbot = gr.Chatbot(sanitize_html=False, height=520, elem_id="main-chatbot")
with gr.Row():
msg_input = gr.Textbox(
show_label=False, lines=1, scale=4,
placeholder="Type a message... (Enter to send, Shift+Enter for newline)",
)
with gr.Column(scale=1, min_width=80):
send_btn = gr.Button("Send", variant="primary")
clear_btn = gr.Button("Clear", variant="secondary")
# ── Context panel ─────────────────────────────────────────
with gr.Column(scale=2):
with gr.Accordion("Retrieved Documents", open=True):
docs_display = gr.HTML(
value='<div style="color:#6b7280;font-size:0.9em;">Documents will appear here after retrieval...</div>',
)
with gr.Accordion("Full Context (prompts)", open=False):
context_display = gr.HTML(
value='<div style="color:#6b7280;font-size:0.9em;">Context will appear here after sending a message...</div>',
)
all_inputs = [
msg_input, chat_state,
cb_guardian, cb_policy, cb_qr, cb_retrieval, cb_answerability,
cb_citations, cb_hallucination, cb_fact_det, cb_fact_cor,
cb_context_attr, cb_uncertainty, cb_requirement,
cfg_guardian_criteria, cfg_guardian_custom, cfg_policy_text,
cfg_requirements, cfg_context_docs,
cfg_exit_guardian, cfg_exit_policy, cfg_exit_answerability,
max_tokens_slider,
]
all_outputs = [chatbot, chat_state, msg_input, docs_display, context_display]
def _update_guardian_desc(choice):
if choice == "Custom":
return gr.update(visible=False), gr.update(visible=True)
text = GUARDIAN_CRITERIA_BANK.get(choice, "")
return gr.update(value=text, visible=True), gr.update(visible=False)
cfg_guardian_criteria.change(
_update_guardian_desc,
inputs=[cfg_guardian_criteria],
outputs=[cfg_guardian_desc, cfg_guardian_custom],
)
send_btn.click(handle_submit, inputs=all_inputs, outputs=all_outputs)
msg_input.submit(handle_submit, inputs=all_inputs, outputs=all_outputs)
clear_btn.click(
lambda: (
[], [],
'<div style="color:#6b7280;font-size:0.9em;">Documents will appear here after retrieval...</div>',
'<div style="color:#6b7280;font-size:0.9em;">Context will appear here after sending a message...</div>',
),
outputs=[chatbot, chat_state, docs_display, context_display],
)
scenario_outputs = [
cb_guardian, cb_policy, cb_qr, cb_retrieval, cb_answerability,
cb_citations, cb_hallucination, cb_fact_det, cb_fact_cor,
cb_context_attr, cb_uncertainty, cb_requirement,
cfg_guardian_criteria,
cfg_exit_guardian, cfg_exit_policy, cfg_exit_answerability,
chatbot, chat_state, msg_input, docs_display, context_display,
]
btn_scenario_1.click(lambda: run_scenario("User Validation"), outputs=scenario_outputs)
btn_scenario_2.click(lambda: run_scenario("Basic RAG"), outputs=scenario_outputs)
btn_scenario_3.click(lambda: run_scenario("Full Pipeline"), outputs=scenario_outputs)
demo.load(get_db_status, outputs=db_status_html)
timer = gr.Timer(5)
timer.tick(get_db_status, outputs=db_status_html)
# ── Startup ───────────────────────────────────────────────────────────────
db_manager.start_loading()
if __name__ == "__main__":
print(f"[INFO] Inference backend: {_inference_backend_summary()}", flush=True)
demo.launch(
server_name="0.0.0.0",
server_port=int(os.getenv("PORT", "7860")),
css=CSS,
ssr_mode=False,
)