""" ============================================================ Conversation Manager ============================================================ هنا أهم module يشتغل على المتطلبات الجديدة: 1. لا يخرج تشخيص بعد كل جملة: - يجمّع كل رسائل المستخدم - يطلب حد أدنى من الرسائل/الكلمات - يتحقق من اتساق التنبؤات على آخر N رسائل 2. يخرج التشخيص لما "متأكد 100%": - confidence threshold (75%) - consistency check (66% من آخر 3 تنبؤات لازم نفس النتيجة) - high-risk override (Suicide بيطلع فوراً) 3. يدّي نصايح حسب التشخيص النهائي """ import time import uuid from collections import Counter, deque from typing import List, Dict, Optional from dataclasses import dataclass, field from config import DIAGNOSIS_CONFIG, SEVERITY_LEVELS @dataclass class Message: role: str # 'user' | 'assistant' content: str timestamp: float = field(default_factory=time.time) prediction: Optional[dict] = None # تنبؤ النموذج لو رسالة من user @dataclass class ConversationState: """حالة المحادثة - بيتخزّن لكل جلسة""" session_id: str messages: List[Message] = field(default_factory=list) diagnosis_confirmed: bool = False final_diagnosis: Optional[str] = None final_confidence: float = 0.0 created_at: float = field(default_factory=time.time) class ConversationManager: """ يدير عدة جلسات في الذاكرة (in-memory). لـ production استبدل بـ Redis أو DB. """ def __init__(self): self.sessions: Dict[str, ConversationState] = {} def create_session(self) -> str: sid = str(uuid.uuid4()) self.sessions[sid] = ConversationState(session_id=sid) return sid def get_session(self, session_id: str) -> ConversationState: if session_id not in self.sessions: self.sessions[session_id] = ConversationState(session_id=session_id) return self.sessions[session_id] def add_user_message(self, session_id: str, text: str, prediction: dict = None): state = self.get_session(session_id) msg = Message(role='user', content=text, prediction=prediction) state.messages.append(msg) # Trim history max_msgs = DIAGNOSIS_CONFIG['max_history_messages'] if len(state.messages) > max_msgs: state.messages = state.messages[-max_msgs:] def add_assistant_message(self, session_id: str, text: str): state = self.get_session(session_id) state.messages.append(Message(role='assistant', content=text)) def get_user_messages(self, session_id: str) -> List[Message]: state = self.get_session(session_id) return [m for m in state.messages if m.role == 'user'] def get_combined_text(self, session_id: str) -> str: """يجمع كل رسائل المستخدم لتحليل شامل""" user_msgs = self.get_user_messages(session_id) return ' '.join(m.content for m in user_msgs) def get_recent_predictions(self, session_id: str, n: int = None) -> List[dict]: if n is None: n = DIAGNOSIS_CONFIG['consistency_window'] user_msgs = self.get_user_messages(session_id) return [m.prediction for m in user_msgs[-n:] if m.prediction is not None] def should_diagnose(self, session_id: str, latest_pred: dict) -> Dict: """ القلب النابض للنظام: يقرر هل نطلع تشخيص ولا لسه نكمّل أسئلة. الشروط: 1. حالة طارئة (Suicide) → فوري 2. عدد الرسائل >= min_messages 3. عدد الكلمات الكلي >= min_total_words 4. الـ confidence >= threshold 5. آخر 3 تنبؤات متفقة على نفس الـ label """ state = self.get_session(session_id) # ===== Rule 1: High-risk override ===== if latest_pred and latest_pred['label'] in DIAGNOSIS_CONFIG['high_risk_immediate']: if latest_pred['confidence'] >= 0.5: # حد أدنى للحالة الطارئة return { 'diagnose': True, 'reason': 'high_risk', 'urgency': 'critical', 'label': latest_pred['label'], 'confidence': latest_pred['confidence'] } # ===== Rule 2: Min messages ===== user_msgs = self.get_user_messages(session_id) if len(user_msgs) < DIAGNOSIS_CONFIG['min_messages_before_diagnosis']: return { 'diagnose': False, 'reason': 'insufficient_messages', 'progress': { 'messages': len(user_msgs), 'required': DIAGNOSIS_CONFIG['min_messages_before_diagnosis'] } } # ===== Rule 3: Min words ===== combined = self.get_combined_text(session_id) word_count = len(combined.split()) if word_count < DIAGNOSIS_CONFIG['min_total_words']: return { 'diagnose': False, 'reason': 'insufficient_words', 'progress': { 'words': word_count, 'required': DIAGNOSIS_CONFIG['min_total_words'] } } # ===== Rule 4: Confidence threshold ===== if latest_pred['confidence'] < DIAGNOSIS_CONFIG['confidence_threshold']: return { 'diagnose': False, 'reason': 'low_confidence', 'progress': { 'confidence': latest_pred['confidence'], 'required': DIAGNOSIS_CONFIG['confidence_threshold'] } } # ===== Rule 5: Consistency ===== recent_preds = self.get_recent_predictions(session_id) if len(recent_preds) < DIAGNOSIS_CONFIG['consistency_window']: return { 'diagnose': False, 'reason': 'insufficient_history_for_consistency' } labels_history = [p['label'] for p in recent_preds] most_common_label, count = Counter(labels_history).most_common(1)[0] consistency_ratio = count / len(labels_history) if consistency_ratio < DIAGNOSIS_CONFIG['consistency_threshold']: return { 'diagnose': False, 'reason': 'inconsistent_predictions', 'progress': { 'consistency': consistency_ratio, 'required': DIAGNOSIS_CONFIG['consistency_threshold'] } } # ===== كل الشروط متحققة → نطلع تشخيص ===== # نحسب الـ confidence المتوسط للـ label الفائز avg_conf = sum(p['confidence'] for p in recent_preds if p['label'] == most_common_label) / count return { 'diagnose': True, 'reason': 'criteria_met', 'urgency': 'normal', 'label': most_common_label, 'confidence': avg_conf, 'consistency': consistency_ratio } def confirm_diagnosis(self, session_id: str, label: str, confidence: float): state = self.get_session(session_id) state.diagnosis_confirmed = True state.final_diagnosis = label state.final_confidence = confidence def reset(self, session_id: str): if session_id in self.sessions: del self.sessions[session_id]