Repository navigation
Expand file tree
/
Copy pathapp.py
More file actions
171 lines (142 loc) · 5.71 KB
/
Copy pathapp.py
File metadata and controls
171 lines (142 loc) · 5.71 KB
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
import streamlit as st
import sqlite3
import pandas as pd
from langchain_community.utilities import SQLDatabase
from langchain_groq import ChatGroq
from langchain_experimental.sql import SQLDatabaseChain
from langchain_core.prompts import PromptTemplate
st.set_page_config(
page_title="Medic Query",
layout="wide",
)
DB_PATH = "hospital_warehouse.db"
EXAMPLE_QUERIES = [
"Show total treatment cost by department",
"List top 5 patients by length of stay",
"How many patients are in each status category?",
"What is the average rating per department?",
]
# --- API Key from Streamlit secrets ---
try:
api_key = st.secrets["GROQ_API_KEY"]
except (KeyError, FileNotFoundError):
st.error(
"GROQ_API_KEY not found. "
"Add it to `.streamlit/secrets.toml` for local development, "
"or configure it in Streamlit Cloud app settings."
)
st.stop()
# --- Cached resources (loaded once per session) ---
@st.cache_resource
def load_db():
return SQLDatabase.from_uri(f"sqlite:///{DB_PATH}")
@st.cache_resource
def load_chain():
db = load_db()
llm = ChatGroq(api_key=api_key, model="openai/gpt-oss-120b")
prompt = PromptTemplate(
input_variables=["input", "table_info"],
template=(
"You are an expert data analyst. Given the following hospital database schema, "
"write a correct SQL query that answers the user's question. "
"Return ONLY the SQL query, no explanations.\n\n"
"Database Schema:\n{table_info}\n\n"
"User Question:\n{input}\n\n"
"SQL Query:"
),
)
return SQLDatabaseChain.from_llm(llm, db, prompt=prompt, verbose=False, return_sql=True)
# --- Sidebar: database stats ---
def get_db_stats():
conn = sqlite3.connect(DB_PATH)
stats = {
"Total Patients": conn.execute("SELECT COUNT(*) FROM dim_patient").fetchone()[0],
"Departments": conn.execute("SELECT COUNT(*) FROM dim_dept").fetchone()[0],
"Staff Members": conn.execute("SELECT COUNT(*) FROM dim_staff").fetchone()[0],
"Treatment Records": conn.execute("SELECT COUNT(*) FROM fact_treatment").fetchone()[0],
}
conn.close()
return stats
with st.sidebar:
st.header("Database Overview")
for label, value in get_db_stats().items():
st.metric(label, f"{value:,}")
st.divider()
st.caption("Source: hospital_warehouse.db")
# --- Main UI ---
st.title("Medic Query")
st.markdown(
"Ask questions about hospital data in plain English. "
"The app generates SQL from your question and returns results from the hospital database."
)
with st.expander("How it works"):
st.markdown(
"""
**Pipeline:** Natural Language → SQL → Database → Results
1. You type a question in plain English
2. **Groq (Llama 3.1 8B)** generates the SQL query via LangChain
3. The SQL runs against the local **SQLite** hospital database
4. Results are displayed as an interactive table
**Tech Stack:** Streamlit · LangChain · Groq (Llama 3.1 8B) · SQLite · Python
**Database Tables:**
- `dim_patient` — patient demographics
- `dim_staff` — staff records
- `dim_dept` — hospital departments
- `dim_bed` — bed assignments
- `fact_treatment` — treatment records (cost, LOS, rating, feedback)
"""
)
# --- Example query buttons ---
st.subheader("Try an example")
if "query_input" not in st.session_state:
st.session_state.query_input = ""
cols = st.columns(len(EXAMPLE_QUERIES))
for col, example in zip(cols, EXAMPLE_QUERIES):
if col.button(example, use_container_width=True):
st.session_state.query_input = example
st.rerun()
# --- Query input ---
st.subheader("Ask your question")
user_query = st.text_area(
"Question",
value=st.session_state.query_input,
placeholder="e.g., Show total treatment cost by department",
height=80,
label_visibility="collapsed",
)
if st.button("Run Query", type="primary"):
if not user_query.strip():
st.warning("Please enter a question before running.")
else:
with st.spinner("Generating SQL and fetching results..."):
try:
chain = load_chain()
response = chain.run(user_query)
# Extract SQL from LLM response
if "SELECT" in response.upper():
sql_start = response.upper().find("SELECT")
sql_query = response[sql_start:].strip().rstrip(";")
else:
sql_query = response.strip()
st.markdown("**Generated SQL**")
st.code(sql_query, language="sql")
conn = sqlite3.connect(DB_PATH)
try:
cursor = conn.execute(sql_query)
rows = cursor.fetchall()
columns = [desc[0] for desc in cursor.description]
conn.close()
st.markdown("**Results**")
if rows:
df = pd.DataFrame(rows, columns=columns)
st.dataframe(df, use_container_width=True)
st.caption(f"{len(rows)} row(s) returned")
else:
st.info("The query ran successfully but returned no results.")
except Exception as sql_err:
conn.close()
st.error(f"SQL execution failed: {sql_err}")
st.info("Try rephrasing your question or use one of the examples above.")
except Exception as e:
st.error(f"Failed to process your question: {e}")
st.info("Try rephrasing your question or use one of the examples above.")