byteastra / test_10_qna_hf.py
risu1012's picture
feat: optimize deployment payload using zipped database index
9fca47f
Raw History Blame Contribute Delete
11.6 kB
import asyncio
import time
import json
import urllib.request
import ssl
import sys
sys.stdout.reconfigure(encoding='utf-8', write_through=True)
# Wrap print to flush immediately
import builtins
_print = builtins.print
def print(*args, **kwargs):
kwargs['flush'] = True
_print(*args, **kwargs)
builtins.print = print
QUESTIONS = [
# ── EASY (1-5) ──────────────────────────────────────────────────────────
("easy", 1, "What is Ayurveda?"),
("easy", 2, "What are the three Doshas?"),
("easy", 3, "What is Vata dosha?"),
("easy", 4, "What are the seven Dhatus?"),
("easy", 5, "What is Agni in Ayurveda?"),
# ── MEDIUM (6-10) ───────────────────────────────────────────────────────
("medium", 6, "List the five subtypes of Vata with their locations and functions."),
("medium", 7, "List the five subtypes of Pitta with their locations and functions."),
("medium", 8, "List the five subtypes of Kapha with their locations and functions."),
("medium", 9, "Explain the 13 types of Agni with their roles in digestion."),
("medium", 10, "What is Basti therapy and what dosha does it primarily treat?"),
# ── HARD (11-15) ────────────────────────────────────────────────────────
("hard", 11, "Explain the Dhatu-Agni pathway: how Rasa Dhatu transforms into Shukra Dhatu step by step."),
("hard", 12, "Describe Samprapti (pathogenesis) with all six stages of disease development."),
("hard", 13, "Explain the role of Ojas in immunity and its relationship to all seven Dhatus."),
("hard", 14, "Describe the Ashtavidha Pariksha (eightfold examination) in detail."),
("hard", 15, "What are the contraindications of Vamana therapy?"),
# ── VERY TOUGH (16-20) ──────────────────────────────────────────────────
("very_tough", 16, "Explain the concept of Tridosha Siddhanta and how the interaction of three Doshas maintains homeostasis."),
("very_tough", 17, "Describe the complete pathogenesis (Samprapti) of Jwara (fever) including all stages."),
("very_tough", 18, "Describe the Srotas, their Mula (root), Marga (path) and Mukha (opening) for all 13 Srotamsi."),
("very_tough", 19, "Explain the theory of Loka-Purusha Samya (macrocosm-microcosm relationship)."),
("very_tough", 20, "What is Swasthavritta? Describe all its components including Dinacharya, Ritucharya and Sadvritta."),
]
BASE_URL = "https://jarnil-byteastra.hf.space"
RESULTS = []
# SSL Context to bypass cert verification for local curl/urllib if needed
ctx = ssl.create_default_context()
ctx.check_hostname = False
ctx.verify_mode = ssl.CERT_NONE
def create_session() -> str:
print(f"Creating test session on HF Space...")
url = f"{BASE_URL}/sessions"
data = json.dumps({
"domain": "ayurveda",
"title": "HF QNA Stress Test Session",
"user_id": "hf_tester"
}).encode("utf-8")
req = urllib.request.Request(url, data=data, headers={"Content-Type": "application/json"}, method="POST")
with urllib.request.urlopen(req, context=ctx) as response:
res = json.loads(response.read().decode("utf-8"))
print(f"Session created: {res['id']}")
return res["id"]
async def run_one(session_id, level, qnum, query, log_file):
url = f"{BASE_URL}/chat"
payload = {
"session_id": session_id,
"message": query,
"stream": True
}
start = time.time()
first_tok = None
words = 0
full = []
grounded = False
citations = 0
error = None
try:
# We run the request inside a thread pool because urllib is blocking,
# or we use standard async fetch if possible.
# To keep it standard library without dependencies (like httpx or aiohttp),
# we do a synchronous streaming read in a thread.
loop = asyncio.get_event_loop()
def perform_request():
nonlocal first_tok, words, grounded, citations, error
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(url, data=data, headers={"Content-Type": "application/json"}, method="POST")
try:
with urllib.request.urlopen(req, context=ctx, timeout=60.0) as response:
# Read line by line (SSE format)
for line in response:
line_str = line.decode("utf-8").strip()
if line_str.startswith("data:"):
evt_data = json.loads(line_str[5:])
evt_type = evt_data.get("type")
if evt_type == "citations":
grounded = evt_data.get("is_grounded", False)
citations = len(evt_data.get("citations", []))
elif evt_type == "delta":
c = evt_data.get("content", "")
if c:
if first_tok is None:
first_tok = time.time() - start
words += len(c.split())
full.append(c)
elif evt_type == "error":
error = evt_data.get("error")
except Exception as e:
error = str(e)
await loop.run_in_executor(None, perform_request)
except Exception as e:
error = str(e)
total = time.time() - start
response = "".join(full).strip()
# Quality heuristics
has_heading = "##" in response or "###" in response
has_list = "- " in response or "1." in response
has_sanskrit = any(w in response for w in ["Agni","Dosha","Dhatu","Vata","Pitta","Kapha","Ama","Ojas"])
score_hint = "GOOD" if (words > 80 and has_heading and has_list and has_sanskrit) else ("PARTIAL" if words > 40 else "SHORT")
ttft_str = f"{first_tok:.2f}s" if first_tok else "N/A"
result = {
"level": level,
"q": qnum,
"query": query,
"grounded": grounded,
"citations": citations,
"ttft": round(first_tok, 2) if first_tok else None,
"total_s": round(total, 2),
"words": words,
"quality": score_hint,
"error": error,
"response": response,
"response_preview": response[:200] + "..." if len(response) > 200 else response,
}
RESULTS.append(result)
print(f"[Q{qnum:02d} | {level.upper():<10}] TTFT: {ttft_str:<6} | Total: {total:.1f}s | Words: {words:<4} | Quality: {score_hint:<7} | Grounded: {str(grounded):<5} | Citations: {citations:<2} | Query: {query}")
# Log to file
log_file.write(f"======================================================================\n")
log_file.write(f"QUESTION {qnum:02d} | Level: {level.upper()} | Citations: {citations} | Grounded: {grounded}\n")
log_file.write(f"TTFT: {ttft_str} | Total Time: {total:.2f}s | Word Count: {words} | Quality: {score_hint}\n")
log_file.write(f"Query: {query}\n")
log_file.write(f"----------------------------------------------------------------------\n")
if error:
log_file.write(f"ERROR: {error}\n")
else:
log_file.write(f"{response}\n")
log_file.write(f"======================================================================\n\n")
log_file.flush()
return result
async def main():
print("="*85)
print("ByteAstra 20-Question BAMS Stress Test (Hugging Face Space API Endpoint)")
print("="*85)
try:
session_id = create_session()
except Exception as e:
print(f"Failed to start session on Hugging Face Space: {e}")
sys.exit(1)
log_path = r"c:\Risu Solutions\ByteAstra\backend\stress_qna_log_hf.txt"
with open(log_path, "w", encoding="utf-8") as log_file:
log_file.write("======================================================================\n")
log_file.write("BYTEASTRA BAMS 20-QUESTION STRESS TEST QNA LOG (HF SPACE)\n")
log_file.write(f"Executed At: {time.strftime('%Y-%m-%d %H:%M:%S')}\n")
log_file.write("======================================================================\n\n")
for level, qnum, query in QUESTIONS:
await run_one(session_id, level, qnum, query, log_file)
await asyncio.sleep(1.0) # brief pause between queries
# Final summary table
print("\n" + "="*85)
print("FINAL METRICS TABLE (HUGGING FACE SPACE)")
print("="*85)
print(f"| {'Q#':<2} | {'Difficulty':<10} | {'TTFT':<7} | {'Total Time':<10} | {'Words':<5} | {'Grounded':<8} | {'Citations':<9} | {'Quality':<7} |")
print(f"|---|------------|---------|------------|-------|----------|-----------|---------|")
for r in RESULTS:
ttft_val = f"{r['ttft']:.2f}s" if r['ttft'] else "N/A"
tot_val = f"{r['total_s']:.1f}s"
print(f"| {r['q']:02d} | {r['level']:<10} | {ttft_val:<7} | {tot_val:<10} | {r['words']:<5} | {str(r['grounded']):<8} | {r['citations']:<9} | {r['quality']:<7} |")
# Aggregate levels
print("\n" + "="*85)
print("AGGREGATED METRICS BY LEVEL")
print("="*85)
by_level = {}
for r in RESULTS:
by_level.setdefault(r["level"], []).append(r)
for level in ["easy", "medium", "hard", "very_tough"]:
group = by_level.get(level, [])
if not group:
continue
grounded_pct = sum(1 for r in group if r["grounded"]) / len(group) * 100
good_pct = sum(1 for r in group if r["quality"] == "GOOD") / len(group) * 100
avg_time = sum(r["total_s"] for r in group) / len(group)
valid_ttfts = [r["ttft"] for r in group if r["ttft"] is not None]
avg_ttft = sum(valid_ttfts) / len(valid_ttfts) if valid_ttfts else 0.0
avg_words = sum(r["words"] for r in group) / len(group)
errors = sum(1 for r in group if r["error"])
print(f"{level.upper():<10} | Grounded={grounded_pct:.0f}% | Good={good_pct:.0f}% | AvgTTFT={avg_ttft:.2f}s | AvgTime={avg_time:.1f}s | AvgWords={avg_words:.0f} | Errors={errors}")
total_grounded = sum(1 for r in RESULTS if r["grounded"])
total_good = sum(1 for r in RESULTS if r["quality"] == "GOOD")
total_errors = sum(1 for r in RESULTS if r["error"])
avg_t = sum(r["total_s"] for r in RESULTS) / len(RESULTS)
all_ttfts = [r["ttft"] for r in RESULTS if r["ttft"] is not None]
avg_ttft_all = sum(all_ttfts) / len(all_ttfts) if all_ttfts else 0.0
print(f"\nOVERALL | Grounded={total_grounded/len(RESULTS)*100:.0f}% | Good={total_good/len(RESULTS)*100:.0f}% | AvgTTFT={avg_ttft_all:.2f}s | AvgTime={avg_t:.1f}s | Errors={total_errors}")
# Save JSON
with open(r"c:\Risu Solutions\ByteAstra\backend\stress_results_hf.json", "w", encoding="utf-8") as f:
json.dump(RESULTS, f, indent=2)
print(f"\nRaw metrics saved to: c:\\Risu Solutions\\ByteAstra\\backend\\stress_results_hf.json")
print(f"Full Q&A logs saved to: {log_path}")
if __name__ == "__main__":
asyncio.run(main())