dvwn/nl2sql-api
0
1# Path: frontend/components/chat.py2# Chat component for handling user interactions and displaying chat history.3import streamlit as st4import pandas as pd5from utils.api import process_userPrompt6from components.auth import save_chat_history7 8def render_chat_interface():9 st.title(":material/sql: NL2SQL Assistant")10 st.divider()11 12 if not st.session_state.messages:13 for _ in range(5):14 st.write("")15 _, center_col, _ = st.columns([1, 2, 1])16 with center_col:17 st.subheader(":material/calendar_view_month: Ask a question about your data (e.g., 'Which country has the highest revenue? Give country name and amount.')", anchor=False)18 19 for idx, message in enumerate(st.session_state.messages):20 role = message.get("role", "assistant")21 content = message.get("content", "")22 23 with st.chat_message(role):24 st.markdown(content)25 26 if message.get("dataframe") is not None:27 raw_data = message["dataframe"]28 29 if isinstance(raw_data, list):30 display_df = pd.DataFrame(raw_data)31 else:32 display_df = raw_data33 34 st.dataframe(display_df, use_container_width=True)35 # display_df = message["dataframe"]36 # st.dataframe(message["dataframe"], use_container_width=True)37 38 csv_data = display_df.to_csv().encode('utf-8')39 if st.download_button(40 label=":material/download: Download CSV",41 data=csv_data,42 file_name=f'query_results_{idx}.csv',43 mime='text/csv',44 key=f"download_{idx}"45 ):46 st.toast("The file has been downloaded!")47 48 if prompt := st.chat_input("Ask a question about yout data..."):49 # Append user message50 st.session_state.messages.append({"role": "user", "content": prompt})51 st.rerun()52 53 if st.session_state.messages and st.session_state.messages[-1]["role"] == "user":54 with st.chat_message("assistant"):55 with st.spinner("Analyzing schema metadata and generating execution context..."):56 payload = process_userPrompt(57 question=st.session_state.messages[-1]["content"],58 model_id=st.session_state.current_model59 )60 61 if payload["status"] == "error":62 st.error(f"Interrupted:\n{payload['error']}")63 response_text = f"Failed to compute response due to error: {payload['error']}"64 st.session_state.messages.append({"role": "assistant", "content": response_text})65 else:66 st.markdown(payload["answer"])67 st.code(payload["sql"], language="sql")68 69 display_df = payload["data"]70 if not display_df.empty:71 display_df = display_df.copy()72 display_df.index = range(1, len(display_df) + 1)73 display_df.index.name = "No."74 75 st.session_state.messages.append({76 "role": "assistant",77 "content": f"{payload['answer']}\n\n```sql\n{payload['sql']}\n```",78 "dataframe": display_df if not display_df.empty else None79 })80 81 if st.session_state.auth_stat != 'guest':82 save_chat_history(st.session_state.username, st.session_state.messages)83 84 st.rerun()