mirror of
https://github.com/Syngnat/GoNavi.git
synced 2026-08-10 16:53:35 +08:00
🐛 fix(sqlserver): 修复存储过程执行消息丢失
- 按消息协议仅在 MsgNext 时扫描结果集,避免空结果边界提前耗尽消息 - 保留连续 DONE 边界后的 PRINT 输出与纯消息结果 - 更新 SQL Server driver-agent revision 并补充回归测试
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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, "]", "]]")
|
||||
|
||||
@@ -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() {
|
||||
|
||||
Reference in New Issue
Block a user