#!/usr/bin/env python3
"""
CHIEF MEDICAL DATA OFFICER — Gesundheitsmanager
Delegierter Subagent für alle Gesundheitsdaten-Analysen.

Rolle:
- Verarbeitet und analysiert alle Gesundheitsdaten
- Extrahiert Laborwerte aus PDFs (Camelot, PyMuPDF, pdfplumber, pytesseract)
- Führt statistische Analysen durch (scipy, statsmodels, pandas)
- Erstellt Visualisierungen (matplotlib, seaborn, plotly)
- Generiert Reports und Korrelationen
- Überwacht Health_Inbox via watchdog für automatische PDF-Verarbeitung

Datenbank: /home/agent/.hermes/assets/Gesundheit/health_data.db
Archiv: /home/agent/.hermes/assets/Gesundheit/archiv/
Inbox: /home/agent/.hermes/assets/Gesundheit/inbox/
"""

import pandas as pd
import numpy as np
import sqlalchemy
from sqlalchemy import create_engine, text
import scipy.stats as stats
import statsmodels.api as sm
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import seaborn as sns
import plotly.express as px
import plotly.io as pio
import camelot
import fitz  # PyMuPDF
import pdfplumber
import pytesseract
from watchdog.observers import Observer
from watchdog.events import FileSystemEventHandler
import os
import sys
import json
import hashlib
from datetime import datetime, timedelta
from pathlib import Path
import warnings
warnings.filterwarnings('ignore')

# Konfiguration
DB_PATH = '/home/agent/.hermes/assets/Gesundheit/health_data.db'
ARCHIV_DIR = '/home/agent/.hermes/assets/Gesundheit/archiv'
INBOX_DIR = '/home/agent/.hermes/assets/Gesundheit/inbox'
REPORTS_DIR = '/home/agent/.hermes/assets/Gesundheit/reports'
EXPORTS_DIR = '/home/agent/.hermes/assets/Gesundheit/exports'

# Erstelle Reports-Ordner falls nicht vorhanden
os.makedirs(REPORTS_DIR, exist_ok=True)
os.makedirs(EXPORTS_DIR, exist_ok=True)


