生成sql查询语句,并调用工具执行
(state: SQLQueryState, *, config: RunnableConfig)
| 7 | from db.db_schema import DB_SCHEMA |
| 8 | |
| 9 | async def generate_sql(state: SQLQueryState, *, config: RunnableConfig) -> dict: |
| 10 | """ |
| 11 | 生成sql查询语句,并调用工具执行 |
| 12 | """ |
| 13 | # 绑定工具,让模型可以调用 execute_sql_query |
| 14 | model = init_chat_model( |
| 15 | name="generate_sql", |
| 16 | model=Settings.app_settings.inference_model, |
| 17 | temperature=Settings.app_settings.temperature, |
| 18 | streaming=Settings.app_settings.streaming, |
| 19 | openai_api_base=Settings.app_settings.openai_api_base, |
| 20 | openai_api_key=Settings.app_settings.openai_api_key, |
| 21 | ) |
| 22 | |
| 23 | prompt = ChatPromptTemplate.from_messages([ |
| 24 | ("system", GENERATE_SQL_PROMPT), |
| 25 | MessagesPlaceholder("history"), |
| 26 | ]) |
| 27 | chain = prompt | model |
| 28 | response = await chain.ainvoke({ |
| 29 | "history": state.messages, |
| 30 | "DB_SCHEMA": DB_SCHEMA |
| 31 | }, config) |
| 32 | print(f"\033[92mUsing generate_sql: {response}\033[0m") # 绿色输出 |
| 33 | |
| 34 | # 提取纯净的 SQL,去除 Markdown 代码块标记 |
| 35 | sql = response.content.strip() |
| 36 | if sql.startswith("```sql"): |
| 37 | sql = sql[6:].strip() |
| 38 | elif sql.startswith("```"): |
| 39 | sql = sql[3:].strip() |
| 40 | if sql.endswith("```"): |
| 41 | sql = sql[:-3].strip() |
| 42 | return {"sql": sql} |
nothing calls this directly
no outgoing calls
no test coverage detected