File size: 7,981 Bytes
4eb1ddb
 
 
 
 
 
 
 
 
 
 
e4aaa90
4eb1ddb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e4aaa90
 
4eb1ddb
e4aaa90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4eb1ddb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e4aaa90
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
"""
============================================================
Mentallico API v2 - Production-ready FastAPI
============================================================
Endpoints:

POST /api/v2/session/create        - يفتح جلسة جديدة
POST /api/v2/session/{id}/message  - يرسل رسالة (نصاً) ويحصل على الرد
POST /api/v2/session/{id}/audio    - يرسل ملف صوت + يرد عليه
POST /api/v2/session/{id}/reset    - يصفّي الجلسة
GET  /api/v2/session/{id}/state    - حالة الجلسة (debugging)
POST /api/v2/upload_documents      - يضيف PDF للـ knowledge base
GET  /api/v2/health                - فحص النظام

كل response بيتبع schema موحّد عشان frontend (Flutter / Web) يقدر يبني UI متّسق.
"""
import os
import shutil
import tempfile
from typing import List, Optional
from fastapi import FastAPI, HTTPException, UploadFile, File, Form
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
import uvicorn

from orchestrator import get_brain

# ============= Schemas =============
class MessageInput(BaseModel):
    text: str = Field(..., description="User message text")


class MessageResponse(BaseModel):
    type: str = Field(..., description="chat | probing | diagnosis | error")
    response: str
    diagnosis: Optional[str] = None
    confidence: float = 0.0
    urgency: str = "normal"
    session_id: str
    audio_info: Optional[dict] = None
    internal: Optional[dict] = None


class SessionResponse(BaseModel):
    session_id: str


class HealthResponse(BaseModel):
    status: str
    components: dict


class UploadResult(BaseModel):
    filename: str
    status: str
    message: str


class UploadResponse(BaseModel):
    details: List[UploadResult]


# ============= App =============
app = FastAPI(
    title="Mentallico AI v2",
    description="Integrated mental-health diagnosis & support system "
                "(RAG + Multi-Model Classifier + STT)",
    version="2.0.0"
)

# CORS - frontend Flutter ممكن يحتاجه
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],  # في production غيّرها لـ domain الـ frontend
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)


# ============= Endpoints =============

@app.get("/api/v2/health", response_model=HealthResponse)
async def health():
    """فحص حالة المكونات"""
    try:
        brain = get_brain()
        components = {
            "classifier_primary": brain.diagnoser.primary_model is not None,
            "classifier_secondary": brain.diagnoser.secondary_model is not None,
            "classifier_emotion": brain.diagnoser.emotion_pipe is not None,
            "rag": brain.rag is not None and brain.rag.vector_db is not None,
            "stt": brain.transcriber is not None,
        }
        all_ok = all(components.values())
        # نسمح بفقدان بعض المكونات بس ميكونش fail بالكامل
        critical = components["rag"] or any([
            components["classifier_primary"],
            components["classifier_secondary"],
            components["classifier_emotion"]
        ])
        return HealthResponse(
            status="ok" if critical else "degraded",
            components=components
        )
    except Exception as e:
        return HealthResponse(
            status="error",
            components={"error": str(e)}
        )


@app.post("/api/v2/session/create", response_model=SessionResponse)
async def create_session():
    """ينشئ جلسة جديدة - بترجع session_id يستخدمه الـ frontend في الرسائل التالية"""
    brain = get_brain()
    sid = brain.create_session()
    return SessionResponse(session_id=sid)


@app.post("/api/v2/session/{session_id}/message", response_model=MessageResponse)
async def send_message(session_id: str, payload: MessageInput):
    """رسالة نصية"""
    if not payload.text or not payload.text.strip():
        raise HTTPException(status_code=400, detail="Empty message")
    brain = get_brain()
    result = brain.process_message(session_id, payload.text.strip())
    return MessageResponse(**result)


@app.post("/api/v2/session/{session_id}/audio", response_model=MessageResponse)
async def send_audio(
    session_id: str,
    audio: UploadFile = File(...),
    language: Optional[str] = Form(None)
):
    """رسالة صوتية - بتعدي STT الأول ثم باقي الـ pipeline"""
    if not audio.filename:
        raise HTTPException(status_code=400, detail="No audio file provided")

    # حفظ مؤقت
    suffix = os.path.splitext(audio.filename)[1] or ".wav"
    with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
        shutil.copyfileobj(audio.file, tmp)
        tmp_path = tmp.name

    try:
        brain = get_brain()
        # نعدي بالـ audio path - الـ orchestrator هيعمل STT ثم باقي الـ flow
        result = brain.process_message(session_id, user_text="", user_audio_path=tmp_path)
        return MessageResponse(**result)
    finally:
        if os.path.exists(tmp_path):
            os.remove(tmp_path)


@app.post("/api/v2/session/{session_id}/reset")
async def reset_session(session_id: str):
    """يمسح الجلسة"""
    brain = get_brain()
    brain.reset_session(session_id)
    return {"status": "ok", "message": f"Session {session_id} reset"}


@app.get("/api/v2/session/{session_id}/state")
async def session_state(session_id: str):
    """حالة الجلسة - مفيد للـ debugging"""
    brain = get_brain()
    state = brain.conv_mgr.get_session(session_id)
    return {
        "session_id": state.session_id,
        "messages_count": len(state.messages),
        "user_messages_count": len([m for m in state.messages if m.role == 'user']),
        "diagnosis_confirmed": state.diagnosis_confirmed,
        "final_diagnosis": state.final_diagnosis,
        "final_confidence": state.final_confidence,
        "messages": [
            {"role": m.role, "content": m.content,
             "prediction": m.prediction['label'] if m.prediction else None}
            for m in state.messages
        ]
    }


@app.post("/api/v2/upload_documents", response_model=UploadResponse)
async def upload_documents(file: UploadFile = File(...)):
    """رفع ملف PDF واحد للـ knowledge base"""
    brain = get_brain()
    
    if not file.filename.lower().endswith(".pdf"):
        return UploadResponse(details=[UploadResult(
            filename=file.filename,
            status="error",
            message="Only PDF files are allowed"
        )])

    # حفظ مؤقت
    with tempfile.NamedTemporaryFile(delete=False, suffix=".pdf") as tmp:
        shutil.copyfileobj(file.file, tmp)
        tmp_path = tmp.name

    try:
        result = brain.add_pdf(tmp_path)
        return UploadResponse(details=[UploadResult(
            filename=file.filename,
            status=result['status'],
            message=result['message']
        )])
    finally:
        if os.path.exists(tmp_path):
            os.remove(tmp_path)


# ============= Backwards-compat (الـ endpoint القديم) =============
@app.post("/api/v1/diagnose")
async def diagnose_legacy(payload: MessageInput):
    """
    Backwards-compatible endpoint - بيعمل single-shot diagnosis
    (يفتح session مؤقت، يبعت رسالة، يرجع الرد، يقفل الـ session).
    موجود عشان الـ frontend القديم يكمل شغل من غير ما نعدّل عليه دلوقتي.
    """
    brain = get_brain()
    sid = brain.create_session()
    try:
        result = brain.process_message(sid, payload.text)
        return {"answer": result['response']}
    finally:
        brain.reset_session(sid)


# ============= Run =============
if __name__ == "__main__":
    uvicorn.run("app:app", host="0.0.0.0", port=7860, reload=False)