Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
116 changes: 102 additions & 14 deletions apps/extraction-service/llm_extract.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,29 @@
import json
import os
import time
from pathlib import Path

from dotenv import load_dotenv
from pydantic import BaseModel
from google import genai
from google.genai import types

from groq import Groq # pip install groq

env_path = Path(__file__).resolve().parent / ".env"
load_dotenv(env_path)

GEMINI_API_KEY = os.getenv("GEMINI_API_KEY")
GROQ_API_KEY = os.getenv("GROQ_API_KEY") # get a free key at https://console.groq.com/keys

# Fallback model config
GROQ_MODEL = "openai/gpt-oss-120b" # llama-3.3-70b-versatile was deprecated by Groq in June 2026
GEMINI_MAX_RETRIES = 2 # retries before falling back
GEMINI_RETRY_DELAY_SECONDS = 2 # backoff between retries
GROQ_MAX_RETRIES = 2 # retries before falling back to Ollama
GROQ_RETRY_DELAY_SECONDS = 2

groq_client = Groq(api_key=GROQ_API_KEY) if GROQ_API_KEY else None


# ---------------- Discharge Summary schema ----------------
Expand Down Expand Up @@ -213,7 +226,8 @@ class ExtractionOutput(BaseModel):
Rules:
- If a field cannot be found, use null (or false/0 for booleans/numbers where structurally required) -- do not guess or fabricate values.
- Dates must be YYYY-MM-DD, times HH:MM (24-hour).
- Return both objects nested under top-level keys "discharge_summary" and "claim_part_b".
- Return ONLY valid JSON with both objects nested under top-level keys "discharge_summary" and "claim_part_b".
- Do not include any explanation, markdown formatting, or code fences -- raw JSON only.

CLINICAL DOCUMENT TEXT (discharge summary source):
\"\"\"
Expand All @@ -229,25 +243,99 @@ class ExtractionOutput(BaseModel):
client = genai.Client(api_key=GEMINI_API_KEY)


def _call_gemini(prompt: str) -> dict:
response = client.models.generate_content(
Comment on lines 243 to +247
model="gemini-3.6-flash",
contents=prompt,
config=types.GenerateContentConfig(
response_mime_type="application/json",
response_schema=ExtractionOutput,
temperature=0.0,
),
)
return json.loads(response.text)


def _strip_code_fences(text: str) -> str:
text = text.strip()
if text.startswith("```"):
# remove ```json ... ``` or ``` ... ```
text = text.split("```", 2)[1] if text.count("```") >= 2 else text
if text.lstrip().startswith("json"):
text = text.lstrip()[4:]
return text.strip().strip("`").strip()


def _call_groq(prompt: str) -> dict:
if groq_client is None:
raise RuntimeError("GROQ_API_KEY not set in .env")

response = groq_client.chat.completions.create(
model="llama-3.3-70b-versatileL",
messages=[{"role": "user", "content": prompt}],
temperature=0.0,
response_format={"type": "json_object"},
)
raw = response.choices[0].message.content
raw = _strip_code_fences(raw)
return json.loads(raw)


def run(clinical_text: str, claim_text: str) -> dict:
prompt = PROMPT_TEMPLATE.format(clinical_text=clinical_text, claim_text=claim_text)

try:
response = client.models.generate_content(
model="gemini-3.6-flash",
contents=prompt,
config=types.GenerateContentConfig(
response_mime_type="application/json",
response_schema=ExtractionOutput,
temperature=0.0,
),
result = None
gemini_error = None
groq_error = None

# --- Tier 1: Gemini ---
for attempt in range(1, GEMINI_MAX_RETRIES + 1):
try:
result = _call_gemini(prompt)
break
except Exception as e:
gemini_error = e
print(f"[Gemini] attempt {attempt}/{GEMINI_MAX_RETRIES} failed: {e}")
if attempt < GEMINI_MAX_RETRIES:
time.sleep(GEMINI_RETRY_DELAY_SECONDS * attempt)

# --- Tier 2: Groq ---
if result is None:
print(f"[Gemini] exhausted retries ({gemini_error}). Trying Groq '{GROQ_MODEL}'...")
for attempt in range(1, GROQ_MAX_RETRIES + 1):
try:
result = _call_groq(prompt)
break
except Exception as e:
groq_error = e
print(f"[Groq] attempt {attempt}/{GROQ_MAX_RETRIES} failed: {e}")
if attempt < GROQ_MAX_RETRIES:
time.sleep(GROQ_RETRY_DELAY_SECONDS * attempt)

# --- Both providers failed ---
if result is None:
raise RuntimeError(
f"All providers failed.\n"
f"Gemini error: {gemini_error}\n"
f"Groq error: {groq_error}"
Comment on lines +316 to +320
)
except Exception as e:
raise RuntimeError(f"Gemini API call failed: {e}") from e

result = json.loads(response.text)
# --- Validate against schema regardless of source ---
try:
validated = ExtractionOutput(**result)
result = validated.model_dump()
except Exception as e:
raise RuntimeError(f"Extraction output failed schema validation: {e}") from e

return {
"discharge_summary": result["discharge_summary"],
"claim_part_b": result["claim_part_b"],
}
}


