diff --git a/pkg/sqlcmd/format.go b/pkg/sqlcmd/format.go index 3ffea3a8..f65c55ae 100644 --- a/pkg/sqlcmd/format.go +++ b/pkg/sqlcmd/format.go @@ -249,11 +249,7 @@ func (f *sqlCmdFormatterType) AddError(err error) { switch e := (err).(type) { case mssql.Error: if print = f.vars.ErrorLevel() <= 0 || e.Class >= uint8(f.vars.ErrorLevel()); print { - if len(e.ProcName) > 0 { - b.WriteString(localizer.Sprintf("Msg %#v, Level %d, State %d, Server %s, Procedure %s, Line %#v%s", e.Number, e.Class, e.State, e.ServerName, e.ProcName, e.LineNo, SqlcmdEol)) - } else { - b.WriteString(localizer.Sprintf("Msg %#v, Level %d, State %d, Server %s, Line %#v%s", e.Number, e.Class, e.State, e.ServerName, e.LineNo, SqlcmdEol)) - } + b.WriteString(errorHeader(e)) if !f.rawErrors { msg = strings.TrimPrefix(msg, "mssql: ") } @@ -266,6 +262,15 @@ func (f *sqlCmdFormatterType) AddError(err error) { } } +// errorHeader returns the "Msg N, Level N, State N, ..." line that precedes a +// server error message. +func errorHeader(e mssql.Error) string { + if len(e.ProcName) > 0 { + return localizer.Sprintf("Msg %#v, Level %d, State %d, Server %s, Procedure %s, Line %#v%s", e.Number, e.Class, e.State, e.ServerName, e.ProcName, e.LineNo, SqlcmdEol) + } + return localizer.Sprintf("Msg %#v, Level %d, State %d, Server %s, Line %#v%s", e.Number, e.Class, e.State, e.ServerName, e.LineNo, SqlcmdEol) +} + // XmlMode enables or disables XML mode func (f *sqlCmdFormatterType) XmlMode(enable bool) { f.xml = enable diff --git a/pkg/sqlcmd/sqlcmd.go b/pkg/sqlcmd/sqlcmd.go index 93637a02..ec89d45b 100644 --- a/pkg/sqlcmd/sqlcmd.go +++ b/pkg/sqlcmd/sqlcmd.go @@ -81,6 +81,7 @@ type Sqlcmd struct { // Cmd provides the implementation of commands like :list and GO Cmd Commands // PrintError allows the host to redirect errors away from the default output. Returns false if the error is not redirected by the host. + // For server errors msg includes the "Msg N, Level N, State N, ..." header line, as the default output would show it. PrintError func(msg string, severity uint8) bool // UnicodeOutputFile is true when UTF16 file output is needed UnicodeOutputFile bool @@ -485,7 +486,7 @@ func (s *Sqlcmd) runQuery(query string) (int, error) { case sqlexp.MsgError: switch e := m.Error.(type) { case mssql.Error: - if !s.PrintError(e.Message, e.Class) { + if !s.PrintError(errorHeader(e)+e.Message, e.Class) { s.Format.AddError(m.Error) } } diff --git a/pkg/sqlcmd/sqlcmd_test.go b/pkg/sqlcmd/sqlcmd_test.go index 2c325fed..87069599 100644 --- a/pkg/sqlcmd/sqlcmd_test.go +++ b/pkg/sqlcmd/sqlcmd_test.go @@ -505,6 +505,26 @@ func TestSqlCmdDefersToPrintError(t *testing.T) { } } +// A host that redirects errors (sqlcmd -r) must receive the same +// "Msg N, Level N, State N, ..." header the default output shows. +func TestPrintErrorIncludesErrorHeader(t *testing.T) { + s, buf := setupSqlCmdWithMemoryOutput(t) + defer buf.Close() + redirected := "" + s.PrintError = func(msg string, severity uint8) bool { + if severity > 10 { + redirected += msg + return true + } + return false + } + err := runSqlCmd(t, s, []string{"RAISERROR (N'Testing!' , 11, 1)", "GO"}) + if assert.NoError(t, err, "runSqlCmd failed") { + assert.Regexp(t, "^Msg 50000, Level 11, State 1, Server .+, Line 1"+SqlcmdEol+"Testing!$", redirected) + assert.Empty(t, buf.buf.String(), "redirected errors should not reach the default output") + } +} + func TestSqlCmdMaintainsConnectionBetweenBatches(t *testing.T) { s, buf := setupSqlCmdWithMemoryOutput(t) defer buf.Close()