diff --git a/internal/app/renderer/formatter/table.go b/internal/app/renderer/formatter/table.go index 5c15aaa..255825e 100644 --- a/internal/app/renderer/formatter/table.go +++ b/internal/app/renderer/formatter/table.go @@ -1,6 +1,7 @@ package formatter import ( + "encoding/csv" "io" "github.com/balajz/pgxcli/internal/config" @@ -12,8 +13,9 @@ import ( ) type TableFormatter struct { - rows int - table *tablewriter.Table + rows int + table *tablewriter.Table + streamWriter *csv.Writer tableConfig *config.TableConfig } @@ -34,6 +36,18 @@ func (p *TableFormatter) Column(_ io.Writer, cols []string) error { } func (p *TableFormatter) Iter(_, ew io.Writer, row []string) error { + if p.streamWriter != nil { + if err := p.streamWriter.Write(row); err != nil { + return perrors.Wrap(err, perrors.WithMessage("failed to stream row")) + } + p.streamWriter.Flush() + if err := p.streamWriter.Error(); err != nil { + return perrors.Wrap(err, perrors.WithMessage("failed to flush streamed row")) + } + p.rows++ + return nil + } + if p.table == nil { return nil } @@ -46,6 +60,29 @@ func (p *TableFormatter) Iter(_, ew io.Writer, row []string) error { return nil } +// StartStreaming abandons the buffered table once it grows too large and +// emits the retained rows as tab-separated records. This keeps the query path +// bounded while preserving all values, including values containing tabs or +// newlines through csv quoting. +func (p *TableFormatter) StartStreaming(w io.Writer, cols []string, bufferedRows [][]string) error { + p.table = nil + p.streamWriter = csv.NewWriter(w) + p.streamWriter.Comma = '\t' + if err := p.streamWriter.Write(cols); err != nil { + return perrors.Wrap(err, perrors.WithMessage("failed to stream header")) + } + for _, row := range bufferedRows { + if err := p.streamWriter.Write(row); err != nil { + return perrors.Wrap(err, perrors.WithMessage("failed to stream buffered row")) + } + } + p.streamWriter.Flush() + if err := p.streamWriter.Error(); err != nil { + return perrors.Wrap(err, perrors.WithMessage("failed to flush streamed table")) + } + return nil +} + func (p *TableFormatter) Caption(w io.Writer, caption string) error { if p.table == nil { return nil @@ -62,6 +99,13 @@ func (p *TableFormatter) Caption(w io.Writer, caption string) error { } func (p *TableFormatter) Render(_ io.Writer, _ int) error { + if p.streamWriter != nil { + p.streamWriter.Flush() + if err := p.streamWriter.Error(); err != nil { + return perrors.Wrap(err, perrors.WithMessage("failed to flush streamed table")) + } + return nil + } if err := p.table.Render(); err != nil { return perrors.Wrap(err, perrors.WithMessage("failed to render table")) } @@ -70,5 +114,6 @@ func (p *TableFormatter) Render(_ io.Writer, _ int) error { func (p *TableFormatter) Done(_ io.Writer) error { p.table = nil + p.streamWriter = nil return nil } diff --git a/internal/app/renderer/renderer.go b/internal/app/renderer/renderer.go index 5b37da2..67d6a4a 100644 --- a/internal/app/renderer/renderer.go +++ b/internal/app/renderer/renderer.go @@ -14,6 +14,15 @@ type Formatter interface { Done(w io.Writer) error } +// StreamingFormatter can switch from table buffering to bounded row output +// once a result exceeds the in-memory formatting limit. +type StreamingFormatter interface { + Formatter + StartStreaming(w io.Writer, cols []string, bufferedRows [][]string) error +} + +const bufferedRowLimit = 500 + func TableRender(cols []string, rowrowIter RowStrIter, caption string, w, Ew io.Writer, c *config.Config) error { tf := formatter.NewTableFormatter(w, &c.Table) return Render(w, Ew, tf, cols, rowrowIter) @@ -24,6 +33,9 @@ func Render(w, Ew io.Writer, formatter Formatter, cols []string, row RowStrIter) return err } + streamingFormatter, canStream := formatter.(StreamingFormatter) + bufferedRows := make([][]string, 0, bufferedRowLimit) + streaming := false nRows := 0 for { r, err := row.Next() @@ -34,9 +46,20 @@ func Render(w, Ew io.Writer, formatter Formatter, cols []string, row RowStrIter) return err } + if canStream && !streaming && len(bufferedRows) >= bufferedRowLimit { + if err := streamingFormatter.StartStreaming(w, cols, bufferedRows); err != nil { + return err + } + bufferedRows = nil + streaming = true + } + if err := formatter.Iter(w, Ew, r); err != nil { return err } + if canStream && !streaming { + bufferedRows = append(bufferedRows, append([]string(nil), r...)) + } nRows++ } diff --git a/internal/app/renderer/streaming_test.go b/internal/app/renderer/streaming_test.go new file mode 100644 index 0000000..aa77adf --- /dev/null +++ b/internal/app/renderer/streaming_test.go @@ -0,0 +1,38 @@ +package renderer + +import ( + "fmt" + "strings" + "testing" + + "github.com/balajz/pgxcli/internal/app/renderer/formatter" + "github.com/balajz/pgxcli/internal/config" +) + +func TestTableRenderSwitchesToBoundedStreamingOutput(t *testing.T) { + rows := make([][]string, bufferedRowLimit+1) + for i := range rows { + rows[i] = []string{fmt.Sprintf("row-%d", i), "value"} + } + + var out strings.Builder + err := Render(&out, &out, formatter.NewTableFormatter(&out, &config.TableConfig{}), []string{"name", "value"}, NewRowSliceIter(rows)) + if err != nil { + t.Fatal(err) + } + + result := out.String() + if !strings.Contains(result, "row-0\tvalue\n") { + t.Fatalf("expected streaming TSV output, got %q", result[len(result)-min(len(result), 200):]) + } + if !strings.Contains(result, fmt.Sprintf("row-%d\tvalue\n", bufferedRowLimit)) { + t.Fatalf("expected final streamed row record in output, got %q", result[len(result)-min(len(result), 200):]) + } +} + +func min(a, b int) int { + if a < b { + return a + } + return b +} diff --git a/internal/app/run_query.go b/internal/app/run_query.go index 1acfd0f..64b923a 100644 --- a/internal/app/run_query.go +++ b/internal/app/run_query.go @@ -3,7 +3,7 @@ package app import ( "context" "fmt" - "strings" + "io" "time" tea "charm.land/bubbletea/v2" @@ -53,41 +53,40 @@ StatementsLoop: } func (p *pgxCLI) handleQueryResult(r database.Rows, execDuration time.Duration) (cmd tea.Cmd, err error) { - var s strings.Builder - cols := renderer.GetColumnStrings(r, true) - if len(cols) > 0 { - rowIter := renderer.NewRowIter(r, true) - if err := renderer.TableRender(cols, rowIter, "", &s, &s, p.config); err != nil { - r.Close() // Ensure closed on error - return nil, err - } - } - - // We must close the rows before reading the tag - if closeErr := r.Close(); closeErr != nil { - return nil, closeErr - } - - tag, err := r.Tag() - if err != nil { - return nil, err - } - tagStr := tag.String() - if tagStr == "" { - tagStr = "OK" - } - - output := s.String() - if len(cols) == 0 { - output = tagStr - } else { - output += tagStr - } + return func() tea.Msg { + streamErr := p.Printer.StreamViaPager(func(w io.Writer) error { + if len(cols) > 0 { + rowIter := renderer.NewRowIter(r, true) + if err := renderer.TableRender(cols, rowIter, "", w, w, p.config); err != nil { + _ = r.Close() + return err + } + } - // Append timing info to the output - timingInfo := fmt.Sprintf("\nTime %.3fs", execDuration.Seconds()) - output += timingInfo + // We must close the rows before reading the tag. + if closeErr := r.Close(); closeErr != nil { + return closeErr + } - return p.printViaPager(output), nil + tag, err := r.Tag() + if err != nil { + return err + } + tagStr := tag.String() + if tagStr == "" { + tagStr = "OK" + } + if _, err := fmt.Fprintln(w, tagStr); err != nil { + return err + } + _, err = fmt.Fprintf(w, "Time %.3fs\n", execDuration.Seconds()) + return err + }) + if streamErr != nil { + p.logger.Error("error streaming query result", "error", streamErr) + return ui.PrintErrCmd(streamErr, ui.DefaultStyles().ErrorOutput) + } + return nil + }, nil } diff --git a/internal/cliio/printer.go b/internal/cliio/printer.go index a2e0f39..662408c 100644 --- a/internal/cliio/printer.go +++ b/internal/cliio/printer.go @@ -3,6 +3,7 @@ package cliio import ( + "bytes" "errors" "fmt" "io" @@ -47,6 +48,7 @@ type Printer interface { PrintError(err error) PrintTime(time time.Duration) PrintViaPager(str string) + StreamViaPager(writeFn func(io.Writer) error) error ShouldUsePager(str string) bool } @@ -206,6 +208,102 @@ func (p *pgxPrinter) PrintViaPager(str string) { } } +// StreamViaPager sends output directly to the configured destination or pager +// without first building the complete result in memory. Auto pager mode uses a +// temporary file only when it must measure the output before choosing a pager. +func (p *pgxPrinter) StreamViaPager(writeFn func(io.Writer) error) error { + if p.pagerMode == pagerModeAlways && p.isTerminal && p.pagerSupported { + if p.tryPipePager(writeFn) || p.tryTempfilePager(writeFn) { + return nil + } + } + + if p.pagerMode != pagerModeAuto || !p.isTerminal || !p.pagerSupported { + return writeFn(p.out) + } + + tmp, err := os.CreateTemp("", "pgxcli-output-*") + if err != nil { + return writeFn(p.out) + } + tmpName := tmp.Name() + defer func() { + _ = os.Remove(tmpName) + }() + + if err := writeFn(tmp); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Close(); err != nil { + return err + } + + file, err := os.Open(tmpName) + if err != nil { + return err + } + defer file.Close() + + info, err := file.Stat() + if err != nil { + return err + } + lines, err := countFileLines(file) + if err != nil { + return err + } + if !p.shouldUsePagerMetrics(info.Size(), lines) { + _, err = io.Copy(p.out, file) + return err + } + + if _, err := file.Seek(0, io.SeekStart); err != nil { + return err + } + cmd := exec.Command(p.pagerPath, p.pagerArgs...) + cmd.Stdin = file + cmd.Stdout = os.Stdout + cmd.Stderr = os.Stderr + return waitIgnoringInterrupt(cmd) +} + +func countFileLines(file *os.File) (int, error) { + if _, err := file.Seek(0, io.SeekStart); err != nil { + return 0, err + } + buf := make([]byte, 32*1024) + lines := 0 + for { + n, err := file.Read(buf) + lines += bytes.Count(buf[:n], []byte{'\n'}) + if err == io.EOF { + break + } + if err != nil { + return 0, err + } + } + if _, err := file.Seek(0, io.SeekStart); err != nil { + return 0, err + } + return lines, nil +} + +func (p *pgxPrinter) shouldUsePagerMetrics(size int64, lines int) bool { + switch p.pagerMode { + case pagerModeNever: + return false + case pagerModeAlways: + return p.isTerminal && p.pagerSupported + default: + if !p.isTerminal || !p.pagerSupported { + return false + } + return size >= autoPagerMinBytes || lines > p.autoPagerLineThreshold() + } +} + func (p *pgxPrinter) shouldUsePager(str string) bool { switch p.pagerMode { case pagerModeNever: diff --git a/internal/cliio/printer_test.go b/internal/cliio/printer_test.go index 193e7c8..39496c6 100644 --- a/internal/cliio/printer_test.go +++ b/internal/cliio/printer_test.go @@ -109,3 +109,19 @@ func TestPrintViaPager_AppendsNewlineWhenNotPresent(t *testing.T) { p.PrintViaPager("SELECT 9") assert.Equal(t, "SELECT 9\n", out.String()) } + +func TestStreamViaPager_WritesWithoutBuildingOutputString(t *testing.T) { + out := &bytes.Buffer{} + p := &pgxPrinter{ + out: out, + errOut: io.Discard, + pagerMode: pagerModeNever, + } + + err := p.StreamViaPager(func(w io.Writer) error { + _, err := io.WriteString(w, "header\nrow\n") + return err + }) + assert.NoError(t, err) + assert.Equal(t, "header\nrow\n", out.String()) +}