feat(ai): add guarded SQL auto execution

This commit is contained in:
t8y2 2026-05-11 02:27:56 +08:00
parent 9355fa7e93
commit e03ce491af
6 changed files with 242 additions and 3 deletions

View File

@ -36,6 +36,7 @@ import { isTauriRuntime } from "@/lib/tauriRuntime";
import { isCloseTabShortcut, isExecuteSqlShortcut } from "@/lib/keyboardShortcuts";
import { isPreviewTab } from "@/lib/tabPresentation";
import { SQL_FILE_UNSUPPORTED_TYPES } from "@/lib/databaseCapabilities";
import { classifyAiSqlExecution } from "@/lib/aiSqlExecutionPolicy";
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
@ -119,9 +120,16 @@ const executableSql = computed(() => {
: "";
});
const { dangerSql, showDangerDialog, tryExecute, cancelActiveExecution, tryExplain, onDangerConfirm } = useSqlExecution(
{ activeTab, activeConnection, executableSql, activeOutputView },
);
const {
dangerSql,
pendingDangerSql,
showDangerDialog,
tryExecute,
doExecute,
cancelActiveExecution,
tryExplain,
onDangerConfirm,
} = useSqlExecution({ activeTab, activeConnection, executableSql, activeOutputView });
const dialogs = useDialogSources();
const { getDatabaseOptions } = useDatabaseOptions();
@ -379,6 +387,28 @@ function onAiExecuteSql(sql: string) {
nextTick(() => tryExecute(sql));
}
function onAiRequestAutoExecuteSql(sql: string) {
const tabId = ensureQueryTab();
queryStore.updateSql(tabId, sql);
selectedSql.value = "";
const decision = classifyAiSqlExecution(sql, activeConnection.value);
if (decision.action === "block") {
toast(t("ai.autoSqlBlocked"), 5000);
return;
}
nextTick(() => {
if (decision.action === "auto_execute") {
void doExecute(sql);
return;
}
dangerSql.value = sql;
pendingDangerSql.value = sql;
showDangerDialog.value = true;
});
}
function handleKeydown(e: KeyboardEvent) {
if (isCloseTabShortcut(e)) {
e.preventDefault();
@ -647,6 +677,7 @@ onUnmounted(() => {
:connection="activeConnection"
@replace-sql="onAiReplaceSql"
@execute-sql="onAiExecuteSql"
@request-auto-execute-sql="onAiRequestAutoExecuteSql"
@close="toggleAiPanel"
/>
</div>

View File

@ -46,6 +46,7 @@ import DatabaseIcon from "@/components/icons/DatabaseIcon.vue";
import { useQueryStore } from "@/stores/queryStore";
import { useToast } from "@/composables/useToast";
import { buildAiContext, runAiStream, type AiAction } from "@/lib/ai";
import { extractFirstSqlCodeBlock, shouldAttemptAiAutoExecute } from "@/lib/aiSqlExecutionPolicy";
import { Marked } from "marked";
import {
aiTestConnection,
@ -81,6 +82,7 @@ const props = defineProps<{
const emit = defineEmits<{
replaceSql: [sql: string];
executeSql: [sql: string];
requestAutoExecuteSql: [sql: string];
close: [];
}>();
@ -281,6 +283,8 @@ async function send() {
prompt.value = "";
scrollToBottom();
const requestedAction = activeAction.value;
const shouldAutoExecute = shouldAttemptAiAutoExecute(text, requestedAction);
isGenerating.value = true;
messages.value.push({ role: "assistant", content: "" });
const assistantIdx = messages.value.length - 1;
@ -314,6 +318,10 @@ async function send() {
const msg = messages.value[assistantIdx];
if (msg) msg.isThinking = false;
isGenerating.value = false;
if (shouldAutoExecute) {
const sql = extractFirstSqlCodeBlock(msg?.content || "");
if (sql) emit("requestAutoExecuteSql", sql);
}
activeAction.value = "generate";
currentSessionId.value = "";
persistConversation();

View File

@ -462,6 +462,7 @@ export default {
fixWithAi: "Fix with AI",
truncated: "Context truncated",
contextSummary: "{database} · {tables} tables",
autoSqlBlocked: "The AI-generated SQL looked too risky to auto-execute. Review it manually before running.",
settingsHint:
"The config is stored in the local app data directory. Requests are sent by the Tauri backend instead of directly from the frontend.",
actions: {

View File

@ -453,6 +453,7 @@ export default {
fixWithAi: "用 AI 修复",
truncated: "上下文已截断",
contextSummary: "{database} · {tables} 张表",
autoSqlBlocked: "AI 生成的 SQL 风险较高,已阻止自动执行,请手动检查后再运行。",
settingsHint: "配置会保存在本机应用数据目录中。请求由 Tauri 后端发出,避免在前端直接暴露给模型服务。",
actions: {
generate: "生成 SQL",

View File

@ -0,0 +1,131 @@
import type { ConnectionConfig } from "@/types/database";
export type ConnectionEnvironment = "production" | "non_production" | "unknown";
export type AiSqlExecutionAction = "auto_execute" | "confirm" | "block";
export type AiSqlExecutionCategory = "read" | "low_risk_write" | "write" | "schema_change" | "dangerous" | "unknown";
export interface AiSqlExecutionDecision {
action: AiSqlExecutionAction;
environment: ConnectionEnvironment;
category: AiSqlExecutionCategory;
reasons: string[];
}
const READ_RE = /^(SELECT|WITH|SHOW|DESCRIBE|DESC|EXPLAIN)\b/i;
const INSERT_RE = /^INSERT\b/i;
const UPDATE_RE = /^UPDATE\b/i;
const DELETE_RE = /^DELETE\b/i;
const CONFIRM_WRITE_RE = /^(MERGE|REPLACE)\b/i;
const BLOCK_RE = /^(DROP|TRUNCATE|ALTER|RENAME)\b/i;
const SCHEMA_RE = /^(CREATE)\b/i;
const PRODUCTION_RE = /\b(prod|prd|production)\b|生产|正式/i;
const NON_PRODUCTION_RE =
/\b(local|localhost|dev|develop|development|test|testing|stage|staging|sandbox|demo)\b|本地|开发|测试|预发/i;
const LOCAL_HOST_RE = /^(localhost|127(?:\.\d{1,3}){3}|0\.0\.0\.0|::1)$/i;
const NEGATIVE_EXECUTION_RE =
/(不要|别|不用|禁止|只生成|仅生成|只写|仅写).{0,12}(执行|运行)|do\s+not\s+execute|don't\s+execute|dont\s+execute|without\s+executing|only\s+(generate|write|return)/i;
const ACTION_INTENT_RE =
/(查一下|查询|查看|看一下|统计|列出|执行|运行|跑一下|show|list|query|find|count|get|fetch|execute|run)/i;
export function stripAiSqlComments(sql: string): string {
return sql
.replace(/\/\*[\s\S]*?\*\//g, " ")
.replace(/--.*$/gm, " ")
.replace(/#.*$/gm, " ");
}
function sqlStatements(sql: string): string[] {
return stripAiSqlComments(sql)
.split(";")
.map((stmt) => stmt.trim())
.filter(Boolean);
}
function classifyStatement(statement: string): AiSqlExecutionCategory {
if (READ_RE.test(statement)) return "read";
if (BLOCK_RE.test(statement)) return "dangerous";
if (SCHEMA_RE.test(statement)) return "schema_change";
if (INSERT_RE.test(statement)) return "low_risk_write";
if (UPDATE_RE.test(statement)) return isScopedUpdate(statement) ? "low_risk_write" : "dangerous";
if (DELETE_RE.test(statement) || CONFIRM_WRITE_RE.test(statement)) return "write";
return "unknown";
}
function isScopedUpdate(statement: string): boolean {
const whereMatch = statement.match(/\bWHERE\b([\s\S]*)$/i);
if (!whereMatch) return false;
const where = whereMatch[1];
if (/\b1\s*=\s*1\b|\btrue\b/i.test(where)) return false;
return /\b[\w"`.[\]]*(?:id|_id|uuid|key)[\w"`.[\]]*\s*=\s*(?:'[^']+'|"[^"]+"|`[^`]+`|[\w.-]+)/i.test(where);
}
export function classifyConnectionEnvironment(connection?: ConnectionConfig): ConnectionEnvironment {
if (!connection) return "unknown";
const parts = [connection.name, connection.host, connection.database, connection.connection_string].filter(Boolean);
const signal = parts.join(" ");
if (PRODUCTION_RE.test(signal)) return "production";
if (LOCAL_HOST_RE.test(connection.host) || NON_PRODUCTION_RE.test(signal)) return "non_production";
return "unknown";
}
export function classifyAiSqlExecution(sql: string, connection?: ConnectionConfig): AiSqlExecutionDecision {
const environment = classifyConnectionEnvironment(connection);
const statements = sqlStatements(sql);
const reasons: string[] = [];
if (!statements.length) {
return { action: "block", environment, category: "unknown", reasons: ["empty_sql"] };
}
const categories = statements.map(classifyStatement);
const hasMultipleStatements = statements.length > 1;
if (hasMultipleStatements) reasons.push("multi_statement");
if (categories.includes("dangerous")) {
return { action: "block", environment, category: "dangerous", reasons };
}
if (categories.includes("unknown")) {
return { action: "confirm", environment, category: "unknown", reasons };
}
if (categories.every((category) => category === "read")) {
return { action: "auto_execute", environment, category: "read", reasons };
}
if (hasMultipleStatements) {
return { action: "confirm", environment, category: "write", reasons };
}
const [category] = categories;
if (category === "low_risk_write") {
return {
action: environment === "non_production" ? "auto_execute" : "confirm",
environment,
category,
reasons,
};
}
return {
action: "confirm",
environment,
category,
reasons,
};
}
export function shouldAttemptAiAutoExecute(instruction: string, action: string): boolean {
if (action !== "generate") return false;
const normalized = instruction.trim();
if (!normalized || NEGATIVE_EXECUTION_RE.test(normalized)) return false;
return ACTION_INTENT_RE.test(normalized);
}
export function extractFirstSqlCodeBlock(content: string): string | undefined {
const match = content.match(/```(?:sql|mysql|postgresql|sqlite|tsql|clickhouse)?\s*\n([\s\S]*?)```/i);
const sql = match?.[1]?.trim();
return sql || undefined;
}

View File

@ -0,0 +1,67 @@
import { strict as assert } from "node:assert";
import test from "node:test";
import {
classifyAiSqlExecution,
classifyConnectionEnvironment,
shouldAttemptAiAutoExecute,
} from "../src/lib/aiSqlExecutionPolicy.ts";
import type { ConnectionConfig } from "../src/types/database.ts";
function conn(overrides: Partial<ConnectionConfig> = {}): ConnectionConfig {
return {
id: "c1",
name: "local-pg",
db_type: "postgres",
host: "127.0.0.1",
port: 5432,
username: "postgres",
password: "",
database: "app_dev",
...overrides,
};
}
test("classifyConnectionEnvironment treats local and dev targets as non-production", () => {
assert.equal(classifyConnectionEnvironment(conn()), "non_production");
assert.equal(classifyConnectionEnvironment(conn({ name: "staging-db", host: "10.0.0.8" })), "non_production");
});
test("classifyConnectionEnvironment treats production signals and unknown targets as production-like", () => {
assert.equal(classifyConnectionEnvironment(conn({ name: "prod-db", host: "10.0.0.9" })), "production");
assert.equal(classifyConnectionEnvironment(conn({ name: "analytics", host: "10.0.0.9", database: "warehouse" })), "unknown");
});
test("read SQL auto-executes on production and non-production", () => {
assert.equal(classifyAiSqlExecution("SELECT * FROM users", conn()).action, "auto_execute");
assert.equal(classifyAiSqlExecution("SHOW TABLES", conn({ name: "prod-db" })).action, "auto_execute");
});
test("single insert auto-executes only on non-production targets", () => {
assert.equal(classifyAiSqlExecution("INSERT INTO users(name) VALUES ('a')", conn()).action, "auto_execute");
assert.equal(classifyAiSqlExecution("INSERT INTO users(name) VALUES ('a')", conn({ name: "prod-db" })).action, "confirm");
});
test("scoped single update auto-executes only on non-production targets", () => {
const sql = "UPDATE users SET name = 'a' WHERE id = 1";
assert.equal(classifyAiSqlExecution(sql, conn()).action, "auto_execute");
assert.equal(classifyAiSqlExecution(sql, conn({ name: "prod-db" })).action, "confirm");
});
test("broad or destructive writes do not auto-execute", () => {
assert.equal(classifyAiSqlExecution("UPDATE users SET name = 'a'", conn()).action, "block");
assert.equal(classifyAiSqlExecution("UPDATE users SET name = 'a' WHERE 1=1", conn()).action, "block");
assert.equal(classifyAiSqlExecution("DELETE FROM users WHERE id = 1", conn()).action, "confirm");
assert.equal(classifyAiSqlExecution("DROP TABLE users", conn()).action, "block");
});
test("comments and multi-statement writes do not bypass policy", () => {
assert.equal(classifyAiSqlExecution("-- SELECT\nDROP TABLE users", conn()).action, "block");
assert.equal(classifyAiSqlExecution("INSERT INTO users(name) VALUES ('a'); UPDATE users SET name='b' WHERE id=1", conn()).action, "confirm");
});
test("AI auto-execution only attempts action-oriented generate requests", () => {
assert.equal(shouldAttemptAiAutoExecute("查一下用户数量", "generate"), true);
assert.equal(shouldAttemptAiAutoExecute("show me recent orders", "generate"), true);
assert.equal(shouldAttemptAiAutoExecute("只生成 SQL不要执行", "generate"), false);
assert.equal(shouldAttemptAiAutoExecute("优化这条 SQL", "optimize"), false);
});