Initial commit
This commit is contained in:
@@ -0,0 +1,113 @@
|
||||
import os
|
||||
from typing import List, Dict, Any
|
||||
from langchain_deepseek import ChatDeepSeek
|
||||
from langchain.schema import HumanMessage, SystemMessage
|
||||
|
||||
|
||||
class SimpleSQLGenerator:
|
||||
"""简化的SQL生成器"""
|
||||
|
||||
def __init__(self, api_key: str = None):
|
||||
self.llm = ChatDeepSeek(
|
||||
model="deepseek-chat",
|
||||
temperature=0,
|
||||
api_key=api_key or os.getenv("DEEPSEEK_API_KEY")
|
||||
)
|
||||
|
||||
def generate_sql(self, user_query: str, knowledge_results: List[Dict[str, Any]]) -> str:
|
||||
"""生成SQL语句"""
|
||||
# 构建上下文
|
||||
context = self._build_context(knowledge_results)
|
||||
|
||||
# 构建提示
|
||||
prompt = f"""你是一个SQL专家。请根据以下信息将用户问题转换为SQL查询语句。
|
||||
|
||||
数据库信息:
|
||||
{context}
|
||||
|
||||
用户问题:{user_query}
|
||||
|
||||
要求:
|
||||
1. 只返回SQL语句,不要包含任何解释
|
||||
2. 确保SQL语法正确
|
||||
3. 使用上下文中提供的表名和字段名
|
||||
4. 如果需要JOIN,请根据表结构进行合理关联
|
||||
|
||||
SQL语句:"""
|
||||
|
||||
messages = [HumanMessage(content=prompt)]
|
||||
response = self.llm.invoke(messages)
|
||||
|
||||
# 清理SQL语句
|
||||
sql = response.content.strip()
|
||||
if sql.startswith("```sql"):
|
||||
sql = sql[6:]
|
||||
if sql.startswith("```"):
|
||||
sql = sql[3:]
|
||||
if sql.endswith("```"):
|
||||
sql = sql[:-3]
|
||||
|
||||
return sql.strip()
|
||||
|
||||
def fix_sql(self, original_sql: str, error_message: str, knowledge_results: List[Dict[str, Any]]) -> str:
|
||||
"""修复SQL语句"""
|
||||
context = self._build_context(knowledge_results)
|
||||
|
||||
prompt = f"""请修复以下SQL语句的错误。
|
||||
|
||||
数据库信息:
|
||||
{context}
|
||||
|
||||
原始SQL:
|
||||
{original_sql}
|
||||
|
||||
错误信息:
|
||||
{error_message}
|
||||
|
||||
请返回修复后的SQL语句(只返回SQL,不要解释):"""
|
||||
|
||||
messages = [HumanMessage(content=prompt)]
|
||||
response = self.llm.invoke(messages)
|
||||
|
||||
# 清理SQL语句
|
||||
fixed_sql = response.content.strip()
|
||||
if fixed_sql.startswith("```sql"):
|
||||
fixed_sql = fixed_sql[6:]
|
||||
if fixed_sql.startswith("```"):
|
||||
fixed_sql = fixed_sql[3:]
|
||||
if fixed_sql.endswith("```"):
|
||||
fixed_sql = fixed_sql[:-3]
|
||||
|
||||
return fixed_sql.strip()
|
||||
|
||||
def _build_context(self, knowledge_results: List[Dict[str, Any]]) -> str:
|
||||
"""构建上下文信息"""
|
||||
context = ""
|
||||
|
||||
# 按类型分组
|
||||
ddl_info = []
|
||||
qsql_examples = []
|
||||
descriptions = []
|
||||
|
||||
for result in knowledge_results:
|
||||
if result["type"] == "ddl":
|
||||
ddl_info.append(result["content"])
|
||||
elif result["type"] == "qsql":
|
||||
qsql_examples.append(result["content"])
|
||||
elif result["type"] == "description":
|
||||
descriptions.append(result["content"])
|
||||
|
||||
# 构建上下文
|
||||
if ddl_info:
|
||||
context += "=== 表结构信息 ===\n"
|
||||
context += "\n".join(ddl_info) + "\n\n"
|
||||
|
||||
if descriptions:
|
||||
context += "=== 表和字段描述 ===\n"
|
||||
context += "\n".join(descriptions) + "\n\n"
|
||||
|
||||
if qsql_examples:
|
||||
context += "=== 查询示例 ===\n"
|
||||
context += "\n".join(qsql_examples) + "\n\n"
|
||||
|
||||
return context
|
||||
Reference in New Issue
Block a user