Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 10 additions & 5 deletions pkg/sqlcmd/format.go
Original file line number Diff line number Diff line change
Expand Up @@ -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: ")
}
Expand All @@ -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
Expand Down
3 changes: 2 additions & 1 deletion pkg/sqlcmd/sqlcmd.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
}
}
Expand Down
20 changes: 20 additions & 0 deletions pkg/sqlcmd/sqlcmd_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down