🐛 fix(sqlserver): 修复托管事务下 UPDATE 误报执行失败

- 统一处理 SQL Server Exec 路径的 RowsAffected 返回
- 兼容 BEGIN/COMMIT/ROLLBACK/SAVE 等事务控制语句无影响行数场景
- 补充 SQL Server 事务控制语句与 DML 的回归测试
This commit is contained in:
Syngnat
2026-06-14 18:03:06 +08:00
parent 5f892d29c8
commit a750266e1c
2 changed files with 102 additions and 4 deletions

View File

@@ -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) {

View 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)
}
}