class HealthManager:
    """Hauptklasse für die Gesundheitsdaten-Verwaltung und -Analyse."""
    
    def __init__(self):
        self.engine = create_engine(f'sqlite:///{DB_PATH}')
        self.conn = self.engine.connect()
        self._validate_schema()
    
    def _validate_schema(self):
        """Prüfe ob alle erforderlichen Tabellen existieren."""
        required_tables = [
            'laborwerte', 'health_events', 'dokumente', 'medikamente',
            'dokumente_status', 'symptome', 'symptom_log', 'vitalzeichen',
            'arztbesuche', 'ernaehrung', 'auswertungen'
        ]
        
        existing = pd.read_sql("SELECT name FROM sqlite_master WHERE type='table'", self.conn)
        existing_tables = set(existing['name'])
        
        missing = [t for t in required_tables if t not in existing_tables]
        if missing:
            print(f"[WARN] Fehlende Tabellen: {missing}")
            self._create_missing_tables(missing)
    
    def _create_missing_tables(self, missing_tables):
        """Erstelle fehlende Tabellen."""
        table_definitions = {
            'symptome': """
                CREATE TABLE IF NOT EXISTS symptome (
                    id INTEGER PRIMARY KEY AUTOINCREMENT,
                    datum DATE NOT NULL,
                    symptombeschreibung TEXT,
                    schwerigkeit INTEGER CHECK(schwerigkeit BETWEEN 1 AND 10),
                    bereich TEXT,
                    dokumentiert BOOLEAN DEFAULT 1
                )
            """,
            'symptom_log': """
                CREATE TABLE IF NOT EXISTS symptom_log (
                    id INTEGER PRIMARY KEY AUTOINCREMENT,
                    datum TIMESTAMP NOT NULL,
                    symptombeschreibung TEXT NOT NULL,
                    schwerigkeit INTEGER CHECK(schwerigkeit BETWEEN 1 AND 10),
                    bereich TEXT,
                    einflussfaktoren TEXT,
                    medikamente_eingenommen TEXT,
                    notizen TEXT,
                    status TEXT DEFAULT 'aktiv'
                )
            """,
            'vitalzeichen': """
                CREATE TABLE IF NOT EXISTS vitalzeichen (
                    id INTEGER PRIMARY KEY AUTOINCREMENT,
                    datum TIMESTAMP NOT NULL,
                    typ TEXT NOT NULL,
                    wert REAL NOT NULL,
                    einheit TEXT,
                    referenzbereich TEXT,
                    gerat TEXT,
                    notizen TEXT
                )
            """,
            'arztbesuche': """
                CREATE TABLE IF NOT EXISTS arztbesuche (
                    id INTEGER PRIMARY KEY AUTOINCREMENT,
                    datum DATE NOT NULL,
                    arzt_name TEXT,
                    spezialisation TEXT,
                    grund TEXT,
                    zusammenfassung TEXT,
                    naechster_termin DATE,
                    notizen TEXT
                )
            """,
            'ernaehrung': """
                CREATE TABLE IF NOT EXISTS ernaehrung (
                    id INTEGER PRIMARY KEY AUTOINCREMENT,
                    datum DATE NOT NULL,
                    mahlzeit TEXT,
                    kalorien INTEGER,
                    protein_g REAL,
                    kohlenhydrate_g REAL,
                    fett_g REAL,
                    ballaststoffe_g REAL,
                    histamin_score INTEGER CHECK(histamin_score BETWEEN 0 AND 10),
                    notizen TEXT
                )
            """,
            'auswertungen': """
                CREATE TABLE IF NOT EXISTS auswertungen (
                    id INTEGER PRIMARY KEY AUTOINCREMENT,
                    datum TIMESTAMP NOT NULL,
                    typ TEXT NOT NULL,
                    name TEXT NOT NULL,
                    ergebnis TEXT,
                    dateipfad TEXT,
                    parameter TEXT,
                    metadata TEXT
                )
            """
        }
        
        for table in missing_tables:
            if table in table_definitions:
                self.conn.execute(text(table_definitions[table]))
                self.conn.commit()
                print(f"[INFO] Tabelle '{table}' erstellt.")
    
    def get_laborwerte(self, parameter=None, start_date=None, end_date=None):
        """Laborwerte aus der Datenbank laden (FRIDAY-Schema)."""
        query = "SELECT * FROM laborwerte WHERE 1=1"
        params = {}
        
        if parameter:
            query += " AND parameter_name LIKE :param"
            params['param'] = f"%{parameter}%"
        if start_date:
            query += " AND ermittlung_datum >= :start"
            params['start'] = start_date
        if end_date:
            query += " AND ermittlung_datum <= :end"
            params['end'] = end_date
        
        df = pd.read_sql(query, self.engine, params=params)
        df['datum'] = pd.to_datetime(df['ermittlung_datum'])
        df['referenzbereich'] = None  # FRIDAY hat keine Referenzbereiche
        return df
    
    def get_dokumente(self, status=None, kategorie=None):
        """Dokumente aus der Datenbank laden (FRIDAY-Schema)."""
        query = "SELECT * FROM dokumente WHERE 1=1"
        params = {}
        if status:
            query += " AND status = :status"
            params['status'] = status
        if kategorie:
            query += " AND kategorie = :kategorie"
            params['kategorie'] = kategorie
        
        return pd.read_sql(query, self.engine, params=params)
    
    def extract_laborwerte_from_pdf(self, pdf_path):
        """Laborwerte aus einem PDF extrahieren (Camelot + PyMuPDF)."""
        extracted_values = []
        
        # Versuche Camelot für Tabellen
        try:
            tables = camelot.read_pdf(pdf_path, pages='all')
            for table in tables:
                df_table = table.df
                if len(df_table) > 2:  # Mindestens 3 Zeilen (Header + 2 Daten)
                    # Erkenne Laborwert-Spalten
                    for i in range(1, len(df_table)):
                        row = df_table.iloc[i]
                        if len(row) >= 3:
                            parameter = str(row[0]).strip() if pd.notna(row[0]) else None
                            wert = str(row[1]).strip() if pd.notna(row[1]) else None
                            einheit = str(row[2]).strip() if pd.notna(row[2]) else None
                            referenz = str(row[3]).strip() if len(row) > 3 and pd.notna(row[3]) else None
                            
                            if parameter and wert:
                                extracted_values.append({
                                    'parameter': parameter,
                                    'wert': wert,
                                    'einheit': einheit,
                                    'referenzbereich': referenz,
                                    'quelle': pdf_path
                                })
        except Exception as e:
            print(f"[WARN] Camelot fehlgeschlagen: {e}")
        
        # Fallback: PyMuPDF für Fließtext
        if not extracted_values:
            try:
                doc = fitz.open(pdf_path)
                pdf_text = ""
                for page in doc:
                    pdf_text += page.get_text()
                
                # Einfache Regex für Laborwerte
                import re
                pattern = r'(\w[\w\s\.]*)\s+([\d\.]+)\s*(\w+/\w*|\w+|\d+)\s*([<>=\d\.\s]+)?'
                matches = re.findall(pattern, pdf_text)
                for match in matches:
                    extracted_values.append({
                        'parameter': match[0].strip(),
                        'wert': match[1].strip(),
                        'einheit': match[2].strip(),
                        'referenzbereich': match[3].strip() if match[3] else None,
                        'quelle': pdf_path
                    })
            except Exception as e:
                print(f"[WARN] PyMuPDF fehlgeschlagen: {e}")
        
        return extracted_values
    
    def extract_laborwerte_from_pdf_plumber(self, pdf_path):
        """Fallback: pdfplumber für PDF-Extraktion."""
        extracted_values = []
        try:
            with pdfplumber.open(pdf_path) as pdf:
                for page in pdf.pages:
                    page_text = page.extract_text()
                    if page_text:
                        import re
                        pattern = r'(\w[\w\s\.]*)\s+([\d\.]+)\s*(\w+/\w*|\w+|\d+)\s*([<>=\d\.\s]+)?'
                        matches = re.findall(pattern, page_text)
                        for match in matches:
                            extracted_values.append({
                                'parameter': match[0].strip(),
                                'wert': match[1].strip(),
                                'einheit': match[2].strip(),
                                'referenzbereich': match[3].strip() if match[3] else None,
                                'quelle': pdf_path
                            })
        except Exception as e:
            print(f"[WARN] pdfplumber fehlgeschlagen: {e}")
        return extracted_values
    
    def extract_from_scanned_pdf(self, pdf_path):
        """OCR für gescannte PDFs (pytesseract)."""
        import fitz
        extracted_text = []
        try:
            doc = fitz.open(pdf_path)
            for page_num in range(len(doc)):
                page = doc[page_num]
                pix = page.get_pixmap(dpi=300)
                img_data = pix.tobytes("png")
                ocr_text = pytesseract.image_to_string(img_data, lang='deu+eng')
                extracted_text.append(ocr_text)
        except Exception as e:
            print(f"[WARN] OCR fehlgeschlagen: {e}")
        return "\n".join(extracted_text)
    
    def compute_crp_trend(self):
        """CRP-Trend berechnen."""
        df = self.get_laborwerte('CRP')
        if df.empty:
            return None
        
        df = df.sort_values('ermittlung_datum')
        
        # Konvertiere Werte zu Floats (handle '<5.0' etc.)
        def parse_val(v):
            try:
                return float(str(v).lstrip('<>='))
            except:
                return 0.0
        
        df['wert_float'] = df['wert'].apply(parse_val)
        values = df['wert_float']
        
        if len(values) < 2:
            return {
                'trend': 'unzureichend Daten',
                'werte': values.tolist(),
                'aktuell': float(values.iloc[-1])
            }
        
        # Lineare Regression für Trend
        x = np.arange(len(values))
        slope, intercept, r_value, p_value, std_err = stats.linregress(x, values)
        
        return {
            'trend': 'steigend' if slope > 0.1 else ('fallend' if slope < -0.1 else 'stabil'),
            'slope': slope,
            'r_squared': r_value ** 2,
            'p_value': p_value,
            'werte': values.tolist(),
            'aktuell': float(values.iloc[-1]),
            'vorhersage': float(slope * len(values) + intercept)
        }
    
    def compute_correlations(self, target_parameter=None):
        """Korrelationsanalyse zwischen Laborparametern."""
        df = self.get_laborwerte()
        if df.empty:
            return {}
        
        # Konvertiere 'wert' zu numerisch
        def parse_val(v):
            try:
                return float(str(v).lstrip('<>='))
            except:
                return np.nan
        
        df['wert_float'] = df['wert'].apply(parse_val)
        
        # Pivot: Parameter als Spalten, Datum als Index
        pivot = df.pivot_table(index='datum', columns='parameter_name', values='wert_float')
        pivot = pivot.dropna(axis=1, thresh=max(5, len(pivot) * 0.5))
        
        if pivot.empty:
            return {}
        
        corr_matrix = pivot.corr()
        
        if target_parameter and target_parameter in corr_matrix.columns:
            # Spezifische Korrelationen zum Target
            target_corrs = corr_matrix[target_parameter].drop(target_parameter).sort_values(ascending=False)
            return {
                'target': target_parameter,
                'korrelationen': {
                    col: round(float(val), 3)
                    for col, val in target_corrs.items()
                    if abs(val) > 0.3
                }
            }
        
        return {
            'korrelationsmatrix': corr_matrix.to_dict()
        }
    
    def generate_crp_chart(self, output_path=None):
        """CRP-Trend-Diagramm erstellen."""
        df = self.get_laborwerte('CRP')
        if df.empty:
            print("[WARN] Keine CRP-Werte gefunden.")
            return None
        
        df = df.sort_values('ermittlung_datum')
        
        # Konvertiere Werte zu Floats
        def parse_val(v):
            try:
                return float(str(v).lstrip('<>='))
            except:
                return 0.0
        
        df['wert_float'] = df['wert'].apply(parse_val)
        
        fig, ax = plt.subplots(figsize=(12, 6))
        ax.plot(df['datum'], df['wert_float'], 'b-o', markersize=6, label='CRP')
        ax.axhline(y=5, color='green', linestyle='--', label='Referenz: <5 mg/l')
        ax.fill_between(df['datum'], df['wert_float'], 5, where=df['wert_float'] >= 5,
                        alpha=0.2, color='red', label='Erhöht')
        ax.set_title('C-reaktives Protein (CRP) — Entzündungsmarker', fontsize=14)
        ax.set_xlabel('Datum')
        ax.set_ylabel('CRP (mg/l)')
        ax.legend()
        ax.grid(True, alpha=0.3)
        plt.xticks(rotation=45)
        plt.tight_layout()
        
        if output_path:
            fig.savefig(output_path, dpi=150)
            plt.close()
            return output_path
        else:
            path = os.path.join(REPORTS_DIR, 'crp_trend.png')
            fig.savefig(path, dpi=150)
            plt.close()
            return path
    
    def generate_cholesterol_chart(self, output_path=None):
        """Cholesterin-Diagramm erstellen."""
        df = self.get_laborwerte('Cholesterin')
        if df.empty:
            print("[WARN] Keine Cholesterin-Werte gefunden.")
            return None
        
        df = df.sort_values('ermittlung_datum')
        
        def parse_val(v):
            try:
                return float(str(v).lstrip('<>='))
            except:
                return 0.0
        
        df['wert_float'] = df['wert'].apply(parse_val)
        
        fig, ax = plt.subplots(figsize=(12, 6))
        
        # Splitte in LDL und Gesamt
        for param in df['parameter_name'].unique():
            subset = df[df['parameter_name'] == param].sort_values('ermittlung_datum')
            label = 'LDL' if 'LDL' in param else ('HDL' if 'HDL' in param else 'Gesamt')
            color = 'red' if label == 'LDL' else ('green' if label == 'HDL' else 'blue')
            ax.plot(subset['datum'], subset['wert_float'], '-o', markersize=6, label=label, color=color)
        
        ax.axhline(y=5.0, color='orange', linestyle='--', label='Referenz: <5.0 mmol/l')
        ax.set_title('Cholesterin-Profile', fontsize=14)
        ax.set_xlabel('Datum')
        ax.set_ylabel('mmol/l')
        ax.legend()
        ax.grid(True, alpha=0.3)
        plt.xticks(rotation=45)
        plt.tight_layout()
        
        if output_path:
            fig.savefig(output_path, dpi=150)
            plt.close()
            return output_path
        else:
            path = os.path.join(REPORTS_DIR, 'cholesterin_profile.png')
            fig.savefig(path, dpi=150)
            plt.close()
            return path
    
    def generate_heatmap(self, parameter_list=None, output_path=None):
        """Heatmap der wichtigsten Laborparameter."""
        df = self.get_laborwerte()
        if df.empty:
            print("[WARN] Keine Laborwerte gefunden.")
            return None
        
        # Konvertiere 'wert' zu numerisch
        def parse_val(v):
            try:
                return float(str(v).lstrip('<>='))
            except:
                return np.nan
        
        df['wert_float'] = df['wert'].apply(parse_val)
        
        pivot = df.pivot_table(index='datum', columns='parameter_name', values='wert_float')
        pivot = pivot.dropna(axis=1, thresh=max(5, len(pivot) * 0.5))
        
        if parameter_list:
            pivot = pivot[[p for p in parameter_list if p in pivot.columns]]
        
        if pivot.empty:
            print("[WARN] Keine Daten für Heatmap.")
            return None
        
        fig, ax = plt.subplots(figsize=(14, 8))
        sns.heatmap(pivot, annot=False, cmap='RdYlGn_r', ax=ax, linewidths=0.5)
        ax.set_title('Laborparameter Heatmap', fontsize=14)
        plt.xticks(rotation=45)
        plt.tight_layout()
        
        if output_path:
            fig.savefig(output_path, dpi=150)
            plt.close()
            return output_path
        else:
            path = os.path.join(REPORTS_DIR, 'labor_heatmap.png')
            fig.savefig(path, dpi=150)
            plt.close()
            return path
    
    def generate_plotly_dashboard(self, output_path=None):
        """Interaktiver Plotly-Dashboard für langfristige Verläufe."""
        df = self.get_laborwerte()
        if df.empty:
            print("[WARN] Keine Laborwerte für Dashboard.")
            return None
        
        fig = px.line(
            df, x='datum', y='wert', color='parameter_name',
            title='Laborwert-Verläufe (Interaktiv)',
            labels={'wert': 'Wert', 'datum': 'Datum', 'parameter_name': 'Parameter'}
        )
        fig.update_layout(
            xaxis_title='Datum',
            yaxis_title='Wert',
            hovermode='x unified',
            height=600
        )
        
        if output_path:
            fig.write_html(output_path)
            return output_path
        else:
            path = os.path.join(REPORTS_DIR, 'labor_dashboard.html')
            fig.write_html(path)
            return path
    
    def generate_profiling_report(self, output_path=None):
        """ydata-profiling HTML-Report generieren."""
        try:
            from ydata_profiling import ProfileReport
        except ImportError:
            print("[WARN] ydata-profiling nicht installiert.")
            return None
        
        df = self.get_laborwerte()
        if df.empty:
            print("[WARN] Keine Daten für Profil-Report.")
            return None
        
        report = ProfileReport(df, title='Gesundheitsdaten-Profil', progress_bar=False)
        output_path = output_path or '/home/agent/.hermes/assets/Gesundheit/reports/health_profile.html'
        report.to_file(output_path)
        print(f"[INFO] ydata-profiling Report: {output_path}")
        return output_path
    
    def process_inbox_pdf(self, pdf_path):
        """Verarbeite ein PDF aus der Inbox."""
        print(f"[INFO] Verarbeite: {pdf_path}")
        
        # Extrahiere Laborwerte
        extracted = self.extract_laborwerte_from_pdf(pdf_path)
        if not extracted:
            extracted = self.extract_laborwerte_from_pdf_plumber(pdf_path)
        
        if not extracted:
            # OCR für gescannte PDFs
            ocr_text = self.extract_from_scanned_pdf(pdf_path)
            print(f"[INFO] OCR-Text extrahiert ({len(ocr_text)} Zeichen)")
            return {'status': 'ocr', 'text': ocr_text}
        
        # Speichere in Datenbank
        for val in extracted:
            self.conn.execute(
                text("""
                    INSERT INTO laborwerte (parameter_name, wert, einheit, ermittlung_datum)
                    VALUES (:param, :wert, :einheit, :datum)
                """),
                {
                    'param': val['parameter'],
                    'wert': val['wert'],
                    'einheit': val.get('einheit', ''),
                    'datum': datetime.now().strftime('%Y-%m-%d %H:%M:%S')
                }
            )
        
        # Dokument in Datenbank eintragen
        file_hash = hashlib.md5(open(pdf_path, 'rb').read()).hexdigest()
        self.conn.execute(
            text("""
                INSERT OR REPLACE INTO dokumente (datei_name, dateipfad, datei_hash, status, kategorie, daten_typ, upload_datum)
                VALUES (:name, :pfad, :hash, :status, :kategorie, :datentyp, :datum)
            """),
            {
                'name': os.path.basename(pdf_path),
                'pfad': pdf_path,
                'hash': file_hash,
                'status': 'eingearbeitet',
                'kategorie': 'LABOR',
                'datentyp': 'pdf',
                'datum': datetime.now().strftime('%Y-%m-%d')
            }
        )
        self.conn.commit()
        
        # Verschiebe aus Inbox
        dest = os.path.join(ARCHIV_DIR, 'laborberichte', os.path.basename(pdf_path))
        os.makedirs(os.path.dirname(dest), exist_ok=True)
        os.rename(pdf_path, dest)
        
        return {
            'status': 'success',
            'extrahierte_werte': len(extracted),
            'werte': extracted[:5]  # Erst 5 Werte anzeigen
        }
    
    def close(self):
        """Verbinde Datenbank."""
        self.conn.close()
        self.engine.dispose()
    
    def __enter__(self):
        return self
    
    def __exit__(self, *args):
        self.close()


