diff --git a/internal/db/driver_agent_revisions_gen.go b/internal/db/driver_agent_revisions_gen.go index e561fcd7..035bd38b 100644 --- a/internal/db/driver_agent_revisions_gen.go +++ b/internal/db/driver_agent_revisions_gen.go @@ -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", diff --git a/internal/db/sqlserver_impl.go b/internal/db/sqlserver_impl.go index b79dd144..2ec4785b 100644 --- a/internal/db/sqlserver_impl.go +++ b/internal/db/sqlserver_impl.go @@ -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, "]", "]]") diff --git a/internal/db/sqlserver_impl_test.go b/internal/db/sqlserver_impl_test.go index 5d301afc..15e2dba7 100644 --- a/internal/db/sqlserver_impl_test.go +++ b/internal/db/sqlserver_impl_test.go @@ -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() {