93 lines
3.5 KiB
Python
93 lines
3.5 KiB
Python
import concurrent.futures
|
||
import json
|
||
import logging
|
||
|
||
from rest_framework.exceptions import ParseError
|
||
|
||
from apps.bi.models import Dataset
|
||
from apps.utils.sql import execute_raw_sql, format_sqldata
|
||
from apps.utils.tools import MyJSONEncoder
|
||
|
||
myLogger = logging.getLogger('log')
|
||
|
||
forbidden_keywords = ["UPDATE", "DELETE", "DROP", "TRUNCATE", "INSERT", "CREATE", "ALTER", "GRANT", "REVOKE", "EXEC", "EXECUTE"]
|
||
|
||
|
||
def check_sql_safe(sql: str):
|
||
"""检查sql安全性
|
||
"""
|
||
sql_upper = sql.upper()
|
||
# 将SQL按空格和分号分割成单词
|
||
words = [word for word in sql_upper.replace(';', ' ').split() if word]
|
||
for kw in forbidden_keywords:
|
||
# 检查关键字是否作为独立单词出现
|
||
if kw in words:
|
||
raise ParseError(f'sql查询有风险-{kw}')
|
||
return sql
|
||
|
||
def format_json_with_placeholders(json_str, **kwargs):
|
||
formatted_json = json_str
|
||
|
||
# 遍历关键字参数,将占位符替换为对应的值
|
||
for key, value in kwargs.items():
|
||
formatted_json = formatted_json.replace("{" + key + "}", json.dumps(value, cls=MyJSONEncoder))
|
||
|
||
# 格式化后的字符串依然是 JSON 字符串,没有使用 json.loads()
|
||
return formatted_json
|
||
|
||
|
||
def render_dataset_sql(dt: Dataset, xquery=None, *, is_test=False):
|
||
"""根据数据集配置和调用参数生成经过安全检查的只读 SQL。"""
|
||
query = dict(dt.default_param or {})
|
||
query.update(dict(dt.test_param or {}) if is_test else dict(xquery or {}))
|
||
if not dt.sql_query:
|
||
return ''
|
||
try:
|
||
return check_sql_safe(dt.sql_query.format(**query))
|
||
except KeyError as exc:
|
||
raise ParseError(f'需指定查询参数_{str(exc)}') from exc
|
||
|
||
|
||
def execute_rendered_dataset(dt: Dataset, full_sql: str, *, raise_exception=True):
|
||
"""执行已经渲染和校验的 SQL,返回可合并到数据集响应的结果。"""
|
||
results = {}
|
||
results2 = {}
|
||
can_cache = True
|
||
sql_list = [sql for sql in full_sql.strip(';').split(';') if sql.strip()]
|
||
if sql_list:
|
||
with concurrent.futures.ThreadPoolExecutor(max_workers=6) as executor:
|
||
futures = {
|
||
executor.submit(execute_raw_sql, sql): (f'ds{index}', sql)
|
||
for index, sql in enumerate(sql_list)
|
||
}
|
||
for future in concurrent.futures.as_completed(futures):
|
||
name, sql = futures[future]
|
||
try:
|
||
res = future.result()
|
||
results[name], results2[name] = format_sqldata(res[0], res[1])
|
||
except Exception as exc:
|
||
myLogger.error(f'bi查询异常:{str(exc)}-{dt.code}--{sql}')
|
||
if raise_exception:
|
||
raise ParseError(f'查询异常:{str(exc)}') from exc
|
||
results[name] = 'error: ' + str(exc)
|
||
can_cache = False
|
||
|
||
response_data = {'data': results, 'data2': results2}
|
||
if dt.echart_options and not dt.echart_options.startswith('function'):
|
||
for result in results.values():
|
||
if isinstance(result, str):
|
||
raise ParseError(result)
|
||
response_data['echart_options'] = format_json_with_placeholders(
|
||
dt.echart_options, **results
|
||
)
|
||
return response_data, can_cache
|
||
|
||
|
||
def exec_dataset(dt: Dataset, xquery=None):
|
||
"""执行数据集
|
||
返回 (sql语句, { rda})
|
||
"""
|
||
full_sql = render_dataset_sql(dt, xquery)
|
||
response_data, _ = execute_rendered_dataset(dt, full_sql)
|
||
return full_sql, response_data
|