if __name__ == "__main__":
# quick manual test
sample_clinical = "Patient Name: John Doe. Admitted 2026-01-01, discharged 2026-01-05..."
sample_claim = "Hospital: ABC Hospital. IP Reg No: 12345..."
out = run(sample_clinical, sample_claim)
print(json.dumps(out, indent=2))
29 changes: 28 additions & 1 deletion apps/extraction-service/rules_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,10 @@ def check_rule_03_diagnosis_match(clinical: Dict[str, Any], claim: Dict[str, Any
if not claim_icd:
return True, "Primary ICD-10 diagnosis code is completely missing from Claim Form Part B."

if clinical_icd and (clinical_icd != claim_icd):
if not clinical_icd:
return True, "Cannot verify diagnosis: Discharge Summary is missing a primary ICD-10 code to cross-check against the claim."

if clinical_icd != claim_icd:
return True, f"Diagnosis mismatch: Claim states ICD '{claim_icd}', but Discharge Summary specifies '{clinical_icd}'."

return False, ""
Expand All @@ -121,6 +124,12 @@ def check_rule_04_min_24hr_stay(clinical: Dict[str, Any], claim: Dict[str, Any])

duration_hours = (dis - adm).total_seconds() / 3600.0

# If duration is negative, the dates are inverted -- that's a data integrity
# problem already caught and reported by RULE_02_DATE_CHRONOLOGY. Don't stack
# a second, nonsensical "-X hours" finding for the same root cause.
if duration_hours < 0:
return False, ""

if duration_hours < 24.0 and "day" not in admission_type:
return True, f"Inpatient stay was {duration_hours:.1f} hours (< 24 hours) but not marked as Day Care."
except Exception:
Expand Down Expand Up @@ -161,6 +170,23 @@ def check_rule_08_hospital_id(clinical: Dict[str, Any], claim: Dict[str, Any]) -
return False, ""


def check_rule_18_bank_details(clinical: Dict[str, Any], claim: Dict[str, Any]) -> Tuple[bool, str]:
"""RULE_18: Ensures insured's bank account number and IFSC code are present for NEFT settlement."""
bank_details = claim.get("section_f_bank_details", {}) or {}
account_no = bank_details.get("account_number")
ifsc_code = bank_details.get("ifsc_code")
Comment on lines +175 to +177

missing = []
if not account_no or not str(account_no).strip():
missing.append("account number")
if not ifsc_code or not str(ifsc_code).strip():
missing.append("IFSC code")

if missing:
return True, f"Insured's bank {' and '.join(missing)} missing from Claim Form Part A Section F -- required for NEFT settlement."
return False, ""


def check_rule_16_arithmetic(clinical: Dict[str, Any], claim: Dict[str, Any]) -> Tuple[bool, str]:
"""RULE_16: Verifies itemized line items sum up to total_claimed_amount."""
finances = claim.get("section_e_financial_summary", {}) or {}
Expand Down Expand Up @@ -195,6 +221,7 @@ def check_rule_16_arithmetic(clinical: Dict[str, Any], claim: Dict[str, Any]) ->
("RULE_07_DOCTOR_REG_MISSING", check_rule_07_doctor_reg),
("RULE_08_HOSPITAL_ID_MISSING", check_rule_08_hospital_id),
("RULE_16_ARITHMETIC_TOTAL_MISMATCH", check_rule_16_arithmetic),
("RULE_18_BANK_DETAILS_MISSING", check_rule_18_bank_details),
]


Expand Down
Loading