from __future__ import annotations from dataclasses import dataclass from inference.forward_chain import ForwardChainingEngine, ForwardChainResult from inference.backward_chain import BackwardChainingEngine, DifferentialResult from inference.symptom_extractor import extract_facts @dataclass class DiagnosisResult: primary_diagnosis: str dosha: str confidence: float forward_trace: list[str] backward_trace: list[str] full_proof_tree: dict all_facts: set[str] fired_rules: list[dict] llm_context: str class VrikshayurvedaInferenceEngine: """ Unified entry point for JAIM's inference system. Replaces the old keyword-based symbolic rule engine entirely. Call diagnose() with the same inputs the old engine received. It returns DiagnosisResult which includes llm_context — a compact paragraph ready to be injected into the Llama prompt. """ def __init__(self, rules_path: str = "vrikshayurveda_rules.yaml") -> None: self.forward_engine = ForwardChainingEngine(rules_path) self.backward_engine = BackwardChainingEngine(rules_path) def diagnose( self, user_query: str, expanded_queries: list[str], retrieved_chunks: list[str], ) -> DiagnosisResult: # Step 1: Extract atomic facts from all text sources facts = extract_facts(user_query, expanded_queries, retrieved_chunks) # Step 2: Forward chain — enrich working memory forward_result: ForwardChainResult = self.forward_engine.run(facts) # Step 3: Backward chain — differential diagnosis diff_result: DifferentialResult = self.backward_engine.differential_diagnosis( forward_result.final_facts ) # Step 4: Determine primary diagnosis label primary_dosha = diff_result.primary confidence = diff_result.primary_confidence dosha_map = { "vata": "Vata disorder", "pitta": "Pitta imbalance", "kapha": "Kapha obstruction", } primary_diagnosis = dosha_map.get(primary_dosha, "Undetermined disorder") # Step 5: Extract remedy recommendations from final facts remedies = [ f.replace("recommend_", "").replace("_", " ") for f in forward_result.final_facts if f.startswith("recommend_") ] # Step 6: Build dosha ranking string ranking_str = ", ".join( f"{r['dosha']} ({round(r['confidence'] * 100)}%)" for r in diff_result.rankings ) # Step 7: Build llm_context paragraph for Llama prompt injection chain_str = " → ".join(forward_result.proof_trace) if forward_result.proof_trace else "No rules fired" remedy_str = ", ".join(remedies) if remedies else "none identified" llm_context = ( f"Inference engine diagnosis: {primary_diagnosis} " f"({round(confidence * 100)}% confidence). " f"Reasoning chain: {chain_str}. " f"Differential dosha ranking: {ranking_str}. " f"Recommended Ayurvedic treatments: {remedy_str}." ) return DiagnosisResult( primary_diagnosis=primary_diagnosis, dosha=primary_dosha, confidence=confidence, forward_trace=forward_result.proof_trace, backward_trace=diff_result.trace, full_proof_tree={"differential": [r["proof_tree"] for r in diff_result.rankings]}, all_facts=forward_result.final_facts, fired_rules=forward_result.fired_rules, llm_context=llm_context, )