class HealthInboxHandler(FileSystemEventHandler):
    """Watchdog-Handler für automatische PDF-Verarbeitung."""
    
    def __init__(self, health_manager):
        super().__init__()
        self.manager = health_manager
    
    def on_created(self, event):
        if event.is_directory:
            return
        if event.src_path.endswith('.pdf'):
            print(f"[WATCHDOG] Neues PDF erkannt: {event.src_path}")
            try:
                result = self.manager.process_inbox_pdf(event.src_path)
                print(f"[WATCHDOG] Verarbeitung abgeschlossen: {result}")
            except Exception as e:
                print(f"[WATCHDOG] Fehler: {e}")


def start_watchdog():
    """Starte den Watchdog für die Health-Inbox."""
    manager = HealthManager()
    handler = HealthInboxHandler(manager)
    
    observer = Observer()
    observer.schedule(handler, INBOX_DIR, recursive=False)
    observer.start()
    print(f"[WATCHDOG] Beobachte: {INBOX_DIR}")
    
    try:
        while True:
            time.sleep(1)
    except KeyboardInterrupt:
        observer.stop()
    observer.join()


if __name__ == '__main__':
    import time
    
    if len(sys.argv) > 1 and sys.argv[1] == 'watchdog':
        start_watchdog()
    else:
        with HealthManager() as hm:
            # Demo: CRP-Trend
            crp = hm.compute_crp_trend()
            if crp:
                print(f"\n[INFO] CRP-Trend: {crp['trend']} (P={crp['p_value']:.4f})")
                print(f"Aktuell: {crp['aktuell']} mg/l")
            
            # Demo: Korrelationen
            corr = hm.compute_correlations()
            if corr:
                print(f"\n[INFO] Korrelationen: {json.dumps(corr, indent=2, ensure_ascii=False)[:500]}")
            
            # Demo: Diagramme
            crp_path = hm.generate_crp_chart()
            if crp_path:
                print(f"\n[INFO] CRP-Diagramm: {crp_path}")
            
            chol_path = hm.generate_cholesterol_chart()
            if chol_path:
                print(f"[INFO] Cholesterin-Diagramm: {chol_path}")
            
            heatmap_path = hm.generate_heatmap()
            if heatmap_path:
                print(f"[INFO] Heatmap: {heatmap_path}")
            
            dashboard_path = hm.generate_plotly_dashboard()
            if dashboard_path:
                print(f"[INFO] Plotly-Dashboard: {dashboard_path}")
            
            profile_path = hm.generate_profiling_report()
            if profile_path:
                print(f"[INFO] ydata-Profile: {profile_path}")
            
            print("\n[INFO] Gesundheitsmanager bereit.")
