🐛 fix(sqlserver): 修复存储过程执行消息丢失

- 按消息协议仅在 MsgNext 时扫描结果集,避免空结果边界提前耗尽消息
- 保留连续 DONE 边界后的 PRINT 输出与纯消息结果
- 更新 SQL Server driver-agent revision 并补充回归测试
This commit is contained in:
Syngnat
2026-07-20 13:40:53 +08:00
parent 258f5bdd64
commit 89e35b087a
3 changed files with 139 additions and 113 deletions

View File

@@ -9,7 +9,7 @@ func init() {
"diros": "src-7d4fe439271d0c56",
"starrocks": "src-ce9ee22641a32f46",
"sphinx": "src-08f5ae54efb3d9df",
"sqlserver": "src-bd1323910fc119ac",
"sqlserver": "src-6c0e98d6d8ba439d",
"sqlite": "src-96dfa25b3042b2d5",
"duckdb": "src-8804eb2cdbc89433",
"dameng": "src-016e77082aea6718",

View File

@@ -40,33 +40,10 @@ func scanSQLServerRowsWithMessages(ctx context.Context, rows *sql.Rows, retmsg *
}
var (
resultSets []connection.ResultSetData
messages []string
allMessages []string
currentResultScanned bool
currentResultStart int
resultSets []connection.ResultSetData
messages []string
allMessages []string
)
// go-mssqldb can emit a result-set boundary without MsgNext. Recover the
// unread rows before advancing and keep them before any affectedRows status.
scanCurrentResultIfNeeded := func() error {
if currentResultScanned {
return nil
}
result, err := scanSQLServerFallbackResultSet(rows)
if err != nil {
return err
}
currentResultScanned = true
if len(result.Columns) == 0 && len(result.Rows) == 0 {
return nil
}
result.Messages = append([]string(nil), messages...)
resultSets = append(resultSets, connection.ResultSetData{})
copy(resultSets[currentResultStart+1:], resultSets[currentResultStart:])
resultSets[currentResultStart] = result
messages = nil
return nil
}
active := true
for active {
raw := retmsg.Message(ctx)
@@ -93,7 +70,6 @@ func scanSQLServerRowsWithMessages(ctx context.Context, rows *sql.Rows, retmsg *
Columns: cols,
Messages: append([]string(nil), messages...),
})
currentResultScanned = true
messages = nil
case sqlexp.MsgRowsAffected:
resultSets = append(resultSets, connection.ResultSetData{
@@ -103,24 +79,15 @@ func scanSQLServerRowsWithMessages(ctx context.Context, rows *sql.Rows, retmsg *
})
messages = nil
case sqlexp.MsgNextResultSet:
if err := scanCurrentResultIfNeeded(); err != nil {
return resultSets, messages, err
}
// Only MsgNext proves a row set is ready. Calling Columns at an empty
// boundary drains later DONE and PRINT tokens in go-mssqldb Rowsq.
active = rows.NextResultSet()
if active {
currentResultScanned = false
currentResultStart = len(resultSets)
}
case sqlexp.MsgError:
return resultSets, messages, msg.Error
default:
active = false
}
}
if err := scanCurrentResultIfNeeded(); err != nil {
return resultSets, messages, err
}
if len(messages) > 0 {
resultSets = append(resultSets, connection.ResultSetData{
Rows: []map[string]interface{}{},
@@ -144,23 +111,6 @@ func emptySQLServerRowsResultSet() connection.ResultSetData {
}
}
func scanSQLServerFallbackResultSet(rows *sql.Rows) (connection.ResultSetData, error) {
data, columns, err := scanRows(rows)
if err != nil {
return emptySQLServerRowsResultSet(), err
}
if data == nil {
data = []map[string]interface{}{}
}
if columns == nil {
columns = []string{}
}
return connection.ResultSetData{
Rows: data,
Columns: columns,
}, nil
}
// quoteBracket escapes ] in identifiers for safe use in SQL Server [bracket] notation
func quoteBracket(name string) string {
return strings.ReplaceAll(name, "]", "]]")

View File

@@ -5,10 +5,13 @@ package db
import (
"context"
"database/sql"
"database/sql/driver"
"errors"
"io"
"os"
"reflect"
"strings"
"sync"
"testing"
"GoNavi-Wails/shared/i18n"
@@ -19,6 +22,97 @@ import (
var rawSQLServerTableNameRequiredText = string([]rune{0x8868, 0x540d, 0x4e0d, 0x80fd, 0x4e3a, 0x7a7a})
const sqlServerPrintOnlyDriverName = "gonavi-sqlserver-print-only"
var registerSQLServerPrintOnlyDriver sync.Once
type sqlServerPrintOnlyDriver struct{}
type sqlServerPrintOnlyConn struct {
retmsg *sqlexp.ReturnMessage
}
type sqlServerPrintOnlyRows struct {
remainingBoundaries int
drained bool
}
type sqlServerPrintOnlyNotice string
func (sqlServerPrintOnlyDriver) Open(string) (driver.Conn, error) {
return &sqlServerPrintOnlyConn{}, nil
}
func (c *sqlServerPrintOnlyConn) Prepare(string) (driver.Stmt, error) {
return nil, errors.New("not implemented")
}
func (c *sqlServerPrintOnlyConn) Close() error {
return nil
}
func (c *sqlServerPrintOnlyConn) Begin() (driver.Tx, error) {
return nil, errors.New("not implemented")
}
func (c *sqlServerPrintOnlyConn) CheckNamedValue(value *driver.NamedValue) error {
retmsg, ok := value.Value.(*sqlexp.ReturnMessage)
if !ok {
return nil
}
sqlexp.ReturnMessageInit(retmsg)
c.retmsg = retmsg
return driver.ErrRemoveArgument
}
func (c *sqlServerPrintOnlyConn) QueryContext(ctx context.Context, _ string, _ []driver.NamedValue) (driver.Rows, error) {
for _, message := range []sqlexp.RawMessage{
sqlexp.MsgNextResultSet{},
sqlexp.MsgNotice{Message: sqlServerPrintOnlyNotice("INSERT c_user(userid) values('168')")},
sqlexp.MsgNextResultSet{},
sqlexp.MsgNotice{Message: sqlServerPrintOnlyNotice("INSERT c_user(userid) values('169')")},
sqlexp.MsgNextResultSet{},
sqlexp.MsgNextResultSet{},
} {
if err := sqlexp.ReturnMessageEnqueue(ctx, c.retmsg, message); err != nil {
return nil, err
}
}
return &sqlServerPrintOnlyRows{remainingBoundaries: 3}, nil
}
func (r *sqlServerPrintOnlyRows) Columns() []string {
// go-mssqldb Rowsq.Columns drains all empty DONE boundaries when no column
// metadata exists. Calling it before the message loop ends loses later PRINTs.
r.drained = true
r.remainingBoundaries = 0
return []string{}
}
func (*sqlServerPrintOnlyRows) Close() error {
return nil
}
func (*sqlServerPrintOnlyRows) Next([]driver.Value) error {
return io.EOF
}
func (*sqlServerPrintOnlyRows) HasNextResultSet() bool {
return true
}
func (r *sqlServerPrintOnlyRows) NextResultSet() error {
if r.drained || r.remainingBoundaries == 0 {
return io.EOF
}
r.remainingBoundaries--
return nil
}
func (m sqlServerPrintOnlyNotice) String() string {
return string(m)
}
type fakeSQLServerExecResult struct {
affected int64
rowErr error
@@ -109,61 +203,7 @@ func TestSQLServerSessionExecerDiscardEvictsPhysicalConnection(t *testing.T) {
}
}
func TestScanSQLServerFallbackResultSetPreservesRowsWhenMessageLoopYieldsNoResult(t *testing.T) {
dbConn, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() {
_ = dbConn.Close()
})
rows, err := dbConn.Query("SELECT 'config:roomType:add' AS menuName")
if err != nil {
t.Fatalf("query rows: %v", err)
}
defer rows.Close()
resultSet, err := scanSQLServerFallbackResultSet(rows)
if err != nil {
t.Fatalf("scanSQLServerFallbackResultSet returned error: %v", err)
}
if !reflect.DeepEqual(resultSet.Columns, []string{"menuName"}) {
t.Fatalf("expected SELECT columns to be preserved, got %#v", resultSet.Columns)
}
if len(resultSet.Rows) != 1 || resultSet.Rows[0]["menuName"] != "config:roomType:add" {
t.Fatalf("expected SELECT rows to be preserved, got %#v", resultSet.Rows)
}
}
func TestScanSQLServerFallbackResultSetPreservesColumnsWhenResultHasNoRows(t *testing.T) {
dbConn, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() {
_ = dbConn.Close()
})
rows, err := dbConn.Query("SELECT 1 AS menuName WHERE 1 = 0")
if err != nil {
t.Fatalf("query empty rows: %v", err)
}
defer rows.Close()
resultSet, err := scanSQLServerFallbackResultSet(rows)
if err != nil {
t.Fatalf("scanSQLServerFallbackResultSet returned error: %v", err)
}
if len(resultSet.Rows) != 0 {
t.Fatalf("expected empty rows, got %#v", resultSet.Rows)
}
if !reflect.DeepEqual(resultSet.Columns, []string{"menuName"}) {
t.Fatalf("expected empty SELECT columns to be preserved, got %#v", resultSet.Columns)
}
}
func TestScanSQLServerRowsWithMessagesRecoversRowsWhenMessageLoopOmitsMsgNext(t *testing.T) {
func TestScanSQLServerRowsWithMessagesPreservesRowsFromMsgNext(t *testing.T) {
dbConn, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatalf("open sqlite: %v", err)
@@ -182,6 +222,7 @@ func TestScanSQLServerRowsWithMessagesRecoversRowsWhenMessageLoopOmitsMsgNext(t
retmsg := &sqlexp.ReturnMessage{}
sqlexp.ReturnMessageInit(retmsg)
for _, message := range []sqlexp.RawMessage{
sqlexp.MsgNext{},
sqlexp.MsgRowsAffected{Count: 1},
sqlexp.MsgNextResultSet{},
} {
@@ -198,12 +239,12 @@ func TestScanSQLServerRowsWithMessagesRecoversRowsWhenMessageLoopOmitsMsgNext(t
t.Fatalf("expected no SQL Server notices, got %#v", messages)
}
if len(resultSets) != 2 {
t.Fatalf("expected recovered SELECT rows plus affected-row status, got %#v", resultSets)
t.Fatalf("expected SELECT rows plus affected-row status, got %#v", resultSets)
}
if !reflect.DeepEqual(resultSets[0].Columns, []string{"menuName"}) ||
len(resultSets[0].Rows) != 1 ||
resultSets[0].Rows[0]["menuName"] != "config:roomType:add" {
t.Fatalf("expected recovered SELECT rows first, got %#v", resultSets)
t.Fatalf("expected SELECT rows first, got %#v", resultSets)
}
if !reflect.DeepEqual(resultSets[1].Columns, []string{"affectedRows"}) ||
len(resultSets[1].Rows) != 1 ||
@@ -212,6 +253,41 @@ func TestScanSQLServerRowsWithMessagesRecoversRowsWhenMessageLoopOmitsMsgNext(t
}
}
func TestSQLServerQueryMultiWithMessagesPreservesPrintsAfterEmptyResultBoundaries(t *testing.T) {
registerSQLServerPrintOnlyDriver.Do(func() {
sql.Register(sqlServerPrintOnlyDriverName, sqlServerPrintOnlyDriver{})
})
dbConn, err := sql.Open(sqlServerPrintOnlyDriverName, "")
if err != nil {
t.Fatalf("open print-only SQL Server driver: %v", err)
}
t.Cleanup(func() {
_ = dbConn.Close()
})
dbInst := &SqlServerDB{conn: dbConn}
resultSets, messages, err := dbInst.QueryMultiWithMessages("p_get_select 'c_user','1=1',1")
if err != nil {
t.Fatalf("QueryMultiWithMessages returned error: %v", err)
}
wantMessages := []string{
"INSERT c_user(userid) values('168')",
"INSERT c_user(userid) values('169')",
}
if !reflect.DeepEqual(messages, wantMessages) {
t.Fatalf("expected all PRINT messages, got %#v", messages)
}
if len(resultSets) != 1 {
t.Fatalf("expected one message-only result set, got %#v", resultSets)
}
if len(resultSets[0].Rows) != 0 || len(resultSets[0].Columns) != 0 {
t.Fatalf("expected message-only result set without tabular data, got %#v", resultSets[0])
}
if !reflect.DeepEqual(resultSets[0].Messages, wantMessages) {
t.Fatalf("expected result set to preserve all PRINT messages, got %#v", resultSets[0].Messages)
}
}
func TestSQLServerMetadataErrorsUseCurrentLanguage(t *testing.T) {
SetBackendLanguage(i18n.LanguageEnUS)
t.Cleanup(func() {