hekaos/Text-to-SQL
0
1"""2Streamlit UI for Recruiting Database Assistant3Chat-based interface for recruiters to query candidate database4"""5 6import streamlit as st7import pandas as pd8from datetime import datetime9from backend import get_agent, DatabaseManager10import json11 12# Page configuration13st.set_page_config(14 page_title="Recruiting Database Assistant",15 page_icon="๐ง ",16 layout="wide",17 initial_sidebar_state="expanded"18)19 20# Custom CSS for better styling21st.markdown("""22<style>23 .main-header {24 font-size: 2.5rem;25 font-weight: bold;26 color: #1f77b4;27 text-align: center;28 margin-bottom: 1rem;29 }30 .chat-message {31 padding: 1rem;32 border-radius: 0.5rem;33 margin-bottom: 1rem;34 border-left: 5px solid;35 }36 .user-message {37 background-color: #e3f2fd;38 border-left-color: #1976d2;39 }40 .agent-message {41 background-color: #f3e5f5;42 border-left-color: #7b1fa2;43 }44 .error-message {45 background-color: #ffebee;46 border-left-color: #c62828;47 }48 .sql-query {49 background-color: #f5f5f5;50 padding: 0.5rem;51 border-radius: 0.25rem;52 font-family: monospace;53 font-size: 0.9rem;54 margin: 0.5rem 0;55 }56 .stats-box {57 background-color: #e8f5e9;58 padding: 1rem;59 border-radius: 0.5rem;60 text-align: center;61 margin-bottom: 1rem;62 }63</style>64""", unsafe_allow_html=True)65 66# Initialize session state67if 'chat_history' not in st.session_state:68 st.session_state.chat_history = []69 70if 'query_count' not in st.session_state:71 st.session_state.query_count = 072 73if 'agent' not in st.session_state:74 try:75 st.session_state.agent = get_agent()76 st.session_state.agent_initialized = True77 except Exception as e:78 st.session_state.agent_initialized = False79 st.session_state.init_error = str(e)80 81 82def format_chat_message(message_type: str, content: str, data=None, sql_query=None, message_id=None):83 """Format a chat message with appropriate styling"""84 timestamp = datetime.now().strftime("%H:%M:%S")85 86 # Generate unique ID for this message if not provided87 if message_id is None:88 message_id = f"{message_type}_{timestamp}_{hash(content) % 10000}"89 90 if message_type == "user":91 st.markdown(f"""92 <div class="chat-message user-message">93 <strong>๐ค You</strong> <small>({timestamp})</small><br/>94 {content}95 </div>96 """, unsafe_allow_html=True)97 98 elif message_type == "agent":99 st.markdown(f"""100 <div class="chat-message agent-message">101 <strong>๐ค Agent</strong> <small>({timestamp})</small><br/>102 {content}103 </div>104 """, unsafe_allow_html=True)105 106 if sql_query:107 st.markdown(f"""108 <div class="sql-query">109 <strong>Generated SQL:</strong><br/>110 <code>{sql_query}</code>111 </div>112 """, unsafe_allow_html=True)113 114 if data and len(data) > 0:115 df = pd.DataFrame(data)116 st.dataframe(df, use_container_width=True)117 118 # Option to download results with unique key119 csv = df.to_csv(index=False)120 st.download_button(121 label="๐ฅ Download Results as CSV",122 data=csv,123 file_name=f"candidates_{datetime.now().strftime('%Y%m%d_%H%M%S')}.csv",124 mime="text/csv",125 key=f"download_{message_id}" # Unique key for each download button126 )127 128 elif message_type == "error":129 st.markdown(f"""130 <div class="chat-message error-message">131 <strong>โ ๏ธ Error</strong> <small>({timestamp})</small><br/>132 {content}133 </div>134 """, unsafe_allow_html=True)135 136 137def display_stats():138 """Display database statistics in the sidebar"""139 try:140 db = DatabaseManager()141 count_result = db.execute_query("SELECT COUNT(*) as total FROM candidates")142 143 if count_result.success:144 total_candidates = count_result.data[0]['total']145 146 st.sidebar.markdown(f"""147 <div class="stats-box">148 <h3 style="margin: 0; color: #2e7d32;">๐ Database Stats</h3>149 <h2 style="margin: 0.5rem 0; color: #1b5e20;">{total_candidates}</h2>150 <p style="margin: 0; color: #558b2f;">Total Candidates</p>151 </div>152 """, unsafe_allow_html=True)153 except Exception as e:154 st.sidebar.error(f"Could not fetch stats: {str(e)}")155 156 157def display_sample_queries():158 """Display sample queries in the sidebar"""159 st.sidebar.markdown("### ๐ก Sample Queries")160 161 sample_queries = [162 "Show me candidates with AWS experience",163 "Find candidates with more than 5 years of experience in Java",164 "List all candidates from Delhi",165 "Who knows Spring Boot?",166 "Find Python developers with AWS skills",167 "Show candidates based in Mumbai or Bangalore",168 "List all candidates with their emails"169 ]170 171 for query in sample_queries:172 if st.sidebar.button(query, key=f"sample_{query}", use_container_width=True):173 process_query(query)174 175 176def display_recent_queries():177 """Display recent queries in the sidebar"""178 if st.session_state.chat_history:179 st.sidebar.markdown("### ๐ Recent Queries")180 181 # Show last 5 queries182 recent = [msg for msg in st.session_state.chat_history if msg['type'] == 'user'][-5:]183 184 for i, msg in enumerate(reversed(recent)):185 st.sidebar.markdown(f"{i+1}. {msg['content'][:50]}...")186 187 188def process_query(query: str):189 """Process a user query through the agent"""190 if not query.strip():191 st.warning("Please enter a query")192 return193 194 # Add user message to chat history195 st.session_state.chat_history.append({196 'type': 'user',197 'content': query,198 'timestamp': datetime.now()199 })200 201 st.session_state.query_count += 1202 203 # Process query with agent204 with st.spinner("๐ค Thinking..."):205 try:206 result = st.session_state.agent.process_query(query)207 208 if result['success']:209 response_content = f"Found {len(result['data'])} candidate(s) matching your query."210 211 st.session_state.chat_history.append({212 'type': 'agent',213 'content': response_content,214 'data': result['data'],215 'sql_query': result['sql_query'],216 'timestamp': datetime.now()217 })218 else:219 error_msg = result.get('error') or result.get('message', 'Unknown error')220 st.session_state.chat_history.append({221 'type': 'error',222 'content': error_msg,223 'timestamp': datetime.now()224 })225 226 except Exception as e:227 st.session_state.chat_history.append({228 'type': 'error',229 'content': f"System error: {str(e)}",230 'timestamp': datetime.now()231 })232 233 234# Main UI235def main():236 # Header237 st.markdown('<h1 class="main-header">๐ง Recruiting Database Assistant</h1>', unsafe_allow_html=True)238 239 # Check if agent is initialized240 if not st.session_state.agent_initialized:241 st.error(f"โ ๏ธ Failed to initialize agent: {st.session_state.get('init_error', 'Unknown error')}")242 st.info("Please check your .env file configuration (OpenAI API key and MySQL credentials)")243 return244 245 # Sidebar246 with st.sidebar:247 st.title("๐ฏ Navigation")248 249 display_stats()250 251 st.markdown("---")252 253 display_sample_queries()254 255 st.markdown("---")256 257 display_recent_queries()258 259 st.markdown("---")260 261 if st.button("๐๏ธ Clear Chat History", use_container_width=True):262 st.session_state.chat_history = []263 st.session_state.query_count = 0264 st.rerun()265 266 st.markdown("---")267 st.markdown("### โน๏ธ About")268 st.markdown("""269 This AI agent helps recruiters query a candidate database using natural language.270 271 **Features:**272 - Natural language to SQL conversion273 - Safe query validation274 - Real-time results275 - Export to CSV276 """)277 278 # Main chat area279 st.markdown("### ๐ฌ Chat with the Agent")280 st.markdown("Ask questions about candidates in natural language")281 282 # Display chat history283 chat_container = st.container()284 with chat_container:285 for idx, message in enumerate(st.session_state.chat_history):286 # Generate unique message ID using index287 msg_id = f"msg_{idx}_{message['timestamp'].strftime('%H%M%S%f')}"288 289 if message['type'] == 'user':290 format_chat_message('user', message['content'], message_id=msg_id)291 elif message['type'] == 'agent':292 format_chat_message(293 'agent',294 message['content'],295 data=message.get('data'),296 sql_query=message.get('sql_query'),297 message_id=msg_id298 )299 elif message['type'] == 'error':300 format_chat_message('error', message['content'], message_id=msg_id)301 302 # Query input303 st.markdown("---")304 305 col1, col2 = st.columns([5, 1])306 307 with col1:308 query_input = st.text_input(309 "Your query:",310 placeholder="e.g., Show me candidates with AWS experience",311 label_visibility="collapsed",312 key="query_input"313 )314 315 with col2:316 submit_button = st.button("๐ Submit", use_container_width=True, type="primary")317 318 # Process query on button click or Enter319 if submit_button and query_input:320 process_query(query_input)321 st.rerun()322 323 # Footer stats324 st.markdown("---")325 col1, col2, col3 = st.columns(3)326 with col1:327 st.metric("Total Queries", st.session_state.query_count)328 with col2:329 st.metric("Chat Messages", len(st.session_state.chat_history))330 with col3:331 if st.session_state.chat_history:332 last_query = st.session_state.chat_history[-1]333 st.metric("Last Query", last_query['timestamp'].strftime("%H:%M:%S"))334 335 336if __name__ == "__main__":337 main()