File size: 7,574 Bytes
3872518
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Simple Knowledge Graph - Disease-Symptom-Drug Relationships

A lightweight knowledge graph for medical relationships without
requiring Neo4j or other graph databases.
"""
from typing import Dict, List
from loguru import logger


class MedicalKnowledgeGraph:
    """
    Simple in-memory knowledge graph for medical relationships.

    Tracks:
    - Diseases and their symptoms
    - Drugs and what they treat
    - Drug side effects
    - Contraindications
    """

    def __init__(self):
        """Initialize knowledge graph with common medical relationships"""

        # Disease -> Symptoms
        self.disease_symptoms = {
            "diabetes": ["increased thirst", "frequent urination", "fatigue", "blurred vision", "slow healing"],
            "hypertension": ["headache", "dizziness", "chest pain", "shortness of breath", "nosebleeds"],
            "asthma": ["wheezing", "shortness of breath", "chest tightness", "coughing"],
            "migraine": ["severe headache", "nausea", "sensitivity to light", "visual disturbances"],
            "flu": ["fever", "cough", "sore throat", "body aches", "fatigue"],
            "covid-19": ["fever", "cough", "loss of taste", "loss of smell", "fatigue", "shortness of breath"],
            "heart disease": ["chest pain", "shortness of breath", "fatigue", "irregular heartbeat"],
            "depression": ["sadness", "loss of interest", "fatigue", "sleep problems", "appetite changes"],
            "anxiety": ["worry", "restlessness", "rapid heartbeat", "sweating", "difficulty concentrating"]
        }

        # Drug -> Treats (diseases)
        self.drug_treats = {
            "metformin": ["diabetes", "prediabetes"],
            "lisinopril": ["hypertension", "heart disease"],
            "albuterol": ["asthma", "copd"],
            "sumatriptan": ["migraine"],
            "ibuprofen": ["pain", "inflammation", "fever"],
            "aspirin": ["pain", "fever", "heart disease prevention"],
            "sertraline": ["depression", "anxiety"],
            "atorvastatin": ["high cholesterol", "heart disease prevention"]
        }

        # Drug -> Side Effects
        self.drug_side_effects = {
            "metformin": ["nausea", "diarrhea", "stomach upset"],
            "lisinopril": ["dizziness", "dry cough", "fatigue"],
            "albuterol": ["tremor", "nervousness", "rapid heartbeat"],
            "sumatriptan": ["dizziness", "drowsiness", "tingling"],
            "ibuprofen": ["stomach upset", "heartburn", "dizziness"],
            "aspirin": ["stomach upset", "bleeding risk"],
            "sertraline": ["nausea", "insomnia", "drowsiness"],
            "atorvastatin": ["muscle pain", "liver problems"]
        }

        # Symptom -> Possible Diseases
        self.symptom_diseases = {}
        for disease, symptoms in self.disease_symptoms.items():
            for symptom in symptoms:
                if symptom not in self.symptom_diseases:
                    self.symptom_diseases[symptom] = []
                self.symptom_diseases[symptom].append(disease)

        logger.info(f"[KnowledgeGraph] Initialized with {len(self.disease_symptoms)} diseases, "
                   f"{len(self.drug_treats)} drugs")

    def get_disease_symptoms(self, disease: str) -> List[str]:
        """Get symptoms for a disease"""
        disease_lower = disease.lower()
        return self.disease_symptoms.get(disease_lower, [])

    def get_possible_diseases(self, symptoms: List[str]) -> Dict[str, int]:
        """
        Get possible diseases based on symptoms.

        Returns dict of disease -> symptom match count
        """
        disease_matches = {}

        for symptom in symptoms:
            symptom_lower = symptom.lower()
            for possible_disease in self.symptom_diseases.get(symptom_lower, []):
                disease_matches[possible_disease] = disease_matches.get(possible_disease, 0) + 1

        # Sort by match count
        return dict(sorted(disease_matches.items(), key=lambda x: x[1], reverse=True))

    def get_drug_info(self, drug: str) -> Dict:
        """Get comprehensive drug information"""
        drug_lower = drug.lower()

        return {
            "drug": drug,
            "treats": self.drug_treats.get(drug_lower, []),
            "side_effects": self.drug_side_effects.get(drug_lower, []),
            "found": drug_lower in self.drug_treats
        }

    def get_treatment_options(self, disease: str) -> List[str]:
        """Get drugs that treat a disease"""
        disease_lower = disease.lower()
        treatments = []

        for drug, treats in self.drug_treats.items():
            if disease_lower in treats:
                treatments.append(drug)

        return treatments

    def enhance_query_with_graph(self, query: str) -> str:
        """
        Enhance a query with knowledge graph context.

        Adds relevant medical relationships to improve RAG retrieval.
        """
        query_lower = query.lower()
        context_parts = []

        # Check for diseases mentioned
        for disease in self.disease_symptoms.keys():
            if disease in query_lower:
                symptoms = self.disease_symptoms[disease]
                context_parts.append(f"{disease} symptoms: {', '.join(symptoms[:5])}")

                # Add treatments
                treatments = self.get_treatment_options(disease)
                if treatments:
                    context_parts.append(f"{disease} treatments: {', '.join(treatments[:3])}")

        # Check for drugs mentioned
        for drug in self.drug_treats.keys():
            if drug in query_lower:
                info = self.get_drug_info(drug)
                if info["treats"]:
                    context_parts.append(f"{drug} treats: {', '.join(info['treats'])}")

        # Check for symptoms mentioned
        mentioned_symptoms = []
        for symptom in self.symptom_diseases.keys():
            if symptom in query_lower:
                mentioned_symptoms.append(symptom)

        if mentioned_symptoms:
            possible_diseases = self.get_possible_diseases(mentioned_symptoms)
            if possible_diseases:
                top_diseases = list(possible_diseases.keys())[:3]
                context_parts.append(f"Symptoms may indicate: {', '.join(top_diseases)}")

        # Combine context
        if context_parts:
            enhanced = f"{query}\n\nMedical context: {' | '.join(context_parts)}"
            logger.info(f"[KnowledgeGraph] Enhanced query with {len(context_parts)} context items")
            return enhanced

        return query

    def add_disease(self, disease: str, symptoms: List[str]):
        """Add or update a disease and its symptoms"""
        disease_lower = disease.lower()
        self.disease_symptoms[disease_lower] = symptoms

        # Update reverse index
        for symptom in symptoms:
            if symptom not in self.symptom_diseases:
                self.symptom_diseases[symptom] = []
            if disease_lower not in self.symptom_diseases[symptom]:
                self.symptom_diseases[symptom].append(disease_lower)

        logger.info(f"[KnowledgeGraph] Added/updated disease: {disease}")

    def add_drug(self, drug: str, treats: List[str], side_effects: List[str]):
        """Add or update a drug"""
        drug_lower = drug.lower()
        self.drug_treats[drug_lower] = treats
        self.drug_side_effects[drug_lower] = side_effects

        logger.info(f"[KnowledgeGraph] Added/updated drug: {drug}")


# Singleton instance
knowledge_graph = MedicalKnowledgeGraph()