import os, uuid, subprocess from flask import Flask, request, jsonify from faster_whisper import WhisperModel from transformers import AutoTokenizer, AutoModelForSeq2SeqLM from deep_translator import GoogleTranslator import torch app = Flask(__name__) # ---------------- ERROR HANDLER ---------------- @app.errorhandler(Exception) def handle_error(e): return jsonify({"error": str(e)}), 500 # ---------------- LOAD MODELS ---------------- print("Loading models...") whisper_model = WhisperModel("small", device="cpu", compute_type="int8") bart_tokenizer = AutoTokenizer.from_pretrained("facebook/bart-base") bart_model = AutoModelForSeq2SeqLM.from_pretrained("facebook/bart-base") print("Models loaded successfully") # ---------------- CLEAN TEXT ---------------- def clean_text(text): with torch.no_grad(): inputs = bart_tokenizer( "Extract key medical symptoms: " + text, return_tensors="pt", truncation=True ) outputs = bart_model.generate(**inputs, max_length=50) return bart_tokenizer.decode(outputs[0], skip_special_tokens=True) # ---------------- SYMPTOM WEIGHTS ---------------- SYMPTOM_WEIGHTS = { # 🚨 Critical "chest pain": 50, "breathing": 50, "shortness of breath": 50, "unconscious": 60, "stroke": 60, "bleeding": 50, # ⚠️ Moderate "fever": 25, "vomiting": 25, "weakness": 20, "body pain": 15, "dizziness": 20, # ✅ Mild "cold": 5, "runny nose": 5, "sneezing": 5, "cough": 5, "gas": 5, "acidity": 5, "headache": 5 } # ---------------- SCORING ---------------- def rule_based_score(text): t = text.lower() score = 0 matched = [] for symptom, weight in SYMPTOM_WEIGHTS.items(): if symptom in t: score += weight matched.append(symptom) return score, matched def map_severity(score): if score >= 70: return "emergency" elif score >= 30: return "urgent" else: return "non-urgent" # ---------------- CORE PIPELINE ---------------- def run_pipeline(text): # Translate try: text_en = GoogleTranslator(source="auto", target="en").translate(text) except: text_en = text # Clean cleaned = clean_text(text_en) # Rule scoring score, matched = rule_based_score(cleaned) severity = map_severity(score) return { "processed_text": cleaned, "matched_symptoms": matched, "severity": severity, "score": min(score, 100) } # ---------------- ROUTES ---------------- @app.route("/", methods=["GET"]) def root(): return jsonify({ "status": "running", "message": "AI Triage API live", "endpoints": { "audio": "/process-audio", "text": "/process-text" } }) # 🎤 AUDIO @app.route("/process-audio", methods=["POST"]) def process_audio(): if "audio" not in request.files: return jsonify({"error": "No audio file"}), 400 file = request.files["audio"] uid = str(uuid.uuid4()) input_path = f"{uid}.webm" wav_path = f"{uid}.wav" file.save(input_path) subprocess.run([ "ffmpeg", "-y", "-i", input_path, "-ac", "1", "-ar", "16000", wav_path ], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) segments, _ = whisper_model.transcribe(wav_path) text = " ".join([seg.text for seg in segments]).strip() # cleanup os.remove(input_path) if os.path.exists(wav_path): os.remove(wav_path) if not text: return jsonify({"error": "No speech detected"}), 400 result = run_pipeline(text) return jsonify({ "input_type": "audio", "original_text": text, **result }) # ⌨️ TEXT @app.route("/process-text", methods=["POST"]) def process_text(): text = request.form.get("text") if not text: return jsonify({"error": "No text provided"}), 400 result = run_pipeline(text) return jsonify({ "input_type": "text", "original_text": text, **result }) # ---------------- RUN ---------------- if __name__ == "__main__": app.run(host="0.0.0.0", port=7860)