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)