| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384 |
- from functools import lru_cache
- from langchain_community.utilities import SQLDatabase
- from langchain_core.tools import tool
- from agent.config import load_config
- @lru_cache(maxsize=1)
- def get_database() -> SQLDatabase:
- config = load_config()
- return SQLDatabase.from_uri(config.database_url)
- def is_readonly_sql(sql: str) -> bool:
- cleaned_sql = sql.strip().rstrip(";").strip()
- lowered_sql = cleaned_sql.lower()
- readonly_prefixes = ("select", "show", "describe", "desc", "explain")
- forbidden_keywords = (
- "insert",
- "update",
- "delete",
- "drop",
- "alter",
- "create",
- "truncate",
- "replace",
- "grant",
- "revoke",
- )
- if not lowered_sql.startswith(readonly_prefixes):
- return False
- if ";" in cleaned_sql:
- return False
- return not any(f" {keyword} " in f" {lowered_sql} " for keyword in forbidden_keywords)
- @tool
- def list_tables() -> str:
- """
- 列出当前数据库中可以使用的所有表名。
- 当用户询问数据库里有哪些表、需要做数据分析、需要查询某个业务数据,
- 或者你不确定应该查询哪张表时,应该先调用这个工具。
- 这个工具不需要参数,只返回表名列表,不会查询表里的具体数据。
- """
- db = get_database()
- table_names = db.get_usable_table_names()
- return "\n".join(table_names) if table_names else "当前数据库中没有可用的数据表。"
- @tool
- def get_table_schema(table_name: str) -> str:
- """
- 查看指定数据表的字段结构、字段类型和部分样例信息。
- 当你准备生成 SQL 之前,应该先调用这个工具确认表结构,
- 不要凭空猜测字段名。参数 table_name 必须是数据库中真实存在的表名,
- 例如:chat_history。一次只传入一个表名。
- """
- db = get_database()
- return db.get_table_info([table_name])
- @tool
- def ask_database(sql: str) -> str:
- """
- 执行只读 SQL 查询并返回查询结果。
- 当你已经知道要查询的表和字段后,调用这个工具执行 SQL。
- 参数 sql 必须是一条完整的只读 SQL 语句,只允许 SELECT、SHOW、
- DESCRIBE、DESC、EXPLAIN。不要传入 INSERT、UPDATE、DELETE、DROP、
- ALTER、CREATE 等会修改数据库的语句。查询数据时建议使用 LIMIT 50
- 限制返回行数,避免一次返回太多数据。
- """
- if not is_readonly_sql(sql):
- return "拒绝执行:只允许单条只读 SQL,例如 SELECT、SHOW、DESCRIBE、EXPLAIN。"
- db = get_database()
- return db.run(sql)
|