feat(sql-audit): 新增 SQL 审计中心并完善事务追踪

- 新增脱敏审计存储、筛选、保留策略、完整性校验与 JSON/CSV 导出
- 覆盖查询编辑器、事务、导入同步、对象操作、AI/MCP 与 Web 运行时入口
- 完善事务完整语句日志、Windows 快捷键映射及 SQL 分析布局
- 补充多语言、健康状态、数据目录迁移和回归测试
This commit is contained in:
Syngnat
2026-07-12 15:24:15 +08:00
parent 0b33b3a683
commit 7c8e8d8dd3
89 changed files with 11421 additions and 194 deletions

View File

@@ -24,6 +24,7 @@ import (
redisbackend "GoNavi-Wails/internal/redis"
"GoNavi-Wails/internal/resultdiff"
"GoNavi-Wails/internal/secretstore"
"GoNavi-Wails/internal/sqlaudit"
syncbackend "GoNavi-Wails/internal/sync"
"GoNavi-Wails/shared/i18n"
"github.com/google/uuid"
@@ -68,20 +69,24 @@ type queryContext struct {
}
type managedSQLTransaction struct {
id string
execer db.StatementExecer
transactor db.TransactionExecer
cancel context.CancelFunc
config connection.ConnectionConfig
dbType string
commitSQL string
rollbackSQL string
createdAt time.Time
mu sync.Mutex
id string
execer db.StatementExecer
transactor db.TransactionExecer
cancel context.CancelFunc
config connection.ConnectionConfig
dbType string
boundaryMode string
commitSQL string
rollbackSQL string
createdAt time.Time
finished bool
}
// App struct
type App struct {
ctx context.Context
webRuntime bool
startedAt time.Time
dbCache map[string]cachedDatabase // Cache for DB connections
connectFailures map[string]cachedConnectFailure
@@ -94,11 +99,26 @@ type App struct {
allowApplicationQuit bool
applicationQuitPromptInFlight bool
queryMu sync.RWMutex
dataRootApplyMu sync.Mutex
configDir string
secretStore secretstore.SecretStore
runningQueries map[string]queryContext // queryID -> cancelFunc and start time
sqlTransactionMu sync.Mutex
sqlTransactions map[string]*managedSQLTransaction
sqlAuditMu sync.RWMutex
sqlAuditStore *sqlaudit.Store
sqlAuditStorePath string
sqlAuditRuntimeActive bool
sqlAuditSuspended bool
sqlAuditAppendMu sync.Mutex
sqlAuditHealthMu sync.RWMutex
sqlAuditHealth sqlAuditHealthState
sqlAuditHealthPath string
sqlAuditHealthRevision uint64
sqlAuditSuspensionDropped int64
sqlAuditSuspensionFirstAt int64
sqlAuditSuspensionLastAt int64
sqlAuditSuspensionLastError string
jvmPreviewTokenMu sync.Mutex
jvmPreviewTokens map[string]jvmPreviewConfirmationToken
jvmPreviewTokenTTL time.Duration
@@ -112,6 +132,15 @@ func NewApp() *App {
return NewAppWithSecretStore(secretstore.NewKeyringStore())
}
// NewWebApp creates the backend used by the authenticated browser server.
// The immutable runtime marker keeps desktop-only Wails APIs from being
// reached through the reflective Web RPC bridge.
func NewWebApp() *App {
app := NewApp()
app.webRuntime = true
return app
}
func NewAppWithSecretStore(store secretstore.SecretStore) *App {
if store == nil {
store = secretstore.NewUnavailableStore("secret store unavailable")
@@ -244,6 +273,7 @@ func (a *App) startup(ctx context.Context) {
if err := migrateLegacyWebKitStorageIfNeeded(a); err != nil {
logger.Warnf("迁移旧 WebKit 连接存储失败:%v", err)
}
a.activateSQLAudit()
if shouldInstallMacNativeWindowDiagnostics() {
installMacNativeWindowDiagnostics(logger.Path())
}
@@ -299,6 +329,7 @@ func (a *App) Shutdown() {
logger.Infof("应用开始关闭,准备释放资源")
a.stopConnectionKeepAliveLoop()
a.rollbackPendingSQLTransactionsOnShutdown()
a.closeSQLAuditStore()
a.mu.Lock()
defer a.mu.Unlock()
for _, dbInst := range a.dbCache {

View File

@@ -318,10 +318,40 @@ func isReadOnlyMongoCommand(query string) bool {
if _, blocked := mongoWriteCommands[commandKey]; blocked {
return false
}
if commandKey == "aggregate" && mongoAggregateHasWriteStage(doc) {
return false
}
_, allowed := mongoReadOnlyCommands[commandKey]
return allowed
}
func mongoAggregateHasWriteStage(doc map[string]interface{}) bool {
var pipeline interface{}
for key, value := range doc {
if strings.EqualFold(strings.TrimSpace(key), "pipeline") {
pipeline = value
break
}
}
stages, ok := pipeline.([]interface{})
if !ok {
return false
}
for _, rawStage := range stages {
stage, ok := rawStage.(map[string]interface{})
if !ok {
continue
}
for key := range stage {
switch strings.ToLower(strings.TrimSpace(key)) {
case "$out", "$merge":
return true
}
}
}
return false
}
func isReadOnlyMilvusCommand(query string) bool {
trimmed := strings.TrimSpace(query)
if !strings.HasPrefix(trimmed, "{") {

View File

@@ -3,6 +3,7 @@ package app
import (
"encoding/json"
"errors"
"fmt"
"io"
"os"
"os/exec"
@@ -124,6 +125,9 @@ func (a *App) SelectDataRootDirectory(currentDir string) connection.QueryResult
}
func (a *App) ApplyDataRootDirectory(directory string, migrate bool) connection.QueryResult {
a.dataRootApplyMu.Lock()
defer a.dataRootApplyMu.Unlock()
currentRoot := appdata.MustResolveActiveRoot()
targetRoot, err := appdata.ResolveRoot(directory)
if err != nil {
@@ -139,6 +143,25 @@ func (a *App) ApplyDataRootDirectory(directory string, migrate bool) connection.
}
}
// The audit database uses SQLite WAL journaling. Pause it and checkpoint/close
// every audit connection before copying or switching the data root so the
// migration never copies a live database or an incomplete WAL sidecar.
resumeSQLAudit, suspendErr := a.suspendSQLAudit()
if suspendErr != nil {
a.resumeSQLAudit(resumeSQLAudit)
return connection.QueryResult{
Success: false,
Message: dataRootErrorWithDetail(
a.appText,
"app.data_root.backend.error.migrate_directory_failed",
suspendErr.Error(),
suspendErr,
map[string]any{"entry": "audit"},
).Error(),
}
}
defer a.resumeSQLAudit(resumeSQLAudit)
if migrate {
if err := migrateDataRootContentsWithText(currentRoot, targetRoot, a.appText); err != nil {
return connection.QueryResult{Success: false, Message: err.Error()}
@@ -222,6 +245,12 @@ func migrateDataRootContentsWithText(sourceRoot string, targetRoot string, text
if _, excluded := dataRootMigrationExcludedEntries[name]; excluded {
continue
}
if name == "audit" {
// The SQLite audit store is copied as an exact directory snapshot
// after every other migration step succeeds. Merging it would leave
// stale WAL/SHM or health sidecars from an existing target root.
continue
}
sourcePath := filepath.Join(sourceRoot, name)
targetPath := filepath.Join(targetRoot, name)
info, err := entry.Info()
@@ -241,6 +270,84 @@ func migrateDataRootContentsWithText(sourceRoot string, targetRoot string, text
if err := rewriteMigratedDataRootStateWithText(targetRoot, text); err != nil {
return err
}
if err := replaceMigratedAuditDirectory(sourceRoot, targetRoot); err != nil {
return dataRootWrapError(text, "app.data_root.backend.error.migrate_directory_failed", err, map[string]any{"entry": "audit"})
}
return nil
}
func replaceMigratedAuditDirectory(sourceRoot string, targetRoot string) error {
sourceAudit := filepath.Join(sourceRoot, "audit")
targetAudit := filepath.Join(targetRoot, "audit")
sourceInfo, sourceErr := os.Stat(sourceAudit)
sourceExists := sourceErr == nil
if sourceErr != nil && !os.IsNotExist(sourceErr) {
return fmt.Errorf("inspect source audit directory: %w", sourceErr)
}
if sourceExists && !sourceInfo.IsDir() {
return errors.New("source audit path is not a directory")
}
stagePath := ""
if sourceExists {
var err error
stagePath, err = os.MkdirTemp(targetRoot, ".gonavi-audit-stage-")
if err != nil {
return fmt.Errorf("create audit migration stage: %w", err)
}
defer func() { _ = os.RemoveAll(stagePath) }()
if err := copyDir(sourceAudit, stagePath); err != nil {
return fmt.Errorf("stage audit directory: %w", err)
}
}
targetInfo, targetErr := os.Lstat(targetAudit)
targetExists := targetErr == nil
if targetErr != nil && !os.IsNotExist(targetErr) {
return fmt.Errorf("inspect target audit directory: %w", targetErr)
}
if targetExists && targetInfo == nil {
return errors.New("target audit path is unavailable")
}
backupPath := ""
if targetExists {
reserved, err := os.MkdirTemp(targetRoot, ".gonavi-audit-backup-")
if err != nil {
return fmt.Errorf("reserve audit migration backup: %w", err)
}
if err := os.Remove(reserved); err != nil {
return fmt.Errorf("prepare audit migration backup: %w", err)
}
backupPath = reserved
if err := os.Rename(targetAudit, backupPath); err != nil {
return fmt.Errorf("backup target audit directory: %w", err)
}
}
rollback := func(cause error) error {
if backupPath == "" {
return cause
}
if restoreErr := os.Rename(backupPath, targetAudit); restoreErr != nil {
return errors.Join(cause, fmt.Errorf("restore target audit directory: %w", restoreErr))
}
backupPath = ""
return cause
}
if sourceExists {
if err := os.Rename(stagePath, targetAudit); err != nil {
return rollback(fmt.Errorf("activate staged audit directory: %w", err))
}
stagePath = ""
}
if backupPath != "" {
if err := os.RemoveAll(backupPath); err != nil {
logger.Warnf("清理数据目录迁移的旧审计备份失败path=%s err=%v", backupPath, err)
}
}
return nil
}

View File

@@ -5,10 +5,38 @@ import (
"os"
"path/filepath"
"testing"
"time"
"GoNavi-Wails/internal/appdata"
"GoNavi-Wails/internal/connection"
)
func TestApplyDataRootDirectorySerializesConcurrentRequests(t *testing.T) {
app := NewApp()
app.dataRootApplyMu.Lock()
done := make(chan connection.QueryResult, 1)
go func() {
done <- app.ApplyDataRootDirectory(appdata.MustResolveActiveRoot(), false)
}()
select {
case <-done:
app.dataRootApplyMu.Unlock()
t.Fatal("ApplyDataRootDirectory bypassed the application-wide serialization lock")
case <-time.After(50 * time.Millisecond):
}
app.dataRootApplyMu.Unlock()
select {
case result := <-done:
if !result.Success {
t.Fatalf("serialized ApplyDataRootDirectory returned failure: %s", result.Message)
}
case <-time.After(5 * time.Second):
t.Fatal("serialized ApplyDataRootDirectory did not resume")
}
}
func TestMigrateDataRootContentsCopiesKnownFilesAndDirectories(t *testing.T) {
sourceRoot := t.TempDir()
targetRoot := filepath.Join(t.TempDir(), "gonavi-data")
@@ -46,6 +74,64 @@ func TestMigrateDataRootContentsCopiesKnownFilesAndDirectories(t *testing.T) {
}
}
func TestMigrateDataRootContentsReplacesAuditDirectoryWithoutStaleSidecars(t *testing.T) {
sourceRoot := t.TempDir()
targetRoot := t.TempDir()
sourceAudit := filepath.Join(sourceRoot, "audit")
targetAudit := filepath.Join(targetRoot, "audit")
if err := os.MkdirAll(sourceAudit, 0o700); err != nil {
t.Fatalf("create source audit directory: %v", err)
}
if err := os.MkdirAll(targetAudit, 0o700); err != nil {
t.Fatalf("create target audit directory: %v", err)
}
if err := os.WriteFile(filepath.Join(sourceAudit, "sql_audit.db"), []byte("source-db"), 0o600); err != nil {
t.Fatalf("write source audit database: %v", err)
}
for name, content := range map[string]string{
"sql_audit.db": "old-db",
"sql_audit.db-wal": "stale-wal",
"sql_audit.db-shm": "stale-shm",
"sql_audit_health.json": "stale-health",
} {
if err := os.WriteFile(filepath.Join(targetAudit, name), []byte(content), 0o600); err != nil {
t.Fatalf("write target audit artifact %s: %v", name, err)
}
}
if err := migrateDataRootContents(sourceRoot, targetRoot); err != nil {
t.Fatalf("migrateDataRootContents returned error: %v", err)
}
payload, err := os.ReadFile(filepath.Join(targetAudit, "sql_audit.db"))
if err != nil || string(payload) != "source-db" {
t.Fatalf("target audit database = %q err=%v, want exact source snapshot", payload, err)
}
for _, name := range []string{"sql_audit.db-wal", "sql_audit.db-shm", "sql_audit_health.json"} {
if _, err := os.Stat(filepath.Join(targetAudit, name)); !os.IsNotExist(err) {
t.Fatalf("stale target audit artifact %s survived exact replacement: %v", name, err)
}
}
}
func TestMigrateDataRootContentsRemovesTargetAuditWhenSourceHasNone(t *testing.T) {
sourceRoot := t.TempDir()
targetRoot := t.TempDir()
targetAudit := filepath.Join(targetRoot, "audit")
if err := os.MkdirAll(targetAudit, 0o700); err != nil {
t.Fatalf("create target audit directory: %v", err)
}
if err := os.WriteFile(filepath.Join(targetAudit, "sql_audit.db-wal"), []byte("stale"), 0o600); err != nil {
t.Fatalf("write stale target audit artifact: %v", err)
}
if err := migrateDataRootContents(sourceRoot, targetRoot); err != nil {
t.Fatalf("migrateDataRootContents returned error: %v", err)
}
if _, err := os.Stat(targetAudit); !os.IsNotExist(err) {
t.Fatalf("target audit directory survived despite absent source snapshot: %v", err)
}
}
func TestMigrateDataRootContentsCopiesSecurityUpdateStateAndRewritesBackupPaths(t *testing.T) {
sourceRoot := t.TempDir()
targetRoot := filepath.Join(t.TempDir(), "gonavi-data")

View File

@@ -11,6 +11,7 @@ import (
"GoNavi-Wails/internal/connection"
"GoNavi-Wails/internal/db"
"GoNavi-Wails/internal/logger"
"GoNavi-Wails/internal/sqlaudit"
"GoNavi-Wails/internal/utils"
"GoNavi-Wails/shared/i18n"
)
@@ -191,7 +192,9 @@ func (a *App) MongoDiscoverMembers(config connection.ConnectionConfig) connectio
}
}
func (a *App) CreateDatabase(config connection.ConnectionConfig, dbName string) connection.QueryResult {
func (a *App) CreateDatabase(config connection.ConnectionConfig, dbName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("CREATE DATABASE %s", strings.TrimSpace(dbName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
dbName = strings.TrimSpace(dbName)
if dbName == "" {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.database_name_required", nil)}
@@ -368,7 +371,9 @@ func resolveSchemaDDLTargetDatabaseWithText(config connection.ConnectionConfig,
return targetDbName, nil
}
func (a *App) CreateSchema(config connection.ConnectionConfig, dbName string, schemaName string) connection.QueryResult {
func (a *App) CreateSchema(config connection.ConnectionConfig, dbName string, schemaName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("CREATE SCHEMA %s", strings.TrimSpace(schemaName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
if err := ensureConnectionAllowsStructureEdit(config, "connection.backend.action.create_schema"); err != nil {
return connection.QueryResult{Success: false, Message: err.Error()}
}
@@ -396,7 +401,9 @@ func (a *App) CreateSchema(config connection.ConnectionConfig, dbName string, sc
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.schema_created", nil)}
}
func (a *App) RenameSchema(config connection.ConnectionConfig, dbName string, oldSchemaName string, newSchemaName string) connection.QueryResult {
func (a *App) RenameSchema(config connection.ConnectionConfig, dbName string, oldSchemaName string, newSchemaName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("ALTER SCHEMA %s RENAME TO %s", strings.TrimSpace(oldSchemaName), strings.TrimSpace(newSchemaName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
if err := ensureConnectionAllowsStructureEdit(config, "connection.backend.action.rename_schema"); err != nil {
return connection.QueryResult{Success: false, Message: err.Error()}
}
@@ -422,7 +429,9 @@ func (a *App) RenameSchema(config connection.ConnectionConfig, dbName string, ol
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.schema_renamed", nil)}
}
func (a *App) DropSchema(config connection.ConnectionConfig, dbName string, schemaName string) connection.QueryResult {
func (a *App) DropSchema(config connection.ConnectionConfig, dbName string, schemaName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("DROP SCHEMA %s", strings.TrimSpace(schemaName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
if err := ensureConnectionAllowsStructureEdit(config, "connection.backend.action.drop_schema"); err != nil {
return connection.QueryResult{Success: false, Message: err.Error()}
}
@@ -680,7 +689,9 @@ func buildRunConfigForDDL(config connection.ConnectionConfig, dbType string, dbN
return runConfig
}
func (a *App) RenameDatabase(config connection.ConnectionConfig, oldName string, newName string) connection.QueryResult {
func (a *App) RenameDatabase(config connection.ConnectionConfig, oldName string, newName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("ALTER DATABASE %s RENAME TO %s", strings.TrimSpace(oldName), strings.TrimSpace(newName))
defer a.beginSQLAuditUserAction(config, oldName, "object_editor", &auditSQL, &result)()
oldName = strings.TrimSpace(oldName)
newName = strings.TrimSpace(newName)
if oldName == "" || newName == "" {
@@ -727,7 +738,9 @@ func (a *App) RenameDatabase(config connection.ConnectionConfig, oldName string,
}
}
func (a *App) DropDatabase(config connection.ConnectionConfig, dbName string) connection.QueryResult {
func (a *App) DropDatabase(config connection.ConnectionConfig, dbName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("DROP DATABASE %s", strings.TrimSpace(dbName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
dbName = strings.TrimSpace(dbName)
if dbName == "" {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.database_name_required", nil)}
@@ -763,7 +776,9 @@ func (a *App) DropDatabase(config connection.ConnectionConfig, dbName string) co
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.database_dropped", nil)}
}
func (a *App) RenameTable(config connection.ConnectionConfig, dbName string, oldTableName string, newTableName string) connection.QueryResult {
func (a *App) RenameTable(config connection.ConnectionConfig, dbName string, oldTableName string, newTableName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("ALTER TABLE %s RENAME TO %s", strings.TrimSpace(oldTableName), strings.TrimSpace(newTableName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
oldTableName = strings.TrimSpace(oldTableName)
newTableName = strings.TrimSpace(newTableName)
if oldTableName == "" || newTableName == "" {
@@ -819,7 +834,9 @@ func (a *App) RenameTable(config connection.ConnectionConfig, dbName string, old
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.table_renamed", nil)}
}
func (a *App) DropTable(config connection.ConnectionConfig, dbName string, tableName string) connection.QueryResult {
func (a *App) DropTable(config connection.ConnectionConfig, dbName string, tableName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("DROP TABLE %s", strings.TrimSpace(tableName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
tableName = strings.TrimSpace(tableName)
if tableName == "" {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.table_name_required", nil)}
@@ -878,16 +895,86 @@ func (a *App) MySQLShowCreateTable(config connection.ConnectionConfig, dbName st
return a.DBShowCreateTable(config, dbName, tableName)
}
func (a *App) DBQuery(config connection.ConnectionConfig, dbName string, query string) connection.QueryResult {
return a.DBQueryWithCancel(config, dbName, query, "")
type dbQueryAuditOptions struct {
trackHistory bool
auditAll bool
auditWrites bool
source string
}
func (a *App) DBQueryWithCancel(config connection.ConnectionConfig, dbName string, query string, queryID string) (result connection.QueryResult) {
// DBQuery() 以及后台元数据读取会传空 queryID只记录 SQL 编辑器显式传入 ID 的查询,
// 避免把表结构探测等内部查询混入用户慢 SQL 历史。
trackQueryHistory := strings.TrimSpace(queryID) != ""
type dbQueryMultiAuditOptions struct {
auditAll bool
auditWrites bool
source string
}
func containsSQLAuditWrite(dbType string, query string) bool {
statements := splitSQLStatements(query)
if len(statements) == 0 {
return !isReadOnlySQLQuery(dbType, query)
}
for _, statement := range statements {
statement = strings.TrimSpace(statement)
if statement != "" && !isReadOnlySQLQuery(dbType, statement) {
return true
}
}
return false
}
func (a *App) DBQuery(config connection.ConnectionConfig, dbName string, query string) connection.QueryResult {
return a.dbQueryWithCancel(config, dbName, query, "", dbQueryAuditOptions{
auditAll: a.webRuntime,
auditWrites: true,
source: "application_api",
})
}
func (a *App) DBQueryWithCancel(config connection.ConnectionConfig, dbName string, query string, queryID string) connection.QueryResult {
explicitQuery := strings.TrimSpace(queryID) != ""
auditSource := "query_editor"
if !explicitQuery {
auditSource = "application_api"
}
return a.dbQueryWithCancel(config, dbName, query, queryID, dbQueryAuditOptions{
trackHistory: explicitQuery,
auditAll: explicitQuery || a.webRuntime,
auditWrites: true,
source: auditSource,
})
}
func (a *App) dbQueryWithCancel(
config connection.ConnectionConfig,
dbName string,
query string,
queryID string,
auditOptions dbQueryAuditOptions,
) (result connection.QueryResult) {
trackQueryHistory := auditOptions.trackHistory
auditStartedAt := time.Now()
var queryExecutionDuration time.Duration
runConfig := normalizeRunConfig(config, dbName)
if queryID == "" {
queryID = generateQueryID()
}
query = sanitizeSQLForPgLike(resolveDDLDBType(config), query)
trackSQLAudit := auditOptions.auditAll || (auditOptions.auditWrites && containsSQLAuditWrite(resolveDDLDBType(runConfig), query))
if trackSQLAudit {
defer func() {
a.recordSQLAuditQuery(sqlAuditQueryInput{
Config: runConfig,
Database: dbName,
DBType: resolveDDLDBType(runConfig),
QueryID: queryID,
SQL: query,
Source: normalizeSQLAuditSource(auditOptions.source),
CommitMode: "auto",
Duration: time.Since(auditStartedAt),
Result: result,
})
}()
}
if trackQueryHistory {
defer func() {
if !result.Success {
@@ -898,12 +985,6 @@ func (a *App) DBQueryWithCancel(config connection.ConnectionConfig, dbName strin
}()
}
// Generate query ID if not provided
if queryID == "" {
queryID = generateQueryID()
}
query = sanitizeSQLForPgLike(resolveDDLDBType(config), query)
if err := ensureConnectionAllowsQuery(config, query); err != nil {
return connection.QueryResult{Success: false, Message: err.Error(), QueryID: queryID}
}
@@ -1010,9 +1091,50 @@ func (a *App) DBQueryWithCancel(config connection.ConnectionConfig, dbName strin
// DBQueryMulti 执行可能包含多条 SQL 语句的查询,返回多个结果集。
// 如果底层驱动支持 MultiResultQuerier一次性执行所有语句
// 否则按分号拆分后逐条执行,模拟多结果集。
func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, query string, queryID string) (result connection.QueryResult) {
func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, query string, queryID string) connection.QueryResult {
explicitQuery := strings.TrimSpace(queryID) != ""
auditSource := "query_editor"
if !explicitQuery {
auditSource = "application_api"
}
return a.dbQueryMulti(config, dbName, query, queryID, dbQueryMultiAuditOptions{
auditAll: explicitQuery || a.webRuntime,
auditWrites: true,
source: auditSource,
})
}
func (a *App) dbQueryMulti(
config connection.ConnectionConfig,
dbName string,
query string,
queryID string,
auditOptions dbQueryMultiAuditOptions,
) (result connection.QueryResult) {
runConfig := normalizeRunConfig(config, dbName)
resolvedDBType := resolveDDLDBType(runConfig)
trackSQLAudit := auditOptions.auditAll || (auditOptions.auditWrites && containsSQLAuditWrite(resolvedDBType, query))
auditSource := normalizeSQLAuditSource(auditOptions.source)
auditStartedAt := time.Now()
var statementAuditEvents []sqlaudit.Event
if trackSQLAudit {
defer func() {
a.recordSQLAuditQuery(sqlAuditQueryInput{
Config: runConfig,
Database: dbName,
DBType: resolvedDBType,
QueryID: queryID,
SQL: query,
Source: auditSource,
CommitMode: "auto",
Duration: time.Since(auditStartedAt),
Result: result,
})
}()
defer func() {
a.appendSQLAuditEvents(statementAuditEvents)
}()
}
// 慢 SQL 埋点:成功执行后记录(低于阈值 500ms 自动跳过)。
// 用 named return + defer 覆盖所有 return path避免遗漏。
var queryExecutionDuration time.Duration
@@ -1083,6 +1205,46 @@ func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, qu
// sql.Rows 不暴露 RowsAffected导致影响行数丢失。
// 因此仅在全部语句皆为读操作时才使用原生路径。
statements := splitSQLStatements(query)
statementCount := 0
for _, statement := range statements {
if strings.TrimSpace(statement) != "" {
statementCount++
}
}
auditSequentialStatements := trackSQLAudit && statementCount > 1
appendStatementAudit := func(
statement string,
statementIndex int,
startedAt time.Time,
rowsAffected int64,
rowsReturned int64,
statementErr error,
) {
if !auditSequentialStatements {
return
}
completedAt := time.Now()
event := buildSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: runConfig,
Database: dbName,
DBType: resolvedDBType,
QueryID: queryID,
EventType: "query_statement",
Status: sqlAuditStatusFromError(statementErr),
Source: auditSource,
CommitMode: "auto",
BoundaryMode: "unknown",
SQL: statement,
StatementIndex: statementIndex,
StatementCount: statementCount,
Duration: completedAt.Sub(startedAt),
RowsAffected: rowsAffected,
RowsReturned: rowsReturned,
Err: statementErr,
})
event.Timestamp = completedAt.UnixMilli()
statementAuditEvents = append(statementAuditEvents, event)
}
allReadOnly := true
for _, stmt := range statements {
if strings.TrimSpace(stmt) != "" && !isReadOnlySQLQuery(runConfig.Type, stmt) {
@@ -1273,6 +1435,7 @@ func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, qu
if stmt == "" {
continue
}
statementStartedAt := time.Now()
isReadStmt := isReadOnlySQLQuery(runConfig.Type, stmt)
tryQueryStmtFirst := shouldTryQueryResultFirst(runConfig.Type, stmt)
@@ -1344,6 +1507,7 @@ func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, qu
}
if err == nil {
if usedMultiResult {
var rowsAffected, rowsReturned int64
if len(statementResults) == 0 && len(messages) > 0 {
statementResults = []connection.ResultSetData{{
Rows: []map[string]interface{}{},
@@ -1359,8 +1523,12 @@ func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, qu
statementResult.Columns = []string{}
}
statementResult.StatementIndex = idx + 1
affected, returned := summarizeManagedSQLResultSet(statementResult)
rowsAffected += affected
rowsReturned += returned
resultSets = append(resultSets, statementResult)
}
appendStatementAudit(stmt, idx+1, statementStartedAt, rowsAffected, rowsReturned, nil)
continue
}
if data == nil {
@@ -1375,11 +1543,13 @@ func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, qu
Messages: messages,
StatementIndex: idx + 1,
})
appendStatementAudit(stmt, idx+1, statementStartedAt, 0, int64(len(data)), nil)
continue
}
if isReadStmt {
logger.Error(err, "DBQueryMulti 逐条查询失败(第 %d/%d 条):%s SQL片段=%q", idx+1, len(statements), formatConnSummary(runConfig), sqlSnippet(stmt))
errMsg := buildStatementExecutionFailedMessage(idx+1, err, len(resultSets))
appendStatementAudit(stmt, idx+1, statementStartedAt, 0, 0, err)
return connection.QueryResult{Success: false, Message: errMsg, QueryID: queryID}
}
}
@@ -1399,6 +1569,7 @@ func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, qu
if err != nil {
logger.Error(err, "DBQueryMulti 逐条执行失败(第 %d/%d 条):%s SQL片段=%q", idx+1, len(statements), formatConnSummary(runConfig), sqlSnippet(stmt))
errMsg := buildStatementExecutionFailedMessage(idx+1, err, len(resultSets))
appendStatementAudit(stmt, idx+1, statementStartedAt, 0, 0, err)
return connection.QueryResult{Success: false, Message: errMsg, QueryID: queryID}
}
resultSets = append(resultSets, connection.ResultSetData{
@@ -1406,6 +1577,7 @@ func (a *App) DBQueryMulti(config connection.ConnectionConfig, dbName string, qu
Columns: []string{"affectedRows"},
StatementIndex: idx + 1,
})
appendStatementAudit(stmt, idx+1, statementStartedAt, affected, 0, nil)
}
if resultSets == nil {
@@ -1518,6 +1690,8 @@ func shouldTryQueryResultFirst(dbType string, query string) bool {
}
keyword := leadingSQLKeyword(query)
switch keyword {
case "explain", "pragma":
return true
case "exec", "execute", "call":
return true
case "set", "print":
@@ -2653,7 +2827,9 @@ func (a *App) DBGetTriggers(config connection.ConnectionConfig, dbName string, t
return connection.QueryResult{Success: true, Data: ensureNonNilSlice(triggers)}
}
func (a *App) DropView(config connection.ConnectionConfig, dbName string, viewName string) connection.QueryResult {
func (a *App) DropView(config connection.ConnectionConfig, dbName string, viewName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("DROP VIEW %s", strings.TrimSpace(viewName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
viewName = strings.TrimSpace(viewName)
if viewName == "" {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.view_name_required", nil)}
@@ -2687,7 +2863,9 @@ func (a *App) DropView(config connection.ConnectionConfig, dbName string, viewNa
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.view_dropped", nil)}
}
func (a *App) DropFunction(config connection.ConnectionConfig, dbName string, routineName string, routineType string) connection.QueryResult {
func (a *App) DropFunction(config connection.ConnectionConfig, dbName string, routineName string, routineType string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("DROP %s %s", strings.ToUpper(strings.TrimSpace(routineType)), strings.TrimSpace(routineName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
routineName = strings.TrimSpace(routineName)
routineType = strings.TrimSpace(strings.ToUpper(routineType))
if routineName == "" {
@@ -2732,7 +2910,9 @@ func (a *App) DropFunction(config connection.ConnectionConfig, dbName string, ro
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.function_dropped", nil)}
}
func (a *App) RenameView(config connection.ConnectionConfig, dbName string, oldName string, newName string) connection.QueryResult {
func (a *App) RenameView(config connection.ConnectionConfig, dbName string, oldName string, newName string) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("ALTER VIEW %s RENAME TO %s", strings.TrimSpace(oldName), strings.TrimSpace(newName))
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
oldName = strings.TrimSpace(oldName)
newName = strings.TrimSpace(newName)
if oldName == "" || newName == "" {

View File

@@ -0,0 +1,140 @@
package app
import (
"strings"
"time"
"GoNavi-Wails/internal/connection"
)
type sqlAuditUserActionOptions struct {
StatementCount *int
SafeError *string
}
// beginSQLAuditUserAction returns a deferred recorder for exported user-facing
// write helpers that already own execution, validation, and dialect handling.
// The supplied text is an operation descriptor or generated statement and must
// never contain row values or message payloads.
func (a *App) beginSQLAuditUserAction(
config connection.ConnectionConfig,
dbName string,
source string,
sqlText *string,
result *connection.QueryResult,
statementCount ...*int,
) func() {
options := sqlAuditUserActionOptions{}
if len(statementCount) > 0 {
options.StatementCount = statementCount[0]
}
return a.beginSQLAuditUserActionWithOptions(config, dbName, source, sqlText, result, options)
}
func (a *App) beginSQLAuditUserActionWithOptions(
config connection.ConnectionConfig,
dbName string,
source string,
sqlText *string,
result *connection.QueryResult,
options sqlAuditUserActionOptions,
) func() {
startedAt := time.Now()
queryID := generateQueryID()
runConfig := normalizeRunConfig(config, dbName)
return func() {
if result == nil {
return
}
statement := ""
if sqlText != nil {
statement = strings.TrimSpace(*sqlText)
}
resolvedStatementCount := 0
if options.StatementCount != nil {
resolvedStatementCount = *options.StatementCount
}
auditResult := *result
if !auditResult.Success && options.SafeError != nil {
auditResult.Message = strings.TrimSpace(*options.SafeError)
}
a.recordSQLAuditQuery(sqlAuditQueryInput{
Config: runConfig,
Database: dbName,
DBType: resolveDDLDBType(runConfig),
QueryID: queryID,
SQL: statement,
Source: source,
CommitMode: "auto",
Duration: time.Since(startedAt),
StatementCount: resolvedStatementCount,
Result: auditResult,
})
}
}
// DBQueryAudited executes one explicit application-level user action and writes
// exactly one SQL audit event without adding it to slow-query history.
func (a *App) DBQueryAudited(
config connection.ConnectionConfig,
dbName string,
query string,
source string,
) connection.QueryResult {
queryID := generateQueryID()
return a.dbQueryWithCancel(config, dbName, query, queryID, dbQueryAuditOptions{
auditAll: true,
auditWrites: true,
source: normalizeSQLAuditUserActionSource(source),
})
}
// DBQueryAI is the dedicated entry point used by GoNavi's built-in AI tool
// runtime. The audit source describes this called entry point; it is not an
// unforgeable actor identity or security provenance claim.
func (a *App) DBQueryAI(
config connection.ConnectionConfig,
dbName string,
query string,
) connection.QueryResult {
return a.dbQueryWithCancel(config, dbName, query, "", dbQueryAuditOptions{
auditAll: true,
auditWrites: true,
source: "ai_action",
})
}
// MCPQueryExecutor is a narrow adapter for the MCP server. Keeping it separate
// from App's Wails-bound method set makes the mcp source a backend-owned fact
// rather than a source string accepted from a browser or desktop caller.
type MCPQueryExecutor struct {
app *App
}
func NewMCPQueryExecutor(app *App) *MCPQueryExecutor {
return &MCPQueryExecutor{app: app}
}
func (executor *MCPQueryExecutor) DBQueryMulti(
config connection.ConnectionConfig,
dbName string,
query string,
) connection.QueryResult {
if executor == nil || executor.app == nil {
return connection.QueryResult{Success: false, Message: "MCP query executor is unavailable"}
}
return executor.app.dbQueryMulti(config, dbName, query, "", dbQueryMultiAuditOptions{
auditAll: true,
auditWrites: true,
source: "mcp",
})
}
func normalizeSQLAuditUserActionSource(source string) string {
switch normalized := strings.ToLower(strings.TrimSpace(source)); normalized {
case "data_editor", "table_designer", "object_editor", "message_publish":
return normalized
default:
return "application_api"
}
}

View File

@@ -0,0 +1,384 @@
package app
import (
"bytes"
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"time"
"GoNavi-Wails/internal/connection"
"GoNavi-Wails/internal/db"
"GoNavi-Wails/internal/sqlaudit"
datasync "GoNavi-Wails/internal/sync"
)
func TestDBQueryAuditedWritesOneUserActionEventWithoutSlowQueryHistory(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
query := "UPDATE users SET display_name = 'private-name' WHERE id = 7"
database := &fakeBatchWriteDB{
execAffected: map[string]int64{query: 2},
execDelay: map[string]time.Duration{query: time.Duration(queryHistorySlowThresholdMs)*time.Millisecond + 25*time.Millisecond},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, Database: "app"}
result := app.DBQueryAudited(config, "app", query, "table_designer")
if !result.Success {
t.Fatalf("DBQueryAudited returned failure: %s", result.Message)
}
if !strings.HasPrefix(result.QueryID, "query-") {
t.Fatalf("query ID = %q, want generated query ID", result.QueryID)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: result.QueryID})
if len(events) != 1 {
t.Fatalf("audit event count = %d, want 1: %#v", len(events), events)
}
event := events[0]
if event.QueryID != result.QueryID || event.Source != "table_designer" || event.Status != "success" || event.RowsAffected != 2 {
t.Fatalf("unexpected application audit event: %#v", event)
}
if strings.Contains(event.SQLText, "private-name") || strings.Contains(event.SQLText, " 7") || !event.SQLRedacted {
t.Fatalf("application audit SQL was not redacted: %#v", event)
}
history := app.GetSlowQueries(config, "app", "recent", 10)
if !history.Success {
t.Fatalf("GetSlowQueries returned failure: %s", history.Message)
}
records, ok := history.Data.([]connection.QueryExecutionRecord)
if !ok {
t.Fatalf("GetSlowQueries data type = %T, want []connection.QueryExecutionRecord", history.Data)
}
if len(records) != 0 {
t.Fatalf("DBQueryAudited must not write slow query history: %#v", records)
}
}
func TestDBQueryAuditedRecordsProtectionDenial(t *testing.T) {
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, Database: "app", ReadOnly: true}
query := "DELETE FROM users WHERE email = 'private@example.test'"
result := app.DBQueryAudited(config, "app", query, "object_editor")
if result.Success || strings.TrimSpace(result.Message) == "" {
t.Fatalf("protected action result = %#v, want rejected error", result)
}
if !strings.HasPrefix(result.QueryID, "query-") {
t.Fatalf("query ID = %q, want generated query ID", result.QueryID)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: result.QueryID})
if len(events) != 1 {
t.Fatalf("audit event count = %d, want 1: %#v", len(events), events)
}
event := events[0]
if event.Status != "error" || event.Source != "object_editor" || strings.TrimSpace(event.Error) == "" {
t.Fatalf("protection denial was not fully audited: %#v", event)
}
if strings.Contains(event.SQLText, "private@example.test") || !event.SQLRedacted {
t.Fatalf("denied audit SQL was not redacted: %#v", event)
}
}
func TestDirectDBQueryCannotBypassWriteAuditWithEmptyQueryID(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
writeSQL := "DELETE FROM users WHERE id = 7"
readSQL := "SELECT id FROM users"
database := &fakeBatchWriteDB{
execAffected: map[string]int64{writeSQL: 1},
queryMap: map[string][]map[string]interface{}{readSQL: {{"id": 7}}},
fieldMap: map[string][]string{readSQL: {"id"}},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "app"}
if result := app.DBQuery(config, "app", writeSQL); !result.Success {
t.Fatalf("direct DBQuery write returned failure: %s", result.Message)
}
if result := app.DBQuery(config, "app", readSQL); !result.Success {
t.Fatalf("direct DBQuery metadata read returned failure: %s", result.Message)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 1 || events[0].Source != "application_api" || events[0].RowsAffected != 1 {
t.Fatalf("direct write was not uniquely audited: %#v", events)
}
}
func TestDirectDBQueryCannotBypassWriteAuditWhenBatchStartsWithRead(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
query := "SELECT 1; UPDATE users SET enabled = false WHERE id = 7"
database := &fakeBatchWriteDB{
queryMap: map[string][]map[string]interface{}{query: {{"value": 1}}},
fieldMap: map[string][]string{query: {"value"}},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "app"}
result := app.DBQuery(config, "app", query)
if !result.Success {
t.Fatalf("direct mixed batch returned failure: %s", result.Message)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 1 || events[0].Source != "application_api" {
t.Fatalf("read-first write batch bypassed audit: %#v", events)
}
}
func TestDirectDBQueryCannotBypassAuditWithNestedWriteSyntax(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
explainWrite := "EXPLAIN ANALYZE UPDATE users SET enabled = false WHERE id = 7"
pragmaWrite := "PRAGMA user_version = 7"
mongoWrite := `{"aggregate":"users","pipeline":[{"$merge":{"into":"users_archive"}}],"cursor":{}}`
database := &fakeBatchWriteDB{
queryMap: map[string][]map[string]interface{}{explainWrite: {{"plan": "ok"}}},
fieldMap: map[string][]string{explainWrite: {"plan"}},
execAffected: map[string]int64{pragmaWrite: 0, mongoWrite: 1},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
for _, testCase := range []struct {
dbType string
query string
}{
{dbType: "postgres", query: explainWrite},
{dbType: "sqlite", query: pragmaWrite},
{dbType: "mongodb", query: mongoWrite},
} {
config := connection.ConnectionConfig{Type: testCase.dbType, Host: "127.0.0.1", Database: "app"}
if result := app.DBQuery(config, "app", testCase.query); !result.Success {
t.Fatalf("direct %s nested write returned failure: %s", testCase.dbType, result.Message)
}
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 3 {
t.Fatalf("nested write syntax bypassed audit: %#v", events)
}
for _, event := range events {
if event.Source != "application_api" {
t.Fatalf("unexpected nested write audit source: %#v", event)
}
}
}
func TestWebRuntimeAuditsGenericReadQueriesWithoutTrustingSQLClassification(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
query := "SELECT mutating_function(7)"
database := &fakeBatchWriteDB{
queryMap: map[string][]map[string]interface{}{query: {{"result": 1}}},
fieldMap: map[string][]string{query: {"result"}},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
app.webRuntime = true
config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, Database: "app"}
result := app.DBQuery(config, "app", query)
if !result.Success {
t.Fatalf("web generic read returned failure: %s", result.Message)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 1 || events[0].Source != "application_api" {
t.Fatalf("web generic query was not audit-all: %#v", events)
}
}
func TestObjectDDLHelperRecordsProtectionDenial(t *testing.T) {
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "app", ReadOnly: true}
result := app.DropTable(config, "app", "private_orders")
if result.Success {
t.Fatalf("DropTable on protected connection returned success: %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Source: "object_editor"})
if len(events) != 1 || events[0].Status != "error" || events[0].Source != "object_editor" {
t.Fatalf("protected object DDL was not audited: %#v", events)
}
if !strings.Contains(events[0].SQLText, "DROP TABLE") {
t.Fatalf("object DDL audit lost operation structure: %#v", events[0])
}
}
func TestApplyChangesRecordsMetadataWithoutRowValues(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
fakeDB := &fakeCreateDatabaseDB{}
newDatabaseFunc = func(string) (db.Database, error) { return fakeDB, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "app"}
result := app.ApplyChanges(config, "app", "users", connection.ChangeSet{
Updates: []connection.UpdateRow{{
Keys: map[string]interface{}{"id": 7},
Values: map[string]interface{}{"email": "private@example.test"},
}},
})
if !result.Success {
t.Fatalf("ApplyChanges returned failure: %s", result.Message)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Source: "data_editor"})
if len(events) != 1 || events[0].Status != "success" || events[0].Source != "data_editor" {
t.Fatalf("data editor action was not audited: %#v", events)
}
if strings.Contains(events[0].SQLText, "private@example.test") || strings.Contains(events[0].SQLText, " 7") {
t.Fatalf("data editor audit leaked row values: %#v", events[0])
}
}
func TestBackendWriteWorkflowsRecordFixedAuditSources(t *testing.T) {
protectedConfig := connection.ConnectionConfig{
Type: "mysql",
Host: "127.0.0.1",
Port: 3306,
Database: "app",
ReadOnly: true,
}
t.Run("sql_file", func(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
newDatabaseFunc = func(string) (db.Database, error) { return &fakeSQLFileBatchDB{}, nil }
app := newSQLAuditTestApp(t)
filePath := filepath.Join(t.TempDir(), "private-job.sql")
if err := os.WriteFile(filePath, []byte("INSERT INTO users(email) VALUES ('private@example.test');"), 0o600); err != nil {
t.Fatalf("write SQL file fixture: %v", err)
}
config := protectedConfig
config.ReadOnly = false
result := app.ExecuteSQLFile(config, "app", filePath, "sql-file-audit-test")
if !result.Success {
t.Fatalf("SQL file execution returned failure: %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Source: "sql_file"})
if len(events) != 1 || events[0].Source != "sql_file" || events[0].Status != "success" {
t.Fatalf("SQL file execution was not audited with its fixed source: %#v", events)
}
if strings.Contains(events[0].SQLText, "private@example.test") || strings.Contains(events[0].SQLText, "private-job.sql") {
t.Fatalf("SQL file audit leaked file contents or path: %#v", events[0])
}
if !strings.Contains(events[0].SQLText, "SHA256_") || !strings.Contains(events[0].SQLText, "EXECUTED_1") || events[0].StatementCount != 1 {
t.Fatalf("SQL file audit lacks a safe content identity or execution counts: %#v", events[0])
}
})
t.Run("sql_file_failure_redaction", func(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
newDatabaseFunc = func(string) (db.Database, error) {
return &fakeSQLFileBatchDB{failExecSQL: "CALL broken_proc"}, nil
}
app := newSQLAuditTestApp(t)
filePath := filepath.Join(t.TempDir(), "private-failure.sql")
if err := os.WriteFile(filePath, []byte("CALL broken_proc('private-secret', 777);"), 0o600); err != nil {
t.Fatalf("write failing SQL file fixture: %v", err)
}
config := protectedConfig
config.ReadOnly = false
result := app.ExecuteSQLFile(config, "app", filePath, "sql-file-failure-audit-test")
if result.Success {
t.Fatalf("failing SQL file execution returned success: %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Source: "sql_file"})
if len(events) != 1 || events[0].Status != "error" {
t.Fatalf("failing SQL file execution was not audited: %#v", events)
}
serialized, _ := json.Marshal(events[0])
for _, secret := range []string{"private-secret", "777", "private-failure.sql", "broken_proc"} {
if bytes.Contains(serialized, []byte(secret)) {
t.Fatalf("failing SQL file audit leaked %q: %s", secret, serialized)
}
}
})
t.Run("data_import", func(t *testing.T) {
app := newSQLAuditTestApp(t)
result := app.ImportDataWithProgress(protectedConfig, "app", "users", "private.csv")
if result.Success {
t.Fatalf("protected import returned success: %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Source: "data_import"})
if len(events) != 1 || events[0].Source != "data_import" || events[0].Status != "error" {
t.Fatalf("data import attempt was not audited with its fixed source: %#v", events)
}
if strings.Contains(events[0].SQLText, "private.csv") {
t.Fatalf("data import audit leaked the local file path: %#v", events[0])
}
if !strings.Contains(strings.ToLower(events[0].SQLText), "users") {
t.Fatalf("data import audit lost the safe target table: %#v", events[0])
}
})
t.Run("sync", func(t *testing.T) {
app := newSQLAuditTestApp(t)
result := app.DataSync(datasync.SyncConfig{
TargetConfig: protectedConfig,
TargetDatabase: "app",
Tables: []string{"users"},
Content: "data",
})
if result.Success {
t.Fatalf("protected data sync returned success: %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Source: "sync"})
if len(events) != 1 || events[0].Source != "sync" || events[0].Status != "error" {
t.Fatalf("data sync attempt was not audited with its fixed source: %#v", events)
}
if events[0].StatementCount != 1 || !strings.Contains(events[0].SQLText, "TABLES_1") {
t.Fatalf("data sync audit lacks its safe task summary: %#v", events[0])
}
})
}
func TestAIEntryPointAndMCPExecutorRecordAuditSources(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
database := &sqlAuditTestDatabase{
rows: []map[string]interface{}{{"id": 7}},
columns: []string{"id"},
affected: 1,
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, Database: "app"}
aiResult := app.DBQueryAI(config, "app", "SELECT id FROM users WHERE email = 'private@example.test'")
if !aiResult.Success {
t.Fatalf("DBQueryAI returned failure: %s", aiResult.Message)
}
mcpResult := NewMCPQueryExecutor(app).DBQueryMulti(config, "app", "SELECT id FROM users")
if !mcpResult.Success {
t.Fatalf("MCPQueryExecutor returned failure: %s", mcpResult.Message)
}
spoofedResult := app.DBQueryAudited(config, "app", "UPDATE users SET email = 'private@example.test' WHERE id = 7", "mcp")
if !spoofedResult.Success {
t.Fatalf("DBQueryAudited returned failure: %s", spoofedResult.Message)
}
aiEvents := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: aiResult.QueryID})
if len(aiEvents) != 1 || aiEvents[0].Source != "ai_action" {
t.Fatalf("AI query did not retain its entry-point source: %#v", aiEvents)
}
if strings.Contains(aiEvents[0].SQLText, "private@example.test") {
t.Fatalf("AI query audit leaked a literal: %#v", aiEvents[0])
}
mcpEvents := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: mcpResult.QueryID})
if len(mcpEvents) != 1 || mcpEvents[0].Source != "mcp" {
t.Fatalf("MCP query did not retain its backend-owned source: %#v", mcpEvents)
}
spoofedEvents := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: spoofedResult.QueryID})
if len(spoofedEvents) != 1 || spoofedEvents[0].Source != "application_api" {
t.Fatalf("public audited query was able to spoof a privileged source: %#v", spoofedEvents)
}
}

View File

@@ -525,7 +525,9 @@ func TestMethodsDBQueryMultiMessagesUseLocalizedText(t *testing.T) {
}
source := string(sourceBytes)
functionSource := methodsDBFunctionSource(t, source, "func (a *App) DBQueryMulti")
// DBQueryMulti only selects the audit policy; the shared execution path owns
// all user-facing multi-statement messages.
functionSource := methodsDBFunctionSource(t, source, "func (a *App) dbQueryMulti")
rawMessages := []string{
`fmt.Sprintf("第 %d 条语句执行失败: %v", idx+1, err)`,
`fmt.Sprintf("(前 %d 条已执行成功)", len(resultSets))`,
@@ -696,7 +698,7 @@ func TestMethodsDBManagedTransactionExecutionMessagesUseLocalizedText(t *testing
"db.backend.error.transaction_rollback_failed",
},
},
"func executeManagedSQLTransactionStatements": {
"func executeManagedSQLTransactionStatementsWithObserver": {
rawMessages: []string{
`fmt.Errorf("当前事务会话不支持查询语句")`,
`fmt.Errorf("第 %d 条语句执行失败: %w", idx+1, err)`,

View File

@@ -7,6 +7,7 @@ import (
"reflect"
"strings"
"testing"
"time"
"GoNavi-Wails/internal/connection"
"GoNavi-Wails/internal/db"
@@ -27,6 +28,9 @@ type fakeBatchWriteDB struct {
queryErr map[string]error
execErr map[string]error
execAffected map[string]int64
execDelay map[string]time.Duration
execStarted chan<- string
execRelease <-chan struct{}
session *fakeBatchWriteSession
}
@@ -173,6 +177,29 @@ func (f *fakeBatchWriteDB) ExecContext(ctx context.Context, query string) (int64
f.lastCtx = ctx
f.execCalls++
f.execQueries = append(f.execQueries, query)
if f.execStarted != nil {
select {
case f.execStarted <- query:
case <-ctx.Done():
return 0, ctx.Err()
}
}
if f.execRelease != nil {
select {
case <-f.execRelease:
case <-ctx.Done():
return 0, ctx.Err()
}
}
if delay := f.execDelay[query]; delay > 0 {
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-timer.C:
case <-ctx.Done():
return 0, ctx.Err()
}
}
if err := f.execErr[query]; err != nil {
return 0, err
}
@@ -356,6 +383,126 @@ func cloneResultSets(input []connection.ResultSetData) []connection.ResultSetDat
return cloned
}
func TestDBQueryMultiInTransactionSerializesCommitWithInFlightStatement(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
initialStatement := "UPDATE users SET active = 1 WHERE id = 1"
followUpStatement := "UPDATE users SET active = 0 WHERE id = 2"
fakeDB := &fakeTransactionalDB{fakeBatchWriteDB: fakeBatchWriteDB{
execAffected: map[string]int64{initialStatement: 1, followUpStatement: 1},
}}
newDatabaseFunc = func(string) (db.Database, error) { return fakeDB, nil }
app := NewAppWithSecretStore(secretstore.NewUnavailableStore("test"))
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "main"}
started := app.DBQueryMultiTransactional(config, "main", initialStatement, "tx-serialize-start")
if !started.Success || started.TransactionID == "" {
t.Fatalf("start managed transaction: %#v", started)
}
execStarted := make(chan string, 1)
execRelease := make(chan struct{})
fakeDB.execStarted = execStarted
fakeDB.execRelease = execRelease
queryDone := make(chan connection.QueryResult, 1)
go func() {
queryDone <- app.DBQueryMultiInTransaction(started.TransactionID, followUpStatement, "tx-serialize-follow-up")
}()
select {
case statement := <-execStarted:
if statement != followUpStatement {
close(execRelease)
t.Fatalf("blocked statement = %q, want %q", statement, followUpStatement)
}
case <-time.After(2 * time.Second):
close(execRelease)
t.Fatal("follow-up statement did not start")
}
commitDone := make(chan connection.QueryResult, 1)
go func() {
commitDone <- app.DBCommitTransaction(started.TransactionID)
}()
select {
case result := <-commitDone:
close(execRelease)
t.Fatalf("commit completed while statement was still in flight: %#v", result)
case <-time.After(75 * time.Millisecond):
}
close(execRelease)
if result := <-queryDone; !result.Success {
t.Fatalf("follow-up statement failed: %#v", result)
}
if result := <-commitDone; !result.Success {
t.Fatalf("commit after follow-up failed: %#v", result)
}
if fakeDB.txSession.commitCalls != 1 || !fakeDB.txSession.closed {
t.Fatalf("transaction was not committed exactly once after execution: commitCalls=%d closed=%v", fakeDB.txSession.commitCalls, fakeDB.txSession.closed)
}
}
func TestRollbackPendingSQLTransactionsWaitsForInFlightStatement(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
initialStatement := "UPDATE users SET active = 1 WHERE id = 1"
followUpStatement := "UPDATE users SET active = 0 WHERE id = 2"
fakeDB := &fakeTransactionalDB{fakeBatchWriteDB: fakeBatchWriteDB{
execAffected: map[string]int64{initialStatement: 1, followUpStatement: 1},
}}
newDatabaseFunc = func(string) (db.Database, error) { return fakeDB, nil }
app := NewAppWithSecretStore(secretstore.NewUnavailableStore("test"))
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "main"}
started := app.DBQueryMultiTransactional(config, "main", initialStatement, "tx-shutdown-start")
if !started.Success || started.TransactionID == "" {
t.Fatalf("start managed transaction: %#v", started)
}
execStarted := make(chan string, 1)
execRelease := make(chan struct{})
fakeDB.execStarted = execStarted
fakeDB.execRelease = execRelease
queryDone := make(chan connection.QueryResult, 1)
go func() {
queryDone <- app.DBQueryMultiInTransaction(started.TransactionID, followUpStatement, "tx-shutdown-follow-up")
}()
select {
case <-execStarted:
case <-time.After(2 * time.Second):
close(execRelease)
t.Fatal("follow-up statement did not start")
}
shutdownDone := make(chan struct{})
go func() {
app.rollbackPendingSQLTransactionsOnShutdown()
close(shutdownDone)
}()
select {
case <-shutdownDone:
close(execRelease)
t.Fatal("shutdown rollback completed while statement was still in flight")
case <-time.After(75 * time.Millisecond):
}
close(execRelease)
if result := <-queryDone; !result.Success {
t.Fatalf("follow-up statement failed: %#v", result)
}
select {
case <-shutdownDone:
case <-time.After(2 * time.Second):
t.Fatal("shutdown rollback did not finish after statement completed")
}
if fakeDB.txSession.rollbackCalls != 1 || !fakeDB.txSession.closed {
t.Fatalf("transaction was not rolled back exactly once after execution: rollbackCalls=%d closed=%v", fakeDB.txSession.rollbackCalls, fakeDB.txSession.closed)
}
}
func TestDBQueryMultiKeepsOracleAnonymousBlockAsSingleStatement(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() {

View File

@@ -9,11 +9,48 @@ import (
"GoNavi-Wails/internal/connection"
"GoNavi-Wails/internal/db"
"GoNavi-Wails/internal/logger"
"GoNavi-Wails/internal/sqlaudit"
"github.com/google/uuid"
)
const sqlEditorTransactionFinishTimeout = 30 * time.Second
type managedSQLStatementObservation struct {
Statement string
StatementIndex int
StatementCount int
StartedAt time.Time
CompletedAt time.Time
Duration time.Duration
RowsAffected int64
RowsReturned int64
Err error
}
type managedSQLStatementObserver func(managedSQLStatementObservation)
func withManagedSQLStatementAuditTimestamp(
observer managedSQLStatementObserver,
events *[]sqlaudit.Event,
) managedSQLStatementObserver {
if observer == nil {
return nil
}
return func(observation managedSQLStatementObservation) {
before := 0
if events != nil {
before = len(*events)
}
observer(observation)
if events == nil || observation.CompletedAt.IsZero() {
return
}
for index := before; index < len(*events); index++ {
(*events)[index].Timestamp = observation.CompletedAt.UnixMilli()
}
}
}
// DBQueryMultiTransactional executes SQL editor DML in a managed transaction.
// The transaction stays open until DBCommitTransaction or DBRollbackTransaction
// is called by the SQL editor UI.
@@ -45,12 +82,35 @@ func (a *App) DBQueryMultiTransactional(config connection.ConnectionConfig, dbNa
}
query = sanitizeSQLForPgLike(transactionDBType, query)
if err := ensureConnectionAllowsQuery(config, query); err != nil {
return connection.QueryResult{Success: false, Message: err.Error(), QueryID: queryID}
}
if !shouldUseManagedSQLTransaction(transactionDBType, query) {
return a.DBQueryMulti(config, dbName, query, queryID)
}
transactionID := "sql-editor-" + uuid.NewString()
transactionAuditOpened := false
transactionBoundaryMode := "unknown"
defer func() {
if transactionAuditOpened {
return
}
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: transactionConfig,
Database: dbName,
DBType: transactionDBType,
QueryID: queryID,
TransactionID: transactionID,
EventType: "transaction_begin",
Status: sqlAuditStatusFromResult(result),
Source: "query_editor",
CommitMode: "pending",
BoundaryMode: transactionBoundaryMode,
SQL: query,
StatementCount: countSQLAuditStatements(query),
Err: sqlAuditErrorFromResult(result),
})
}()
if err := ensureConnectionAllowsQuery(config, query); err != nil {
return connection.QueryResult{Success: false, Message: err.Error(), QueryID: queryID}
}
var queryExecutionDuration time.Duration
defer func() {
if !result.Success {
@@ -97,6 +157,7 @@ func (a *App) DBQueryMultiTransactional(config connection.ConnectionConfig, dbNa
startTextTransaction bool
)
if provider, ok := dbInst.(db.TransactionExecerProvider); ok {
transactionBoundaryMode = "driver_api"
// database/sql rolls back a BeginTx transaction when its context is cancelled.
// SQL editor transactions must outlive the execution RPC and be ended only by
// explicit commit, rollback, or shutdown cleanup.
@@ -111,6 +172,7 @@ func (a *App) DBQueryMultiTransactional(config connection.ConnectionConfig, dbNa
sessionExecer = transactionExecer
transactor = transactionExecer
} else if implicitTextTransaction {
transactionBoundaryMode = "implicit"
provider, ok := dbInst.(db.SessionExecerProvider)
if !ok {
return connection.QueryResult{
@@ -125,6 +187,7 @@ func (a *App) DBQueryMultiTransactional(config connection.ConnectionConfig, dbNa
return connection.QueryResult{Success: false, Message: err.Error(), QueryID: queryID}
}
} else {
transactionBoundaryMode = "text_sql"
if !hasTextTransaction {
return connection.QueryResult{
Success: false,
@@ -166,12 +229,49 @@ func (a *App) DBQueryMultiTransactional(config connection.ConnectionConfig, dbNa
return connection.QueryResult{Success: false, Message: err.Error(), QueryID: queryID}
}
}
transactionAuditOpened = true
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: transactionConfig,
Database: dbName,
DBType: transactionDBType,
QueryID: queryID,
TransactionID: transactionID,
EventType: "transaction_begin",
Status: "success",
Source: "query_editor",
CommitMode: "pending",
BoundaryMode: transactionBoundaryMode,
})
statements := splitSQLStatements(query)
queryStartedAt := time.Now()
resultSets, err := executeManagedSQLTransactionStatements(ctx, sessionExecer, transactionConfig, statements, a.appText)
statementAuditEvents := make([]sqlaudit.Event, 0, len(statements))
resultSets, err := executeManagedSQLTransactionStatementsWithObserver(
ctx,
sessionExecer,
transactionConfig,
statements,
a.appText,
withManagedSQLStatementAuditTimestamp(
a.sqlAuditTransactionStatementObserver(transactionConfig, dbName, transactionDBType, queryID, transactionID, transactionBoundaryMode, &statementAuditEvents),
&statementAuditEvents,
),
)
queryExecutionDuration += time.Since(queryStartedAt)
a.appendSQLAuditEvents(statementAuditEvents)
if err != nil {
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: transactionConfig,
Database: dbName,
DBType: transactionDBType,
QueryID: queryID,
TransactionID: transactionID,
EventType: "transaction_rollback_requested",
Status: "success",
Source: "query_editor",
CommitMode: "auto",
BoundaryMode: transactionBoundaryMode,
})
var rollbackErr error
if transactor != nil {
rollbackErr = transactor.Rollback()
@@ -182,25 +282,38 @@ func (a *App) DBQueryMultiTransactional(config connection.ConnectionConfig, dbNa
logger.Error(rollbackErr, "DBQueryMultiTransactional 执行失败后回滚失败:%s SQL片段=%q", formatConnSummary(runConfig), sqlSnippet(query))
err = appendRollbackFailureMessage(err, rollbackErr)
}
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: transactionConfig,
Database: dbName,
DBType: transactionDBType,
QueryID: queryID,
TransactionID: transactionID,
EventType: "transaction_auto_rollback",
Status: sqlAuditStatusFromError(rollbackErr),
Source: "query_editor",
CommitMode: "auto",
BoundaryMode: transactionBoundaryMode,
Err: rollbackErr,
})
logger.Error(err, "DBQueryMultiTransactional 执行失败:%s SQL片段=%q", formatConnSummary(runConfig), sqlSnippet(query))
return connection.QueryResult{Success: false, Message: err.Error(), QueryID: queryID}
}
transactionID := "sql-editor-" + uuid.NewString()
a.sqlTransactionMu.Lock()
if a.sqlTransactions == nil {
a.sqlTransactions = make(map[string]*managedSQLTransaction)
}
a.sqlTransactions[transactionID] = &managedSQLTransaction{
id: transactionID,
execer: sessionExecer,
transactor: transactor,
cancel: transactionCancel,
config: runConfig,
dbType: transactionDBType,
commitSQL: commitSQL,
rollbackSQL: rollbackSQL,
createdAt: time.Now(),
id: transactionID,
execer: sessionExecer,
transactor: transactor,
cancel: transactionCancel,
config: runConfig,
dbType: transactionDBType,
boundaryMode: transactionBoundaryMode,
commitSQL: commitSQL,
rollbackSQL: rollbackSQL,
createdAt: time.Now(),
}
a.sqlTransactionMu.Unlock()
@@ -231,6 +344,11 @@ func (a *App) DBQueryMultiInTransaction(transactionID string, query string, quer
if !ok || tx == nil || tx.execer == nil {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.transaction_not_found", nil), QueryID: queryID}
}
tx.mu.Lock()
defer tx.mu.Unlock()
if tx.finished || tx.execer == nil {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.transaction_not_found", nil), QueryID: queryID}
}
runConfig := tx.config
if strings.TrimSpace(runConfig.Type) == "" {
@@ -263,8 +381,20 @@ func (a *App) DBQueryMultiInTransaction(transactionID string, query string, quer
}()
queryStartedAt := time.Now()
resultSets, err := executeManagedSQLTransactionStatements(ctx, tx.execer, runConfig, statements, a.appText)
statementAuditEvents := make([]sqlaudit.Event, 0, len(statements))
resultSets, err := executeManagedSQLTransactionStatementsWithObserver(
ctx,
tx.execer,
runConfig,
statements,
a.appText,
withManagedSQLStatementAuditTimestamp(
a.sqlAuditTransactionStatementObserver(runConfig, runConfig.Database, tx.dbType, queryID, transactionID, tx.boundaryMode, &statementAuditEvents),
&statementAuditEvents,
),
)
queryExecutionDuration += time.Since(queryStartedAt)
a.appendSQLAuditEvents(statementAuditEvents)
if err != nil {
logger.Error(err, "DBQueryMultiInTransaction 执行失败id=%s dbType=%s SQL片段=%q", transactionID, tx.dbType, sqlSnippet(query))
return connection.QueryResult{
@@ -286,6 +416,17 @@ func (a *App) DBQueryMultiInTransaction(transactionID string, query string, quer
}
func executeManagedSQLTransactionStatements(ctx context.Context, session db.StatementExecer, runConfig connection.ConnectionConfig, statements []string, text func(string, map[string]any) string) ([]connection.ResultSetData, error) {
return executeManagedSQLTransactionStatementsWithObserver(ctx, session, runConfig, statements, text, nil)
}
func executeManagedSQLTransactionStatementsWithObserver(
ctx context.Context,
session db.StatementExecer,
runConfig connection.ConnectionConfig,
statements []string,
text func(string, map[string]any) string,
observer managedSQLStatementObserver,
) ([]connection.ResultSetData, error) {
if text == nil {
text = defaultDBBackendText
}
@@ -306,11 +447,37 @@ func executeManagedSQLTransactionStatements(ctx context.Context, session db.Stat
sessionMultiQueryTarget, _ := session.(db.StatementMultiResultQueryExecer)
sessionMultiQueryMessageTarget, _ := session.(db.StatementMultiResultQueryMessageExecer)
for idx, stmt := range statements {
statementCount := 0
for _, statement := range statements {
if strings.TrimSpace(statement) != "" {
statementCount++
}
}
statementIndex := 0
for _, stmt := range statements {
stmt = strings.TrimSpace(stmt)
if stmt == "" {
continue
}
statementIndex++
statementStartedAt := time.Now()
emitObservation := func(rowsAffected, rowsReturned int64, err error) {
if observer == nil {
return
}
completedAt := time.Now()
observer(managedSQLStatementObservation{
Statement: stmt,
StatementIndex: statementIndex,
StatementCount: statementCount,
StartedAt: statementStartedAt,
CompletedAt: completedAt,
Duration: completedAt.Sub(statementStartedAt),
RowsAffected: rowsAffected,
RowsReturned: rowsReturned,
Err: err,
})
}
isReadStmt := isReadOnlySQLQuery(runConfig.Type, stmt)
tryQueryStmtFirst := shouldTryQueryResultFirst(runConfig.Type, stmt)
@@ -346,6 +513,7 @@ func executeManagedSQLTransactionStatements(ctx context.Context, session db.Stat
}
if err == nil {
if usedMultiResult {
var rowsAffected, rowsReturned int64
if len(statementResults) == 0 && len(messages) > 0 {
statementResults = []connection.ResultSetData{{
Rows: []map[string]interface{}{},
@@ -360,9 +528,13 @@ func executeManagedSQLTransactionStatements(ctx context.Context, session db.Stat
if statementResult.Columns == nil {
statementResult.Columns = []string{}
}
statementResult.StatementIndex = idx + 1
statementResult.StatementIndex = statementIndex
affected, returned := summarizeManagedSQLResultSet(statementResult)
rowsAffected += affected
rowsReturned += returned
resultSets = append(resultSets, statementResult)
}
emitObservation(rowsAffected, rowsReturned, nil)
continue
}
if data == nil {
@@ -375,24 +547,30 @@ func executeManagedSQLTransactionStatements(ctx context.Context, session db.Stat
Rows: data,
Columns: columns,
Messages: messages,
StatementIndex: idx + 1,
StatementIndex: statementIndex,
})
emitObservation(0, int64(len(data)), nil)
continue
}
if isReadStmt {
return nil, buildStatementExecutionFailedError(idx+1, err)
statementErr := buildStatementExecutionFailedError(statementIndex, err)
emitObservation(0, 0, statementErr)
return nil, statementErr
}
}
affected, err := session.ExecContext(ctx, stmt)
if err != nil {
return nil, buildStatementExecutionFailedError(idx+1, err)
statementErr := buildStatementExecutionFailedError(statementIndex, err)
emitObservation(0, 0, statementErr)
return nil, statementErr
}
resultSets = append(resultSets, connection.ResultSetData{
Rows: []map[string]interface{}{{"affectedRows": affected}},
Columns: []string{"affectedRows"},
StatementIndex: idx + 1,
StatementIndex: statementIndex,
})
emitObservation(affected, 0, nil)
}
if resultSets == nil {
@@ -401,6 +579,46 @@ func executeManagedSQLTransactionStatements(ctx context.Context, session db.Stat
return resultSets, nil
}
func summarizeManagedSQLResultSet(resultSet connection.ResultSetData) (rowsAffected, rowsReturned int64) {
if !isAffectedRowsResultSet(resultSet) {
return 0, int64(len(resultSet.Rows))
}
for _, row := range resultSet.Rows {
value, ok := row["affectedRows"]
if !ok {
for key, candidate := range row {
if strings.EqualFold(strings.TrimSpace(key), "affectedRows") {
value = candidate
ok = true
break
}
}
}
if !ok {
continue
}
switch typed := value.(type) {
case int:
rowsAffected += int64(typed)
case int32:
rowsAffected += int64(typed)
case int64:
rowsAffected += typed
case uint:
rowsAffected += int64(typed)
case uint32:
rowsAffected += int64(typed)
case uint64:
if typed <= uint64(^uint64(0)>>1) {
rowsAffected += int64(typed)
}
case float64:
rowsAffected += int64(typed)
}
}
return rowsAffected, 0
}
func shouldUseManagedSQLTransaction(dbType string, query string) bool {
if isManagedSQLTransactionUnsupportedType(dbType) {
return false
@@ -500,15 +718,24 @@ func isOracleLikeAnonymousBlockManagedWrite(dbType string, stmt string) bool {
}
func (a *App) DBCommitTransaction(transactionID string) connection.QueryResult {
return a.finishManagedSQLTransaction(transactionID, true)
return a.finishManagedSQLTransaction(transactionID, true, "manual")
}
func (a *App) DBRollbackTransaction(transactionID string) connection.QueryResult {
return a.finishManagedSQLTransaction(transactionID, false)
return a.finishManagedSQLTransaction(transactionID, false, "manual")
}
func (a *App) finishManagedSQLTransaction(transactionID string, commit bool) connection.QueryResult {
func (a *App) DBCommitTransactionWithTrigger(transactionID string, trigger string) connection.QueryResult {
return a.finishManagedSQLTransaction(transactionID, true, trigger)
}
func (a *App) DBRollbackTransactionWithTrigger(transactionID string, trigger string) connection.QueryResult {
return a.finishManagedSQLTransaction(transactionID, false, trigger)
}
func (a *App) finishManagedSQLTransaction(transactionID string, commit bool, trigger string) connection.QueryResult {
transactionID = strings.TrimSpace(transactionID)
trigger = normalizeSQLTransactionFinishTrigger(trigger)
if transactionID == "" {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.transaction_id_required", nil)}
}
@@ -522,19 +749,53 @@ func (a *App) finishManagedSQLTransaction(transactionID string, commit bool) con
if !ok || tx == nil || tx.execer == nil {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.transaction_not_found", nil)}
}
tx.mu.Lock()
defer tx.mu.Unlock()
if tx.finished || tx.execer == nil {
return connection.QueryResult{Success: false, Message: a.appText("db.backend.error.transaction_not_found", nil)}
}
tx.finished = true
if tx.cancel != nil {
defer tx.cancel()
}
actionCode := "rollback"
sqlText := tx.rollbackSQL
eventType := "transaction_rollback"
if commit {
actionCode = "commit"
sqlText = tx.commitSQL
eventType = "transaction_commit"
} else if trigger == "tab_close" || trigger == "auto" {
eventType = "transaction_auto_rollback"
}
auditSource := "query_editor"
if trigger == "tab_close" {
auditSource = "tab_close"
}
commitMode := "manual"
if trigger == "auto" {
commitMode = "auto"
}
ctx, cancel := context.WithTimeout(context.Background(), sqlEditorTransactionFinishTimeout)
defer cancel()
requestedEventType := "transaction_rollback_requested"
if commit {
requestedEventType = "transaction_commit_requested"
}
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: tx.config,
Database: tx.config.Database,
DBType: tx.dbType,
TransactionID: transactionID,
EventType: requestedEventType,
Status: "success",
Source: auditSource,
CommitMode: commitMode,
BoundaryMode: tx.boundaryMode,
})
startedAt := time.Now()
var execErr error
if tx.transactor != nil {
@@ -548,6 +809,19 @@ func (a *App) finishManagedSQLTransaction(transactionID string, commit bool) con
}
closeErr := tx.execer.Close()
if execErr != nil {
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: tx.config,
Database: tx.config.Database,
DBType: tx.dbType,
TransactionID: transactionID,
EventType: eventType,
Status: "error",
Source: auditSource,
CommitMode: commitMode,
BoundaryMode: tx.boundaryMode,
Duration: time.Since(startedAt),
Err: execErr,
})
logger.Error(execErr, "SQL 编辑器事务%s失败id=%s dbType=%s", actionCode, transactionID, tx.dbType)
key := "db.backend.error.transaction_rollback_failed"
if commit {
@@ -556,6 +830,21 @@ func (a *App) finishManagedSQLTransaction(transactionID string, commit bool) con
return connection.QueryResult{Success: false, Message: a.appText(key, map[string]any{"detail": execErr.Error()})}
}
if closeErr != nil {
// Commit/Rollback has already succeeded at the database boundary. Record that
// outcome as success while retaining the local session cleanup error.
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: tx.config,
Database: tx.config.Database,
DBType: tx.dbType,
TransactionID: transactionID,
EventType: eventType,
Status: "success",
Source: auditSource,
CommitMode: commitMode,
BoundaryMode: tx.boundaryMode,
Duration: time.Since(startedAt),
Err: closeErr,
})
logger.Error(closeErr, "SQL 编辑器事务%s后关闭会话失败id=%s dbType=%s", actionCode, transactionID, tx.dbType)
key := "db.backend.error.transaction_rollback_close_failed"
if commit {
@@ -563,6 +852,18 @@ func (a *App) finishManagedSQLTransaction(transactionID string, commit bool) con
}
return connection.QueryResult{Success: false, Message: a.appText(key, map[string]any{"detail": closeErr.Error()})}
}
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: tx.config,
Database: tx.config.Database,
DBType: tx.dbType,
TransactionID: transactionID,
EventType: eventType,
Status: "success",
Source: auditSource,
CommitMode: commitMode,
BoundaryMode: tx.boundaryMode,
Duration: time.Since(startedAt),
})
if commit {
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.transaction_committed", nil)}
@@ -570,6 +871,15 @@ func (a *App) finishManagedSQLTransaction(transactionID string, commit bool) con
return connection.QueryResult{Success: true, Message: a.appText("db.backend.message.transaction_rolled_back", nil)}
}
func normalizeSQLTransactionFinishTrigger(trigger string) string {
switch strings.ToLower(strings.TrimSpace(trigger)) {
case "auto", "tab_close":
return strings.ToLower(strings.TrimSpace(trigger))
default:
return "manual"
}
}
func (a *App) rollbackPendingSQLTransactionsOnShutdown() {
a.sqlTransactionMu.Lock()
pending := make([]*managedSQLTransaction, 0, len(a.sqlTransactions))
@@ -582,13 +892,34 @@ func (a *App) rollbackPendingSQLTransactionsOnShutdown() {
a.sqlTransactionMu.Unlock()
for _, tx := range pending {
tx.mu.Lock()
if tx.finished {
tx.mu.Unlock()
continue
}
tx.finished = true
ctx, cancel := context.WithTimeout(context.Background(), sqlEditorTransactionFinishTimeout)
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: tx.config,
Database: tx.config.Database,
DBType: tx.dbType,
TransactionID: tx.id,
EventType: "transaction_rollback_requested",
Status: "success",
Source: "app_shutdown",
CommitMode: "auto",
BoundaryMode: tx.boundaryMode,
})
startedAt := time.Now()
var rollbackErr error
if tx.transactor != nil {
if err := tx.transactor.Rollback(); err != nil {
rollbackErr = err
logger.Warnf("关闭应用时回滚 SQL 编辑器事务失败id=%s dbType=%s err=%v", tx.id, tx.dbType, err)
}
} else if strings.TrimSpace(tx.rollbackSQL) != "" && tx.execer != nil {
if _, err := tx.execer.ExecContext(ctx, tx.rollbackSQL); err != nil {
rollbackErr = err
logger.Warnf("关闭应用时回滚 SQL 编辑器事务失败id=%s dbType=%s err=%v", tx.id, tx.dbType, err)
}
}
@@ -596,10 +927,32 @@ func (a *App) rollbackPendingSQLTransactionsOnShutdown() {
if tx.cancel != nil {
tx.cancel()
}
var closeErr error
if tx.execer != nil {
if err := tx.execer.Close(); err != nil {
closeErr = err
logger.Warnf("关闭应用时关闭 SQL 编辑器事务会话失败id=%s dbType=%s err=%v", tx.id, tx.dbType, err)
}
}
auditErr := rollbackErr
status := sqlAuditStatusFromError(rollbackErr)
if auditErr == nil && closeErr != nil {
// The database rollback succeeded; retain only the cleanup warning.
auditErr = closeErr
}
a.recordSQLAuditTransactionEvent(sqlAuditTransactionEventInput{
Config: tx.config,
Database: tx.config.Database,
DBType: tx.dbType,
TransactionID: tx.id,
EventType: "transaction_auto_rollback",
Status: status,
Source: "app_shutdown",
CommitMode: "auto",
BoundaryMode: tx.boundaryMode,
Duration: time.Since(startedAt),
Err: auditErr,
})
tx.mu.Unlock()
}
}

View File

@@ -3,7 +3,9 @@ package app
import (
"bufio"
"context"
"crypto/sha256"
"encoding/csv"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
@@ -1548,7 +1550,14 @@ func executeSQLFileStream(ctx context.Context, dbInst db.Database, reader io.Rea
// ExecuteSQLFile 在后端流式读取并执行大 SQL 文件,通过事件推送进度。
// 前端通过 EventsOn("sqlfile:progress", ...) 监听进度。
func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string, filePath string, jobID string) connection.QueryResult {
func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string, filePath string, jobID string) (result connection.QueryResult) {
auditSQL := "EXECUTE SQL FILE"
auditStatementCount := 0
auditSafeError := "SQL file task failed before an execution summary was available"
defer a.beginSQLAuditUserActionWithOptions(config, dbName, "sql_file", &auditSQL, &result, sqlAuditUserActionOptions{
StatementCount: &auditStatementCount,
SafeError: &auditSafeError,
})()
if strings.TrimSpace(filePath) == "" {
return connection.QueryResult{Success: false, Message: a.appText("file.backend.error.file_path_empty", nil)}
}
@@ -1574,8 +1583,12 @@ func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string,
defer f.Close()
// 获取文件大小用于计算进度
fi, _ := f.Stat()
totalSize := fi.Size()
var totalSize int64
totalSizeKnown := false
if fi, statErr := f.Stat(); statErr == nil {
totalSize = fi.Size()
totalSizeKnown = true
}
// 设置取消上下文
ctx, cancel := context.WithCancel(context.Background())
@@ -1619,7 +1632,8 @@ func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string,
emitProgress("running", 0, 0, 0, 0, "", "")
// 使用 countingReader 追踪已读取字节数
cr := &countingReader{r: f}
fileDigest := sha256.New()
cr := &countingReader{r: io.TeeReader(f, fileDigest)}
startTime := time.Now()
execResult, streamErr := executeSQLFileStream(ctx, dbInst, cr, sqlFileExecutionOptions{
@@ -1644,6 +1658,12 @@ func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string,
executedCount := execResult.Executed
failedCount := execResult.Failed
errorLogs := execResult.Errors
auditStatementCount = executedCount + failedCount
auditSQL = fmt.Sprintf("EXECUTE SQL FILE EXECUTED_%d FAILED_%d", executedCount, failedCount)
auditSafeError = fmt.Sprintf("SQL file task failed after executing %d statement(s); %d statement(s) failed", executedCount, failedCount)
if totalSizeKnown && cr.n == totalSize {
auditSQL += " SHA256_" + hex.EncodeToString(fileDigest.Sum(nil))
}
if streamErr != nil && streamErr.Error() == "已取消" {
emitProgress("cancelled", executedCount, failedCount, executedCount+failedCount, cr.n, "", a.appText("file.backend.message.user_cancelled", nil))
@@ -2364,7 +2384,21 @@ func formatImportSQLValue(dbType, columnType string, value interface{}) string {
}
// ImportDataWithProgress 执行导入并发送进度事件
func (a *App) ImportDataWithProgress(config connection.ConnectionConfig, dbName, tableName, filePath string) connection.QueryResult {
func (a *App) ImportDataWithProgress(config connection.ConnectionConfig, dbName, tableName, filePath string) (result connection.QueryResult) {
dbType := resolveDDLDBType(config)
schemaName, pureTableName := normalizeSchemaAndTable(config, dbName, tableName)
auditTarget := strings.TrimSpace(tableName)
if pureTableName != "" {
auditTarget = quoteTableIdentByType(dbType, schemaName, pureTableName)
}
if auditTarget == "" {
auditTarget = "TARGET_TABLE"
}
auditSQL := "IMPORT DATA INTO " + auditTarget
auditSafeError := "data import task failed"
defer a.beginSQLAuditUserActionWithOptions(config, dbName, "data_import", &auditSQL, &result, sqlAuditUserActionOptions{
SafeError: &auditSafeError,
})()
if err := ensureConnectionAllowsDataImport(config, "connection.backend.action.import_data"); err != nil {
return connection.QueryResult{Success: false, Message: err.Error()}
}
@@ -2374,8 +2408,6 @@ func (a *App) ImportDataWithProgress(config connection.ConnectionConfig, dbName,
return connection.QueryResult{Success: false, Message: err.Error()}
}
dbType := resolveDDLDBType(config)
schemaName, pureTableName := normalizeSchemaAndTable(config, dbName, tableName)
columnTypeMap := map[string]string{}
if defs, colErr := dbInst.GetColumns(schemaName, pureTableName); colErr == nil {
columnTypeMap = buildImportColumnTypeMap(defs)
@@ -2406,19 +2438,22 @@ func (a *App) ImportDataWithProgress(config connection.ConnectionConfig, dbName,
"imported": resultData.Success,
"failed": resultData.Failed,
})
result := map[string]interface{}{
resultPayload := map[string]interface{}{
"success": resultData.Success,
"failed": resultData.Failed,
"total": resultData.Total,
"affectedRows": int64(resultData.Success),
"errorLogs": resultData.ErrorLogs,
"errorSummary": summary,
}
maybeReleaseFileTransferMemory("import-finished", int64(resultData.Total), filePath)
return connection.QueryResult{Success: true, Data: result, Message: summary}
return connection.QueryResult{Success: true, Data: resultPayload, Message: summary}
}
func (a *App) ApplyChanges(config connection.ConnectionConfig, dbName, tableName string, changes connection.ChangeSet) connection.QueryResult {
func (a *App) ApplyChanges(config connection.ConnectionConfig, dbName, tableName string, changes connection.ChangeSet) (result connection.QueryResult) {
auditSQL := fmt.Sprintf("APPLY CHANGES TO %s", strings.TrimSpace(tableName))
defer a.beginSQLAuditUserAction(config, dbName, "data_editor", &auditSQL, &result)()
if err := ensureConnectionAllowsDataEdit(config, "connection.backend.action.apply_result_changes"); err != nil {
return connection.QueryResult{Success: false, Message: err.Error()}
}
@@ -2983,7 +3018,13 @@ func tableDataClearMessageKeys(mode tableDataClearMode, partial bool) (failureKe
}
}
func (a *App) runTableDataClear(config connection.ConnectionConfig, dbName string, tableNames []string, mode tableDataClearMode) connection.QueryResult {
func (a *App) runTableDataClear(config connection.ConnectionConfig, dbName string, tableNames []string, mode tableDataClearMode) (result connection.QueryResult) {
auditAction := "DELETE TABLE DATA"
if mode == tableDataClearModeTruncate {
auditAction = "TRUNCATE TABLE DATA"
}
auditSQL := auditAction + " " + strings.Join(tableNames, ", ")
defer a.beginSQLAuditUserAction(config, dbName, "object_editor", &auditSQL, &result)()
actionLabel, progressLabel := tableDataClearActionLabels(mode)
if err := ensureConnectionAllowsDataEdit(config, actionLabel); err != nil {
return connection.QueryResult{Success: false, Message: err.Error()}

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,678 @@
package app
import (
"errors"
"os"
"path/filepath"
"sort"
"strings"
"testing"
"time"
"GoNavi-Wails/internal/connection"
"GoNavi-Wails/internal/db"
"GoNavi-Wails/internal/secretstore"
"GoNavi-Wails/internal/sqlaudit"
)
type sqlAuditTestDatabase struct {
rows []map[string]interface{}
columns []string
queryErr error
affected int64
connected bool
}
func (database *sqlAuditTestDatabase) Connect(connection.ConnectionConfig) error {
database.connected = true
return nil
}
func (database *sqlAuditTestDatabase) Close() error { return nil }
func (database *sqlAuditTestDatabase) Ping() error { return nil }
func (database *sqlAuditTestDatabase) Query(string) ([]map[string]interface{}, []string, error) {
return database.rows, database.columns, database.queryErr
}
func (database *sqlAuditTestDatabase) Exec(string) (int64, error) {
return database.affected, database.queryErr
}
func (database *sqlAuditTestDatabase) GetDatabases() ([]string, error) { return nil, nil }
func (database *sqlAuditTestDatabase) GetTables(string) ([]string, error) {
return nil, nil
}
func (database *sqlAuditTestDatabase) GetCreateStatement(string, string) (string, error) {
return "", nil
}
func (database *sqlAuditTestDatabase) GetColumns(string, string) ([]connection.ColumnDefinition, error) {
return nil, nil
}
func (database *sqlAuditTestDatabase) GetAllColumns(string) ([]connection.ColumnDefinitionWithTable, error) {
return nil, nil
}
func (database *sqlAuditTestDatabase) GetIndexes(string, string) ([]connection.IndexDefinition, error) {
return nil, nil
}
func (database *sqlAuditTestDatabase) GetForeignKeys(string, string) ([]connection.ForeignKeyDefinition, error) {
return nil, nil
}
func (database *sqlAuditTestDatabase) GetTriggers(string, string) ([]connection.TriggerDefinition, error) {
return nil, nil
}
func newSQLAuditTestApp(t *testing.T) *App {
t.Helper()
app := NewAppWithSecretStore(secretstore.NewUnavailableStore("test"))
app.configDir = t.TempDir()
app.activateSQLAudit()
t.Cleanup(func() { app.closeSQLAuditStore() })
return app
}
func loadSQLAuditEvents(t *testing.T, app *App, filter sqlaudit.Filter) []sqlaudit.Event {
t.Helper()
filter.PageSize = 500
result := app.GetSQLAuditEvents(filter)
if !result.Success {
t.Fatalf("GetSQLAuditEvents returned failure: %s", result.Message)
}
page, ok := result.Data.(sqlaudit.Page)
if !ok {
t.Fatalf("GetSQLAuditEvents data type = %T, want sqlaudit.Page", result.Data)
}
sort.Slice(page.Items, func(left, right int) bool {
return page.Items[left].Sequence < page.Items[right].Sequence
})
return page.Items
}
func TestSQLAuditConnectionFingerprintExcludesSecrets(t *testing.T) {
base := connection.ConnectionConfig{
Type: "postgres",
Host: "db.internal",
Port: 5432,
User: "alice",
Password: "primary-secret",
Database: "app",
DSN: "postgres://alice:dsn-secret@db.internal/app",
URI: "postgres://alice:uri-secret@db.internal/app",
UseSSH: true,
SSH: connection.SSHConfig{Host: "jump.internal", User: "ops", Password: "ssh-secret"},
UseProxy: true,
Proxy: connection.ProxyConfig{Host: "proxy.internal", User: "proxy", Password: "proxy-secret"},
UseHTTPTunnel: true,
HTTPTunnel: connection.HTTPTunnelConfig{Host: "tunnel.internal", User: "tunnel", Password: "tunnel-secret"},
}
changedSecrets := base
changedSecrets.User = "bob"
changedSecrets.Password = "changed-primary"
changedSecrets.DSN = "postgres://bob:changed-dsn@db.internal/app"
changedSecrets.URI = "postgres://bob:changed-uri@db.internal/app"
changedSecrets.SSH.Password = "changed-ssh"
changedSecrets.Proxy.Password = "changed-proxy"
changedSecrets.HTTPTunnel.Password = "changed-tunnel"
baseFingerprint := buildSQLAuditConnectionFingerprint(base, "app")
if got := buildSQLAuditConnectionFingerprint(changedSecrets, "app"); got != baseFingerprint {
t.Fatalf("secret-only changes altered audit fingerprint: base=%s changed=%s", baseFingerprint, got)
}
changedEndpoint := base
changedEndpoint.Host = "other.internal"
if got := buildSQLAuditConnectionFingerprint(changedEndpoint, "app"); got == baseFingerprint {
t.Fatal("endpoint change should alter audit fingerprint")
}
saved := base
saved.ID = "connection-1"
savedFingerprint := buildSQLAuditConnectionFingerprint(saved, "app")
savedChangedEndpoint := saved
savedChangedEndpoint.Host = "moved.internal"
if got := buildSQLAuditConnectionFingerprint(savedChangedEndpoint, "app"); got != savedFingerprint {
t.Fatal("saved connection endpoint edit should retain its audit fingerprint")
}
if got := buildSQLAuditConnectionFingerprint(saved, "analytics"); got == savedFingerprint {
t.Fatal("logical database change should alter a saved connection audit fingerprint")
}
dsnOnly := connection.ConnectionConfig{Type: "custom", Driver: "postgres", DSN: "host=opaque-a port=5432 user=alice password=first dbname=app"}
dsnSecretChanged := dsnOnly
dsnSecretChanged.DSN = "host=opaque-a port=5432 user=bob password=second dbname=app"
if got, want := buildSQLAuditConnectionFingerprint(dsnSecretChanged, "app"), buildSQLAuditConnectionFingerprint(dsnOnly, "app"); got != want {
t.Fatal("DSN credential change should not alter temporary connection fingerprint")
}
dsnEndpointChanged := dsnOnly
dsnEndpointChanged.DSN = "host=opaque-b port=5432 user=alice password=first dbname=app"
if got, want := buildSQLAuditConnectionFingerprint(dsnEndpointChanged, "app"), buildSQLAuditConnectionFingerprint(dsnOnly, "app"); got == want {
t.Fatal("safe DSN endpoint change should alter temporary connection fingerprint")
}
uriOnly := connection.ConnectionConfig{Type: "postgres", URI: "postgres://alice:first@db.internal/app?pass=first&key=first"}
uriSecretsChanged := uriOnly
uriSecretsChanged.URI = "postgres://bob:second@db.internal/app?pass=second&key=second&custom_secret=third"
if got, want := buildSQLAuditConnectionFingerprint(uriSecretsChanged, "app"), buildSQLAuditConnectionFingerprint(uriOnly, "app"); got != want {
t.Fatal("URI credentials and query parameters must not alter temporary connection fingerprint")
}
uriEndpointChanged := uriOnly
uriEndpointChanged.URI = "postgres://alice:first@other.internal/app?pass=first"
if got, want := buildSQLAuditConnectionFingerprint(uriEndpointChanged, "app"), buildSQLAuditConnectionFingerprint(uriOnly, "app"); got == want {
t.Fatal("URI authority change should alter temporary connection fingerprint")
}
}
func TestDBQueryWithCancelWritesRedactedSQLAudit(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
database := &sqlAuditTestDatabase{
rows: []map[string]interface{}{{"id": int64(7)}},
columns: []string{"id"},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, Database: "app"}
result := app.DBQueryWithCancel(config, "app", "SELECT * FROM users WHERE email = 'secret@example.test' AND id = 7", "query-audit-1")
if !result.Success {
t.Fatalf("DBQueryWithCancel returned failure: %s", result.Message)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: "query-audit-1"})
if len(events) != 1 {
t.Fatalf("audit event count = %d, want 1: %#v", len(events), events)
}
event := events[0]
if event.EventType != "query" || event.Status != "success" || event.RowsReturned != 1 {
t.Fatalf("unexpected query audit event: %#v", event)
}
if strings.Contains(event.SQLText, "secret@example.test") || strings.Contains(event.SQLText, " 7") {
t.Fatalf("audit SQL leaked literals: %q", event.SQLText)
}
if !event.SQLRedacted || event.ConnectionFingerprint == "" {
t.Fatalf("expected redacted event with connection identity: %#v", event)
}
}
func TestManagedSQLTransactionWritesCompleteAuditTimeline(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
firstStatement := "UPDATE users SET name = 'private-name' WHERE id = 1"
secondStatement := "DELETE FROM audit_logs WHERE user_id = 1"
database := &fakeBatchWriteDB{execAffected: map[string]int64{
firstStatement: 1,
secondStatement: 3,
}}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "main"}
started := app.DBQueryMultiTransactional(config, "main", firstStatement+";\n"+secondStatement+";", "query-tx-audit")
if !started.Success || started.TransactionID == "" {
t.Fatalf("DBQueryMultiTransactional returned %#v", started)
}
committed := app.DBCommitTransactionWithTrigger(started.TransactionID, "auto")
if !committed.Success {
t.Fatalf("DBCommitTransactionWithTrigger returned failure: %s", committed.Message)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{TransactionID: started.TransactionID})
if len(events) != 5 {
t.Fatalf("transaction audit event count = %d, want 5: %#v", len(events), events)
}
wantTypes := []string{"transaction_begin", "transaction_statement", "transaction_statement", "transaction_commit_requested", "transaction_commit"}
for index, wantType := range wantTypes {
if events[index].EventType != wantType || events[index].Status != "success" {
t.Fatalf("event %d = %#v, want type=%s success", index, events[index], wantType)
}
}
if events[1].StatementIndex != 1 || events[1].StatementCount != 2 || events[1].RowsAffected != 1 {
t.Fatalf("unexpected first statement audit metrics: %#v", events[1])
}
if events[2].StatementIndex != 2 || events[2].RowsAffected != 3 {
t.Fatalf("unexpected second statement audit metrics: %#v", events[2])
}
if strings.Contains(events[1].SQLText, "private-name") {
t.Fatalf("transaction audit leaked SQL literal: %q", events[1].SQLText)
}
if events[3].CommitMode != "auto" || events[3].BoundaryMode != "text_sql" {
t.Fatalf("commit request audit lost trigger/boundary metadata: %#v", events[3])
}
if events[4].CommitMode != "auto" || events[4].BoundaryMode != "text_sql" {
t.Fatalf("commit audit lost trigger/boundary metadata: %#v", events[4])
}
}
func TestManagedSQLTransactionFailureAuditsStatementAndAutomaticRollback(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
firstStatement := "UPDATE users SET active = 1 WHERE id = 1"
secondStatement := "DELETE FROM missing_table WHERE id = 1"
database := &fakeBatchWriteDB{
execAffected: map[string]int64{firstStatement: 1},
execErr: map[string]error{secondStatement: errors.New("table 'private_table_name' does not exist")},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "main"}
result := app.DBQueryMultiTransactional(config, "main", firstStatement+";\n"+secondStatement+";", "query-tx-failed")
if result.Success {
t.Fatalf("expected transaction failure, got %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: "query-tx-failed"})
if len(events) != 5 {
t.Fatalf("failed transaction audit event count = %d, want 5: %#v", len(events), events)
}
if events[2].EventType != "transaction_statement" || events[2].Status != "error" {
t.Fatalf("failed statement was not audited: %#v", events[2])
}
if strings.Contains(events[2].Error, "private_table_name") {
t.Fatalf("audit error leaked quoted driver detail: %q", events[2].Error)
}
if events[3].EventType != "transaction_rollback_requested" || events[3].Status != "success" {
t.Fatalf("automatic rollback request was not audited: %#v", events[3])
}
if events[4].EventType != "transaction_auto_rollback" || events[4].Status != "success" {
t.Fatalf("automatic rollback was not audited: %#v", events[4])
}
}
func TestSQLAuditRecordsRealSQLiteExecutions(t *testing.T) {
app := newSQLAuditTestApp(t)
databasePath := filepath.Join(t.TempDir(), "audit-target.sqlite")
config := connection.ConnectionConfig{
Type: "custom",
Driver: "sqlite",
DSN: databasePath,
Database: databasePath,
}
t.Cleanup(func() { app.DBReleaseConnection(config) })
queries := []struct {
id string
sql string
}{
{id: "real-sqlite-create", sql: "CREATE TABLE audit_users (id INTEGER PRIMARY KEY, email TEXT)"},
{id: "real-sqlite-insert", sql: "INSERT INTO audit_users(id, email) VALUES (1, 'private@example.test')"},
{id: "real-sqlite-select", sql: "SELECT id, email FROM audit_users WHERE id = 1"},
}
for _, query := range queries {
result := app.DBQueryMulti(config, databasePath, query.sql, query.id)
if !result.Success {
t.Fatalf("DBQueryMulti(%s) returned failure: %s", query.id, result.Message)
}
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: "real-sqlite-", PageSize: 10})
if len(events) != len(queries) {
t.Fatalf("real SQLite audit event count = %d, want %d: %#v", len(events), len(queries), events)
}
for _, event := range events {
if event.Status != "success" || event.DBType != "sqlite" {
t.Fatalf("unexpected real SQLite audit event: %#v", event)
}
if strings.Contains(event.SQLText, "private@example.test") {
t.Fatalf("real SQLite audit leaked literal: %q", event.SQLText)
}
}
if events[1].RowsAffected != 1 {
t.Fatalf("real SQLite INSERT affected rows = %d, want 1", events[1].RowsAffected)
}
if events[2].RowsReturned != 1 {
t.Fatalf("real SQLite SELECT returned rows = %d, want 1", events[2].RowsReturned)
}
}
func TestSQLAuditHealthRecordsAndClosesPersistenceGap(t *testing.T) {
app := newSQLAuditTestApp(t)
wasActive, suspendErr := app.suspendSQLAudit()
if suspendErr != nil {
t.Fatalf("suspendSQLAudit returned error: %v", suspendErr)
}
if !wasActive {
t.Fatal("expected SQL audit runtime to be active before suspension")
}
app.appendSQLAuditEvent(sqlaudit.Event{EventType: "query", Status: "success", QueryID: "lost-query"})
degradedResult := app.GetSQLAuditHealth()
degraded, ok := degradedResult.Data.(sqlAuditHealthState)
if !ok {
t.Fatalf("GetSQLAuditHealth data type = %T, want sqlAuditHealthState", degradedResult.Data)
}
if degraded.Status != sqlAuditHealthStatusDegraded || degraded.DroppedEvents != 1 {
t.Fatalf("unexpected degraded SQL audit health: %#v", degraded)
}
app.resumeSQLAudit(wasActive)
app.appendSQLAuditEvent(sqlaudit.Event{EventType: "query", Status: "success", QueryID: "recovered-query"})
healthResult := app.GetSQLAuditHealth()
health := healthResult.Data.(sqlAuditHealthState)
if health.Status != sqlAuditHealthStatusHealthy || health.DroppedEvents != 1 || health.LastSuccessAt == 0 {
t.Fatalf("unexpected recovered SQL audit health: %#v", health)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 2 || events[0].QueryID != "recovered-query" || events[1].EventType != "audit_gap" {
t.Fatalf("recovered audit timeline did not persist gap marker and next event: %#v", events)
}
}
func TestSuspendSQLAuditReturnsCheckpointFailureAndCanResume(t *testing.T) {
app := newSQLAuditTestApp(t)
originalClose := closeSQLAuditStoreHandle
closeSQLAuditStoreHandle = func(store *sqlaudit.Store) error {
return errors.Join(originalClose(store), errors.New("simulated checkpoint failure"))
}
wasActive, err := app.suspendSQLAudit()
closeSQLAuditStoreHandle = originalClose
if err == nil || !strings.Contains(err.Error(), "simulated checkpoint failure") {
t.Fatalf("suspendSQLAudit error = %v, want checkpoint failure", err)
}
if !wasActive {
t.Fatal("failed suspension lost the prior active state")
}
app.resumeSQLAudit(wasActive)
app.appendSQLAuditEvent(sqlaudit.Event{EventType: "query", Status: "success", QueryID: "after-failed-suspend"})
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: "after-failed-suspend"})
if len(events) != 1 {
t.Fatalf("audit store did not resume after failed suspension: %#v", events)
}
}
func TestSQLAuditHealthReportsCaptureStateAndMode(t *testing.T) {
app := newSQLAuditTestApp(t)
health := app.GetSQLAuditHealth().Data.(sqlAuditHealthState)
if health.CaptureEnabled == nil || !*health.CaptureEnabled {
t.Fatalf("default capture state was not reported as enabled: %#v", health)
}
if health.CaptureMode != sqlaudit.CaptureModeRedacted {
t.Fatalf("default capture mode = %q, want %q", health.CaptureMode, sqlaudit.CaptureModeRedacted)
}
settings := sqlaudit.DefaultSettings()
settings.Enabled = false
settings.CaptureMode = sqlaudit.CaptureModeMetadata
if result := app.UpdateSQLAuditSettings(settings); !result.Success {
t.Fatalf("UpdateSQLAuditSettings returned failure: %s", result.Message)
}
health = app.GetSQLAuditHealth().Data.(sqlAuditHealthState)
if health.CaptureEnabled == nil || *health.CaptureEnabled {
t.Fatalf("disabled capture state was not reported explicitly: %#v", health)
}
if health.CaptureMode != sqlaudit.CaptureModeMetadata {
t.Fatalf("capture mode = %q, want %q", health.CaptureMode, sqlaudit.CaptureModeMetadata)
}
}
func TestSQLAuditRecoveryRetainsGapMarkerAtMinimumRecordLimit(t *testing.T) {
app := newSQLAuditTestApp(t)
settingsResult := app.UpdateSQLAuditSettings(sqlaudit.Settings{
Enabled: true,
CaptureMode: sqlaudit.CaptureModeRedacted,
RetentionDays: 30,
MaxRecords: 1,
})
if !settingsResult.Success {
t.Fatalf("UpdateSQLAuditSettings returned failure: %s", settingsResult.Message)
}
wasActive, suspendErr := app.suspendSQLAudit()
if suspendErr != nil {
t.Fatalf("suspendSQLAudit returned error: %v", suspendErr)
}
app.appendSQLAuditEvent(sqlaudit.Event{EventType: "query", Status: "success", QueryID: "lost-at-minimum-limit"})
app.resumeSQLAudit(wasActive)
app.appendSQLAuditEvent(sqlaudit.Event{EventType: "query", Status: "success", QueryID: "current-at-minimum-limit"})
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 1 || events[0].EventType != "audit_gap" {
t.Fatalf("minimum retention limit must keep the recovery marker: %#v", events)
}
health := app.GetSQLAuditHealth().Data.(sqlAuditHealthState)
if health.Status != sqlAuditHealthStatusHealthy || health.DroppedEvents != 1 {
t.Fatalf("unexpected health after durable minimum-limit marker: %#v", health)
}
}
func TestSQLAuditOversizedBatchCreatesVisibleHealthGapInsteadOfSilentTail(t *testing.T) {
app := newSQLAuditTestApp(t)
settingsResult := app.UpdateSQLAuditSettings(sqlaudit.Settings{
Enabled: true,
CaptureMode: sqlaudit.CaptureModeRedacted,
RetentionDays: 30,
MaxRecords: 1,
})
if !settingsResult.Success {
t.Fatalf("UpdateSQLAuditSettings returned failure: %s", settingsResult.Message)
}
app.appendSQLAuditEvents([]sqlaudit.Event{
{EventType: "transaction_statement", Status: "success", QueryID: "oversized-1"},
{EventType: "transaction_statement", Status: "success", QueryID: "oversized-2"},
})
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 1 || events[0].EventType != "audit_settings_change" {
t.Fatalf("oversized batch must fail atomically while retaining its settings boundary, got %#v", events)
}
health := app.GetSQLAuditHealth().Data.(sqlAuditHealthState)
if health.Status != sqlAuditHealthStatusDegraded || health.DroppedEvents != 2 {
t.Fatalf("oversized batch was not exposed as a health gap: %#v", health)
}
}
func TestSQLAuditControlEventsSurviveDisableAndClear(t *testing.T) {
app := newSQLAuditTestApp(t)
disabled := sqlaudit.Settings{
Enabled: false,
CaptureMode: sqlaudit.CaptureModeRedacted,
RetentionDays: 30,
MaxRecords: 100,
}
if result := app.UpdateSQLAuditSettings(disabled); !result.Success {
t.Fatalf("disable SQL audit returned failure: %s", result.Message)
}
app.appendSQLAuditEvent(sqlaudit.Event{EventType: "query", Status: "success", QueryID: "disabled-query"})
enabled := disabled
enabled.Enabled = true
if result := app.UpdateSQLAuditSettings(enabled); !result.Success {
t.Fatalf("enable SQL audit returned failure: %s", result.Message)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 2 || events[0].EventType != "audit_settings_change" || events[1].EventType != "audit_settings_change" {
t.Fatalf("disable/enable control boundaries were not persisted: %#v", events)
}
for _, event := range events {
if event.QueryID == "disabled-query" {
t.Fatalf("ordinary event was persisted while auditing was disabled: %#v", events)
}
}
if result := app.ClearSQLAuditEvents(0); !result.Success {
t.Fatalf("ClearSQLAuditEvents returned failure: %s", result.Message)
}
events = loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 1 || events[0].EventType != "audit_clear" || events[0].RowsAffected != 2 {
t.Fatalf("clear boundary did not replace deleted history with a control event: %#v", events)
}
}
func TestSQLAuditSettingsControlCanRecoverDegradedWriterWhileDisablingCapture(t *testing.T) {
app := newSQLAuditTestApp(t)
app.markSQLAuditFailure(1, errors.New("simulated writer failure"))
settings := sqlaudit.DefaultSettings()
settings.Enabled = false
if result := app.UpdateSQLAuditSettings(settings); !result.Success {
t.Fatalf("disable after writer recovery returned failure: %s", result.Message)
}
health := app.GetSQLAuditHealth().Data.(sqlAuditHealthState)
if health.Status != sqlAuditHealthStatusHealthy || health.DroppedEvents != 1 || health.CaptureEnabled == nil || *health.CaptureEnabled {
t.Fatalf("control write did not close the degraded state: %#v", health)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{})
if len(events) != 2 || events[0].EventType != "audit_settings_change" || events[1].EventType != "audit_gap" {
t.Fatalf("settings recovery lacks control and gap boundaries: %#v", events)
}
if result := app.BuildSQLAuditExport(sqlaudit.Filter{}, "json"); !result.Success {
t.Fatalf("recovered disabled audit history should remain exportable: %s", result.Message)
}
}
func TestWriteSQLAuditExportPreservesExistingFileWhenAtomicReplacementFails(t *testing.T) {
directory := t.TempDir()
target := filepath.Join(directory, "audit.json")
if err := os.WriteFile(target, []byte("original"), 0o600); err != nil {
t.Fatalf("write original export: %v", err)
}
originalReplace := replaceSQLAuditFile
t.Cleanup(func() { replaceSQLAuditFile = originalReplace })
replaceSQLAuditFile = func(_, _ string) error {
return errors.New("simulated atomic replacement failure")
}
if err := writeSQLAuditExportAtomically(target, []byte("replacement")); err == nil {
t.Fatal("expected replacement failure")
}
content, err := os.ReadFile(target)
if err != nil {
t.Fatalf("read preserved export: %v", err)
}
if string(content) != "original" {
t.Fatalf("existing export was not preserved: %q", content)
}
}
func TestExportSQLAuditFileRejectsWebRuntimeBeforeOpeningDesktopDialog(t *testing.T) {
app := NewWebApp()
app.configDir = t.TempDir()
app.activateSQLAudit()
t.Cleanup(func() { app.closeSQLAuditStore() })
result := app.ExportSQLAuditFile(sqlaudit.Filter{}, "json")
if result.Success || !strings.Contains(result.Message, "BuildSQLAuditExport") {
t.Fatalf("web runtime desktop export result = %#v, want safe rejection", result)
}
}
func TestSQLAuditExportTargetRejectsInternalStorageFiles(t *testing.T) {
app := newSQLAuditTestApp(t)
for _, protectedPath := range []string{
app.sqlAuditDatabasePath(),
app.sqlAuditDatabasePath() + "-wal",
app.sqlAuditDatabasePath() + "-shm",
app.sqlAuditHealthFilePath(),
} {
if err := app.validateSQLAuditExportTarget(protectedPath); err == nil {
t.Fatalf("protected audit export target %q was accepted", protectedPath)
}
}
if err := app.validateSQLAuditExportTarget(filepath.Join(t.TempDir(), "safe-export.json")); err != nil {
t.Fatalf("safe audit export target was rejected: %v", err)
}
}
func TestSQLAuditExportTargetRejectsMissingSidecarThroughSymlinkedParent(t *testing.T) {
app := newSQLAuditTestApp(t)
auditDirectory := filepath.Dir(app.sqlAuditDatabasePath())
aliasDirectory := filepath.Join(t.TempDir(), "audit-alias")
if err := os.Symlink(auditDirectory, aliasDirectory); err != nil {
t.Skipf("creating a directory symlink is unavailable: %v", err)
}
candidate := filepath.Join(aliasDirectory, filepath.Base(app.sqlAuditHealthFilePath()))
_ = os.Remove(app.sqlAuditHealthFilePath())
if err := app.validateSQLAuditExportTarget(candidate); err == nil {
t.Fatalf("symlinked missing health export target %q was accepted", candidate)
}
}
func TestDBQueryMultiAuditsSuccessfulPrefixBeforeLaterStatementFailure(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
firstStatement := "UPDATE users SET active = 1 WHERE id = 1"
secondStatement := "DELETE FROM missing_table WHERE id = 2"
database := &fakeBatchWriteDB{
execAffected: map[string]int64{firstStatement: 3},
execErr: map[string]error{secondStatement: errors.New("second statement failed")},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "main"}
result := app.DBQueryMulti(config, "main", firstStatement+";\n"+secondStatement+";", "query-partial-audit")
if result.Success {
t.Fatalf("expected second statement failure, got %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: "query-partial-audit"})
if len(events) != 3 {
t.Fatalf("partial batch audit event count = %d, want 3: %#v", len(events), events)
}
if events[0].EventType != "query_statement" || events[0].Status != "success" ||
events[0].StatementIndex != 1 || events[0].RowsAffected != 3 {
t.Fatalf("successful committed prefix was not audited: %#v", events[0])
}
if events[1].EventType != "query_statement" || events[1].Status != "error" ||
events[1].StatementIndex != 2 || !strings.Contains(events[1].Error, "second statement failed") {
t.Fatalf("failed statement was not audited: %#v", events[1])
}
if events[2].EventType != "query" || events[2].Status != "error" || events[2].StatementCount != 2 {
t.Fatalf("batch summary was not retained after statement events: %#v", events[2])
}
}
func TestManagedSQLTransactionStatementAuditUsesActualCompletionTimes(t *testing.T) {
originalNewDatabaseFunc := newDatabaseFunc
t.Cleanup(func() { newDatabaseFunc = originalNewDatabaseFunc })
firstStatement := "UPDATE users SET active = 1 WHERE id = 1"
secondStatement := "UPDATE users SET active = 0 WHERE id = 2"
database := &fakeBatchWriteDB{
execAffected: map[string]int64{firstStatement: 1, secondStatement: 1},
execDelay: map[string]time.Duration{
firstStatement: 25 * time.Millisecond,
secondStatement: 25 * time.Millisecond,
},
}
newDatabaseFunc = func(string) (db.Database, error) { return database, nil }
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{Type: "mysql", Host: "127.0.0.1", Port: 3306, Database: "main"}
started := app.DBQueryMultiTransactional(config, "main", firstStatement+";\n"+secondStatement+";", "query-timestamp-audit")
if !started.Success || started.TransactionID == "" {
t.Fatalf("start managed transaction: %#v", started)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{TransactionID: started.TransactionID})
statements := make([]sqlaudit.Event, 0, 2)
for _, event := range events {
if event.EventType == "transaction_statement" {
statements = append(statements, event)
}
}
if len(statements) != 2 {
t.Fatalf("transaction statement events = %d, want 2: %#v", len(statements), events)
}
if statements[0].Timestamp <= 0 || statements[1].Timestamp <= statements[0].Timestamp {
t.Fatalf("statement completion timestamps were collapsed at batch flush: %#v", statements)
}
if statements[0].DurationMs < 20 || statements[1].DurationMs < 20 {
t.Fatalf("statement durations do not reflect execution time: %#v", statements)
}
if rollback := app.DBRollbackTransaction(started.TransactionID); !rollback.Success {
t.Fatalf("rollback managed transaction: %#v", rollback)
}
}
func TestManagedSQLTransactionProtectionDenialIsAudited(t *testing.T) {
app := newSQLAuditTestApp(t)
config := connection.ConnectionConfig{
Type: "mysql",
Host: "127.0.0.1",
Port: 3306,
Database: "main",
ReadOnly: true,
}
result := app.DBQueryMultiTransactional(config, "main", "UPDATE users SET active = 1 WHERE id = 9", "query-denied-audit")
if result.Success {
t.Fatalf("expected production protection denial, got %#v", result)
}
events := loadSQLAuditEvents(t, app, sqlaudit.Filter{Search: "query-denied-audit"})
if len(events) != 1 {
t.Fatalf("denied managed transaction audit events = %d, want 1: %#v", len(events), events)
}
if events[0].EventType != "transaction_begin" || events[0].Status != "error" || events[0].Error == "" {
t.Fatalf("denied managed transaction was not audited: %#v", events[0])
}
}

View File

@@ -81,7 +81,34 @@ func (a *App) resolveDataSyncEndpointConfig(raw connection.ConnectionConfig, sel
}
// DataSync executes a data synchronization task
func (a *App) DataSync(config sync.SyncConfig) sync.SyncResult {
func (a *App) DataSync(config sync.SyncConfig) (result sync.SyncResult) {
auditStartedAt := time.Now()
defer func() {
runConfig := normalizeRunConfig(config.TargetConfig, config.TargetDatabase)
auditMessage := ""
if !result.Success {
auditMessage = "data synchronization task failed"
}
auditResult := connection.QueryResult{
Success: result.Success,
Message: auditMessage,
Data: map[string]int64{
"affectedRows": int64(result.RowsInserted + result.RowsUpdated + result.RowsDeleted),
},
}
a.recordSQLAuditQuery(sqlAuditQueryInput{
Config: runConfig,
Database: config.TargetDatabase,
DBType: resolveDDLDBType(runConfig),
QueryID: generateQueryID(),
SQL: fmt.Sprintf("SYNC DATA TABLES_%d", len(config.Tables)),
Source: "sync",
CommitMode: "auto",
Duration: time.Since(auditStartedAt),
StatementCount: len(config.Tables),
Result: auditResult,
})
}()
if err := ensureDataSyncTargetProtection(config); err != nil {
return sync.SyncResult{
Success: false,

View File

@@ -0,0 +1,20 @@
//go:build !windows
package app
import (
"errors"
"os"
"path/filepath"
)
func atomicReplaceSQLAuditFile(source, target string) error {
if err := os.Rename(source, target); err != nil {
return err
}
directory, err := os.Open(filepath.Dir(target))
if err != nil {
return err
}
return errors.Join(directory.Sync(), directory.Close())
}

View File

@@ -0,0 +1,21 @@
//go:build windows
package app
import "golang.org/x/sys/windows"
func atomicReplaceSQLAuditFile(source, target string) error {
sourcePath, err := windows.UTF16PtrFromString(source)
if err != nil {
return err
}
targetPath, err := windows.UTF16PtrFromString(target)
if err != nil {
return err
}
return windows.MoveFileEx(
sourcePath,
targetPath,
windows.MOVEFILE_REPLACE_EXISTING|windows.MOVEFILE_WRITE_THROUGH,
)
}

View File

@@ -452,14 +452,164 @@ func isReadOnlySQLQuery(dbType string, query string) bool {
if keyword == "select" && isSQLSelectIntoStatement(query) {
return false
}
if keyword == "explain" && explainAnalyzeMayWrite(query) {
return false
}
if keyword == "pragma" {
return !pragmaMayWrite(query)
}
switch keyword {
case "select", "with", "show", "describe", "desc", "explain", "pragma", "values", "consume":
case "select", "with", "show", "describe", "desc", "explain", "values", "consume":
return true
default:
return false
}
}
func explainAnalyzeMayWrite(query string) bool {
keyword, pos := nextSQLKeyword(query, 0)
if keyword != "explain" {
return false
}
pos = skipSQLTrivia(query, pos)
analyze := false
if pos < len(query) && query[pos] == '(' {
next := skipBalancedSQLParens(query, pos)
if next < 0 {
return false
}
options := query[pos+1 : next-1]
analyze = sqlContainsKeyword(options, "analyze") || sqlContainsKeyword(options, "analyse")
pos = next
} else {
for {
option, next := nextSQLKeyword(query, pos)
switch option {
case "analyze", "analyse":
analyze = true
pos = next
case "verbose":
pos = next
default:
goto optionsDone
}
}
}
optionsDone:
if !analyze {
return false
}
body := query[skipSQLTrivia(query, pos):]
bodyKeyword, withHasWrite := sqlDataOperationInfo(body)
if withHasWrite || isSQLDataWriteKeyword(bodyKeyword) {
return true
}
if bodyKeyword == "select" && isSQLSelectIntoStatement(body) {
return true
}
switch bodyKeyword {
case "create", "execute", "call":
return true
default:
return false
}
}
func pragmaMayWrite(query string) bool {
keyword, pos := nextSQLKeyword(query, 0)
if keyword != "pragma" {
return false
}
name, next, ok := readSQLIdentifierName(query, pos)
if !ok {
return true
}
pos = skipSQLTrivia(query, next)
if pos < len(query) && query[pos] == '.' {
name, next, ok = readSQLIdentifierName(query, pos+1)
if !ok {
return true
}
pos = skipSQLTrivia(query, next)
}
if pos < len(query) && query[pos] == '=' {
return true
}
if pos < len(query) && query[pos] == '(' {
return !isReadOnlyPragmaWithArgument(name)
}
for {
pos = skipSQLTrivia(query, pos)
if pos < len(query) && query[pos] == ';' {
pos++
continue
}
break
}
if pos < len(query) {
return true
}
return !isReadOnlyPragmaWithoutArgument(name)
}
func readSQLIdentifierName(text string, start int) (string, int, bool) {
pos := skipSQLTrivia(text, start)
end, ok := skipSQLIdentifierToken(text, pos)
if !ok || end <= pos {
return "", pos, false
}
token := text[pos:end]
switch token[0] {
case '"', '`':
if len(token) < 2 {
return "", end, false
}
delimiter := string(token[0])
token = strings.ReplaceAll(token[1:len(token)-1], delimiter+delimiter, delimiter)
case '[':
if len(token) < 2 || token[len(token)-1] != ']' {
return "", end, false
}
token = token[1 : len(token)-1]
}
token = strings.ToLower(strings.TrimSpace(token))
return token, end, token != ""
}
func isReadOnlyPragmaWithArgument(name string) bool {
switch name {
case "foreign_key_check", "foreign_key_list", "index_info", "index_xinfo", "index_list",
"integrity_check", "quick_check", "table_info", "table_xinfo":
return true
default:
return false
}
}
func isReadOnlyPragmaWithoutArgument(name string) bool {
switch name {
case "analysis_limit", "application_id", "auto_vacuum", "automatic_index", "busy_timeout",
"cache_size", "cache_spill", "case_sensitive_like", "cell_size_check", "checkpoint_fullfsync",
"collation_list", "compile_options", "data_version", "database_list", "defer_foreign_keys",
"encoding", "foreign_key_check", "foreign_key_list", "foreign_keys", "freelist_count",
"full_column_names", "fullfsync", "function_list", "hard_heap_limit", "ignore_check_constraints",
"index_info", "index_list", "index_xinfo", "integrity_check", "journal_mode", "journal_size_limit",
"legacy_alter_table", "legacy_file_format", "locking_mode", "max_page_count", "mmap_size",
"module_list", "page_count", "page_size", "pragma_list", "query_only", "quick_check",
"read_uncommitted", "recursive_triggers", "reverse_unordered_selects", "schema_version", "secure_delete",
"short_column_names", "soft_heap_limit", "stats", "synchronous", "table_info", "table_list",
"table_xinfo", "temp_store", "threads", "trusted_schema", "user_version", "wal_autocheckpoint",
"writable_schema":
return true
default:
// Unknown/action pragmas are conservative writes. This covers
// no-argument operations such as optimize, incremental_vacuum and
// wal_checkpoint without depending on a perpetually complete list.
return false
}
}
func isBatchableWriteSQLStatement(dbType string, query string) bool {
if isReadOnlySQLQuery(dbType, query) {
return false

View File

@@ -111,6 +111,61 @@ func TestIsReadOnlySQLQuery_TreatsMongoDeleteAsWrite(t *testing.T) {
}
}
func TestIsReadOnlySQLQuery_TreatsMongoAggregateOutputStagesAsWrites(t *testing.T) {
for _, query := range []string{
`{"aggregate":"users","pipeline":[{"$match":{"active":true}},{"$out":"active_users"}],"cursor":{}}`,
`{"aggregate":"users","pipeline":[{"$merge":{"into":"active_users"}}],"cursor":{}}`,
} {
if isReadOnlySQLQuery("mongodb", query) {
t.Fatalf("Mongo aggregate write stage was classified read-only: %s", query)
}
}
if !isReadOnlySQLQuery("mongodb", `{"aggregate":"users","pipeline":[{"$match":{"active":true}}],"cursor":{}}`) {
t.Fatal("read-only Mongo aggregate was classified as write")
}
}
func TestIsReadOnlySQLQuery_TreatsExecutingExplainWritesAsWrites(t *testing.T) {
for _, query := range []string{
"EXPLAIN ANALYZE UPDATE users SET active = false",
"EXPLAIN (ANALYZE true, BUFFERS true) DELETE FROM users",
"EXPLAIN ANALYSE WITH removed AS (DELETE FROM users RETURNING id) SELECT * FROM removed",
} {
if isReadOnlySQLQuery("postgres", query) {
t.Fatalf("executing EXPLAIN write was classified read-only: %s", query)
}
}
if !isReadOnlySQLQuery("postgres", "EXPLAIN UPDATE users SET active = false") {
t.Fatal("non-executing EXPLAIN was classified as write")
}
if !isReadOnlySQLQuery("postgres", "EXPLAIN ANALYZE SELECT * FROM users") {
t.Fatal("EXPLAIN ANALYZE SELECT was classified as write")
}
}
func TestIsReadOnlySQLQuery_TreatsMutablePragmasAsWrites(t *testing.T) {
for _, query := range []string{
"PRAGMA user_version = 7",
"PRAGMA main.application_id(42)",
`PRAGMA "main".user_version = 123`,
"PRAGMA [main].user_version = 124",
"PRAGMA `main`.user_version = 125",
"PRAGMA optimize",
"PRAGMA incremental_vacuum",
"PRAGMA wal_checkpoint",
} {
if isReadOnlySQLQuery("sqlite", query) {
t.Fatalf("mutable PRAGMA was classified read-only: %s", query)
}
}
if !isReadOnlySQLQuery("sqlite", "PRAGMA table_info('users')") {
t.Fatal("metadata PRAGMA was classified as write")
}
if !isReadOnlySQLQuery("sqlite", "PRAGMA database_list") {
t.Fatal("read-only no-argument PRAGMA was classified as write")
}
}
func TestIsReadOnlySQLQuery_TreatsMilvusJSONQueriesAsReadOnly(t *testing.T) {
for _, query := range []string{
`{"list_collections":true}`,