diff --git a/plugins/services/credential_tester.go b/plugins/services/credential_tester.go index b557420..aa302b6 100644 --- a/plugins/services/credential_tester.go +++ b/plugins/services/credential_tester.go @@ -2,6 +2,7 @@ package services import ( "context" + "database/sql" "fmt" "io" "sync" @@ -369,3 +370,17 @@ func matchIgnoreCase(a, b string) bool { } return true } + +// ============================================================================= +// 通用数据库连接包装 +// ============================================================================= + +// SQLDBWrapper 包装 sql.DB 以实现 io.Closer +// 用于 MySQL、PostgreSQL、MSSQL、Oracle 等数据库插件的连接返回 +type SQLDBWrapper struct { + *sql.DB +} + +func (w *SQLDBWrapper) Close() error { + return w.DB.Close() +} diff --git a/plugins/services/mssql.go b/plugins/services/mssql.go index a1d6b48..0cb533c 100644 --- a/plugins/services/mssql.go +++ b/plugins/services/mssql.go @@ -98,21 +98,12 @@ func (p *MSSQLPlugin) doMSSQLAuth(ctx context.Context, info *common.HostInfo, cr return &AuthResult{ Success: true, - Conn: &mssqlDBWrapper{db}, + Conn: &SQLDBWrapper{db}, ErrorType: ErrorTypeUnknown, Error: nil, } } -// mssqlDBWrapper 包装 sql.DB 以实现 io.Closer -type mssqlDBWrapper struct { - *sql.DB -} - -func (w *mssqlDBWrapper) Close() error { - return w.DB.Close() -} - // classifyMSSQLErrorType MSSQL错误分类 func classifyMSSQLErrorType(err error) ErrorType { if err == nil { diff --git a/plugins/services/mysql.go b/plugins/services/mysql.go index 4db1463..543f3a2 100644 --- a/plugins/services/mysql.go +++ b/plugins/services/mysql.go @@ -104,24 +104,14 @@ func (p *MySQLPlugin) doMySQLAuth(ctx context.Context, info *common.HostInfo, cr state.IncrementTCPSuccessPacketCount() - // MySQL 使用 sql.DB,包装为 io.Closer return &AuthResult{ Success: true, - Conn: &sqlDBWrapper{db}, + Conn: &SQLDBWrapper{db}, ErrorType: ErrorTypeUnknown, Error: nil, } } -// sqlDBWrapper 包装 sql.DB 以实现 io.Closer -type sqlDBWrapper struct { - *sql.DB -} - -func (w *sqlDBWrapper) Close() error { - return w.DB.Close() -} - // classifyMySQLErrorType MySQL错误分类 func classifyMySQLErrorType(err error) ErrorType { if err == nil { diff --git a/plugins/services/oracle.go b/plugins/services/oracle.go index 87c2703..4125ea9 100644 --- a/plugins/services/oracle.go +++ b/plugins/services/oracle.go @@ -106,7 +106,7 @@ func (p *OraclePlugin) doOracleAuth(ctx context.Context, info *common.HostInfo, return &AuthResult{ Success: true, - Conn: &oracleDBWrapper{db}, + Conn: &SQLDBWrapper{db}, ErrorType: ErrorTypeUnknown, Error: nil, } @@ -120,15 +120,6 @@ func (p *OraclePlugin) doOracleAuth(ctx context.Context, info *common.HostInfo, } } -// oracleDBWrapper 包装 sql.DB 以实现 io.Closer -type oracleDBWrapper struct { - *sql.DB -} - -func (w *oracleDBWrapper) Close() error { - return w.DB.Close() -} - // classifyOracleErrorType Oracle错误分类 func classifyOracleErrorType(err error) ErrorType { if err == nil { diff --git a/plugins/services/postgresql.go b/plugins/services/postgresql.go index 1777a0d..e5ecc69 100644 --- a/plugins/services/postgresql.go +++ b/plugins/services/postgresql.go @@ -104,21 +104,12 @@ func (p *PostgreSQLPlugin) doPostgreSQLAuth(ctx context.Context, info *common.Ho return &AuthResult{ Success: true, - Conn: &pgDBWrapper{db}, + Conn: &SQLDBWrapper{db}, ErrorType: ErrorTypeUnknown, Error: nil, } } -// pgDBWrapper 包装 sql.DB 以实现 io.Closer -type pgDBWrapper struct { - *sql.DB -} - -func (w *pgDBWrapper) Close() error { - return w.DB.Close() -} - // classifyPostgreSQLErrorType PostgreSQL错误分类 func classifyPostgreSQLErrorType(err error) ErrorType { if err == nil {