mirror of
https://github.com/Syngnat/GoNavi.git
synced 2026-08-22 08:53:46 +08:00
🐛 fix(sqlserver): 修复托管事务下 UPDATE 误报执行失败
- 统一处理 SQL Server Exec 路径的 RowsAffected 返回 - 兼容 BEGIN/COMMIT/ROLLBACK/SAVE 等事务控制语句无影响行数场景 - 补充 SQL Server 事务控制语句与 DML 的回归测试
This commit is contained in:
@@ -293,7 +293,7 @@ func (s *SqlServerDB) ExecContext(ctx context.Context, query string) (int64, err
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
return sqlServerRowsAffected(query, res)
|
||||
}
|
||||
|
||||
func (s *SqlServerDB) ExecBatchContext(ctx context.Context, query string) (int64, error) {
|
||||
@@ -304,7 +304,7 @@ func (s *SqlServerDB) ExecBatchContext(ctx context.Context, query string) (int64
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
return sqlServerRowsAffected(query, res)
|
||||
}
|
||||
|
||||
func (s *SqlServerDB) OpenSessionExecer(ctx context.Context) (StatementExecer, error) {
|
||||
@@ -326,7 +326,7 @@ func (s *SqlServerDB) Exec(query string) (int64, error) {
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
return sqlServerRowsAffected(query, res)
|
||||
}
|
||||
|
||||
func (e *sqlServerSessionExecer) Exec(query string) (int64, error) {
|
||||
@@ -341,7 +341,38 @@ func (e *sqlServerSessionExecer) ExecContext(ctx context.Context, query string)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.RowsAffected()
|
||||
return sqlServerRowsAffected(query, res)
|
||||
}
|
||||
|
||||
func sqlServerRowsAffected(query string, res sql.Result) (int64, error) {
|
||||
if res == nil {
|
||||
return 0, nil
|
||||
}
|
||||
affected, err := res.RowsAffected()
|
||||
if err == nil {
|
||||
return affected, nil
|
||||
}
|
||||
if sqlServerAllowsUnknownRowsAffected(query) {
|
||||
return 0, nil
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
|
||||
func sqlServerAllowsUnknownRowsAffected(query string) bool {
|
||||
trimmed := strings.TrimSpace(query)
|
||||
if trimmed == "" {
|
||||
return false
|
||||
}
|
||||
fields := strings.Fields(trimmed)
|
||||
if len(fields) == 0 {
|
||||
return false
|
||||
}
|
||||
switch strings.ToLower(fields[0]) {
|
||||
case "begin", "commit", "rollback", "save":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (e *sqlServerSessionExecer) Query(query string) ([]map[string]interface{}, []string, error) {
|
||||
|
||||
67
internal/db/sqlserver_impl_test.go
Normal file
67
internal/db/sqlserver_impl_test.go
Normal file
@@ -0,0 +1,67 @@
|
||||
//go:build gonavi_full_drivers || gonavi_sqlserver_driver
|
||||
|
||||
package db
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type fakeSQLServerExecResult struct {
|
||||
affected int64
|
||||
rowErr error
|
||||
}
|
||||
|
||||
func (r fakeSQLServerExecResult) LastInsertId() (int64, error) {
|
||||
return 0, errors.New("not implemented")
|
||||
}
|
||||
|
||||
func (r fakeSQLServerExecResult) RowsAffected() (int64, error) {
|
||||
if r.rowErr != nil {
|
||||
return 0, r.rowErr
|
||||
}
|
||||
return r.affected, nil
|
||||
}
|
||||
|
||||
func TestSQLServerRowsAffectedIgnoresTransactionControlErrors(t *testing.T) {
|
||||
rowErr := errors.New("不支持的方法")
|
||||
for _, query := range []string{
|
||||
"BEGIN TRANSACTION",
|
||||
"COMMIT TRANSACTION",
|
||||
"ROLLBACK TRANSACTION",
|
||||
"SAVE TRANSACTION before_update",
|
||||
"BEGIN TRY\nSELECT 1\nEND TRY",
|
||||
} {
|
||||
affected, err := sqlServerRowsAffected(query, fakeSQLServerExecResult{rowErr: rowErr})
|
||||
if err != nil {
|
||||
t.Fatalf("sqlServerRowsAffected(%q) returned unexpected error: %v", query, err)
|
||||
}
|
||||
if affected != 0 {
|
||||
t.Fatalf("sqlServerRowsAffected(%q) = %d, want 0", query, affected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLServerRowsAffectedPreservesDMLCount(t *testing.T) {
|
||||
affected, err := sqlServerRowsAffected(
|
||||
"UPDATE dbo.users SET name = 'neo' WHERE id = 1",
|
||||
fakeSQLServerExecResult{affected: 3},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("sqlServerRowsAffected returned unexpected error: %v", err)
|
||||
}
|
||||
if affected != 3 {
|
||||
t.Fatalf("sqlServerRowsAffected = %d, want 3", affected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLServerRowsAffectedDoesNotHideDMLRowsAffectedErrors(t *testing.T) {
|
||||
rowErr := errors.New("rows affected unsupported")
|
||||
_, err := sqlServerRowsAffected(
|
||||
"UPDATE dbo.users SET name = 'neo' WHERE id = 1",
|
||||
fakeSQLServerExecResult{rowErr: rowErr},
|
||||
)
|
||||
if !errors.Is(err, rowErr) {
|
||||
t.Fatalf("expected rows affected error to propagate for DML, got %v", err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user