(dbId: Optional[str], sqlGPT: Optional[SqlGPT] = None)
| 40 | |
| 41 | |
| 42 | def sql_chat(dbId: Optional[str], sqlGPT: Optional[SqlGPT] = None): |
| 43 | if sqlGPT is None: |
| 44 | sqlUrl = st.session_state["openai_model"] |
| 45 | sqlGPT = SqlGPT() |
| 46 | |
| 47 | if "openai_model" not in st.session_state: |
| 48 | st.session_state["openai_model"] = "gpt-3.5-turbo" |
| 49 | |
| 50 | if "messages" not in st.session_state: |
| 51 | # TODO 加载历史会话 |
| 52 | st.session_state.messages = [] |
| 53 | |
| 54 | for message in st.session_state.messages: |
| 55 | with st.chat_message(message["role"]): |
| 56 | st.markdown(message["content"]) |
| 57 | |
| 58 | if question := st.chat_input("What is up?"): |
| 59 | st.session_state.messages.append({"role": "user", "content": question}) |
| 60 | with st.chat_message("user"): |
| 61 | st.markdown(question) |
| 62 | |
| 63 | with st.chat_message("assistant"): |
| 64 | message_placeholder = st.empty() |
| 65 | full_response = sqlGPT.generateSQL(question=question) |
| 66 | |
| 67 | message_placeholder.markdown(full_response) |
| 68 | st.session_state.messages.append({"role": "assistant", "content": full_response}) |
| 69 | |
| 70 | |
| 71 | def main(): |
nothing calls this directly
no test coverage detected