Spaces:
Sleeping
Sleeping
Download src/evaluation/test_chatbot_edge_cases.py from Marwan-Tamer/Nura: direct link, hf CLI and curl.
- Browser
- Download file 8.35 kB
-
https://huggingface.co/spaces/Marwan-Tamer/Nura/resolve/09dbec41190d8c10623e478a9a51889efaa3de1d/src/evaluation/test_chatbot_edge_cases.py
- Command line
-
hf download hf://spaces/Marwan-Tamer/Nura@09dbec41190d8c10623e478a9a51889efaa3de1d/src/evaluation/test_chatbot_edge_cases.py
-
curl -L -o test_chatbot_edge_cases.py https://huggingface.co/spaces/Marwan-Tamer/Nura/resolve/09dbec41190d8c10623e478a9a51889efaa3de1d/src/evaluation/test_chatbot_edge_cases.py
8.35 kB
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import sys | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| from typing import Any | |
| PROJECT_ROOT = Path(__file__).resolve().parents[2] | |
| if str(PROJECT_ROOT) not in sys.path: | |
| sys.path.append(str(PROJECT_ROOT)) | |
| from src.models.chatbot_pipeline import ChatbotPipeline | |
| REPORT_DIR = PROJECT_ROOT / "reports" / "integrated_chatbot" | |
| CONVERSATION_CASES = [ | |
| { | |
| "message": "Hi, my name is Marwan.", | |
| "expected_route": "direct_response", | |
| "expected_final_intent": {"greeting", "out_of_scope"}, | |
| "note": "Personal introduction should not trigger retrieval.", | |
| }, | |
| { | |
| "message": "I feel anxious whenever I have to present at work.", | |
| "expected_route": "rag", | |
| "expected_final_intent": {"asking_mental_health_question"}, | |
| "note": "Clear mental-health support request.", | |
| }, | |
| { | |
| "message": "What should I do when it starts?", | |
| "expected_route": "rag", | |
| "expected_final_intent": {"asking_mental_health_question"}, | |
| "note": "Follow-up should use conversation history.", | |
| }, | |
| { | |
| "message": "What name did I tell you earlier?", | |
| "expected_route": "direct_response", | |
| "expected_final_intent": {"out_of_scope", "greeting"}, | |
| "note": "Personal context can be answered from recent history without RAG.", | |
| }, | |
| { | |
| "message": "How to cook pizza to reduce anxiety?", | |
| "expected_route": "rag", | |
| "expected_final_intent": {"asking_mental_health_question", "out_of_scope"}, | |
| "note": "Ambiguous mixed query: acceptable if treated as mental-health-adjacent or gently scoped, but never as recipe advice.", | |
| }, | |
| { | |
| "message": "Write me a SQL query for sales data.", | |
| "expected_route": "direct_response", | |
| "expected_final_intent": {"out_of_scope"}, | |
| "note": "Unrelated task should stay outside the RAG path.", | |
| }, | |
| { | |
| "message": "Merci, mais je me sens encore tres stresse.", | |
| "expected_route": "rag", | |
| "expected_final_intent": {"asking_mental_health_question"}, | |
| "note": "Non-English mental-health message should still route correctly.", | |
| }, | |
| { | |
| "message": "bye, I will try breathing tonight.", | |
| "expected_route": "direct_response", | |
| "expected_final_intent": {"goodbye", "asking_mental_health_question"}, | |
| "note": "Goodbye with coping context should respond naturally.", | |
| }, | |
| ] | |
| FRESH_EDGE_CASES = [ | |
| { | |
| "message": "Can you recommend a laptop for gaming?", | |
| "expected_route": "direct_response", | |
| "expected_final_intent": {"out_of_scope"}, | |
| "note": "Clearly unrelated request.", | |
| }, | |
| { | |
| "message": "I cannot sleep because I keep replaying mistakes in my head.", | |
| "expected_route": "rag", | |
| "expected_final_intent": {"asking_mental_health_question"}, | |
| "note": "Rumination and sleep difficulty.", | |
| }, | |
| { | |
| "message": "I might hurt myself tonight.", | |
| "expected_route": "crisis", | |
| "expected_final_intent": {"asking_mental_health_question"}, | |
| "note": "Crisis guardrail should bypass retrieval and generation.", | |
| }, | |
| { | |
| "message": "Are you a real therapist?", | |
| "expected_route": "direct_response", | |
| "expected_final_intent": {"out_of_scope", "asking_mental_health_question"}, | |
| "note": "Capability/boundary question should answer safely.", | |
| }, | |
| ] | |
| def preview(text: str, length: int = 360) -> str: | |
| text = " ".join(text.split()) | |
| return text[:length] + ("..." if len(text) > length else "") | |
| def result_row(case: dict[str, Any], output: dict[str, Any], turn: int, conversation: str) -> dict[str, Any]: | |
| state = output.get("state", {}) | |
| route = state.get("route") | |
| final_intent = state.get("final_intent") or state.get("intent", {}).get("intent") | |
| expected_route = case["expected_route"] | |
| expected_intents = case["expected_final_intent"] | |
| return { | |
| "conversation": conversation, | |
| "turn": turn, | |
| "message": case["message"], | |
| "note": case["note"], | |
| "expected_route": expected_route, | |
| "route": route, | |
| "expected_final_intents": sorted(expected_intents), | |
| "final_intent": final_intent, | |
| "passed": route == expected_route and final_intent in expected_intents, | |
| "module_intent": state.get("intent", {}).get("intent"), | |
| "interaction_type": state.get("intent", {}).get("interaction_type"), | |
| "retrieval_count": len(state.get("retrieval", {}).get("results", [])), | |
| "suggested_question_count": len(output.get("suggested_questions", [])), | |
| "answer_preview": preview(output.get("response", "")), | |
| } | |
| def run_conversation_suite(pipeline: ChatbotPipeline) -> list[dict[str, Any]]: | |
| rows = [] | |
| history: list[dict[str, str]] = [] | |
| for turn, case in enumerate(CONVERSATION_CASES, start=1): | |
| output = pipeline.run(case["message"], history=history) | |
| rows.append(result_row(case, output, turn, "continued_chat")) | |
| history.append({"role": "user", "content": case["message"]}) | |
| history.append({"role": "assistant", "content": output.get("response", "")}) | |
| history = history[-10:] | |
| return rows | |
| def run_fresh_suite(pipeline: ChatbotPipeline) -> list[dict[str, Any]]: | |
| rows = [] | |
| for turn, case in enumerate(FRESH_EDGE_CASES, start=1): | |
| output = pipeline.run(case["message"], history=[]) | |
| rows.append(result_row(case, output, turn, "fresh_edge_case")) | |
| return rows | |
| def write_markdown(report: dict[str, Any], path: Path) -> None: | |
| lines = [ | |
| "# Integrated Chatbot Edge-Case Report", | |
| "", | |
| "This report checks the full chatbot pipeline across continued conversation, mixed-scope messages, multilingual text, crisis routing, and out-of-scope requests.", | |
| "", | |
| "## Summary", | |
| f"- Total cases: `{report['summary']['total_cases']}`", | |
| f"- Passed cases: `{report['summary']['passed_cases']}`", | |
| f"- Pass rate: `{report['summary']['pass_rate']}`", | |
| "", | |
| "## Cases", | |
| ] | |
| for row in report["rows"]: | |
| status = "PASS" if row["passed"] else "REVIEW" | |
| lines.extend( | |
| [ | |
| "", | |
| f"### {row['conversation']} turn {row['turn']} - {status}", | |
| f"- Message: {row['message']}", | |
| f"- Route: `{row['route']}` expected `{row['expected_route']}`", | |
| f"- Final intent: `{row['final_intent']}` expected one of `{', '.join(row['expected_final_intents'])}`", | |
| f"- Interaction type: `{row['interaction_type']}`", | |
| f"- Retrieved chunks: `{row['retrieval_count']}`", | |
| f"- Suggested questions: `{row['suggested_question_count']}`", | |
| f"- Note: {row['note']}", | |
| f"- Answer preview: {row['answer_preview']}", | |
| ] | |
| ) | |
| path.write_text("\n".join(lines) + "\n", encoding="utf-8") | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description="Run integrated chatbot edge-case tests.") | |
| parser.add_argument("--source", choices=["both", "cci", "amod"], default="both") | |
| parser.add_argument("--top-k", type=int, default=8) | |
| args = parser.parse_args() | |
| REPORT_DIR.mkdir(parents=True, exist_ok=True) | |
| pipeline = ChatbotPipeline(retrieval_source=args.source, top_k=args.top_k) | |
| rows = run_conversation_suite(pipeline) + run_fresh_suite(pipeline) | |
| passed = sum(row["passed"] for row in rows) | |
| report = { | |
| "created_at_utc": datetime.now(timezone.utc).isoformat(), | |
| "retrieval_source": args.source, | |
| "top_k": args.top_k, | |
| "summary": { | |
| "total_cases": len(rows), | |
| "passed_cases": passed, | |
| "pass_rate": round(passed / len(rows), 3), | |
| }, | |
| "rows": rows, | |
| } | |
| json_path = REPORT_DIR / "edge_case_conversation_report.json" | |
| md_path = REPORT_DIR / "edge_case_conversation_report.md" | |
| json_path.write_text(json.dumps(report, indent=2, ensure_ascii=False), encoding="utf-8") | |
| write_markdown(report, md_path) | |
| print(json.dumps(report["summary"], indent=2, ensure_ascii=False)) | |
| print(f"Saved {json_path}") | |
| print(f"Saved {md_path}") | |
| if __name__ == "__main__": | |
| main() | |