diff options
| author | Paul Buetow <paul@buetow.org> | 2026-07-22 23:51:18 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-07-22 23:51:18 +0300 |
| commit | 849951be1d1a7ee9f9302006ccb187bf5b4e36f3 (patch) | |
| tree | 496c924a03a9ea6212e29bb4699e268066ebad81 /internal/clients/handlers/basehandler.go | |
| parent | bf78b3abffee6d49c08ca2980156afc455994969 (diff) | |
feat: DTail fork — server/client feature development
Squashed development of the snonux/dtail fork's product code (internal/, cmd/)
since diverging from mimecast/dtail. Major areas:
- Read/output path: the former "turbo" channel-less path is now the single,
default server-side read/output path for cat/grep/tail and MapReduce; the old
channel-based path and its config/env toggles were removed.
- MapReduce: single aggregate implementation (server + serverless) fed directly
by a processor pipeline, with input-exhausted finalization via the shutdown
coordinator; high-concurrency and data-race fixes.
- Journal source reads (journal:unit.service) via journalctl, Linux-gated behind
a journal-v1 capability.
- Auth-key fast reconnect: in-memory per-user public-key cache with TTL/max-keys,
registered over an authenticated session (AUTHKEY), checked before
authorized_keys.
- Interactive query reload (--interactive-query) with SESSION START/UPDATE
generation boundaries and capability negotiation.
- Client-side deadlines: --timeout / --shutdownAfter as context deadlines;
follow shutdown handling.
- Client logging: diagnostics-only daily log by default, opt-in payload tee via
--log-payload.
- Numerous correctness fixes (buffer-pool double-recycle races, EOF-sentinel
leaks, glob-expansion cap, TOCTOU in CSV parsing) with accompanying unit tests.
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
Diffstat (limited to 'internal/clients/handlers/basehandler.go')
| -rw-r--r-- | internal/clients/handlers/basehandler.go | 319 |
1 files changed, 311 insertions, 8 deletions
diff --git a/internal/clients/handlers/basehandler.go b/internal/clients/handlers/basehandler.go index 6f637a7..188c3aa 100644 --- a/internal/clients/handlers/basehandler.go +++ b/internal/clients/handlers/basehandler.go @@ -5,7 +5,10 @@ import ( "encoding/base64" "fmt" "io" + "sort" + "strconv" "strings" + "sync" "time" "github.com/mimecast/dtail/internal" @@ -19,7 +22,29 @@ type baseHandler struct { shellStarted bool commands chan string receiveBuf bytes.Buffer - status int + + // pendingCommand holds the unsent tail of a dequeued command frame. A + // frame (e.g. a large base64-encoded regex or MapReduce query) can exceed + // the buffer io.Copy hands to Read (32KB), so the remainder is kept here + // and drained by subsequent Read calls instead of being dropped, which + // would corrupt the server-side command stream. Only touched by Read + // (single output-copy goroutine). + pendingCommand []byte + status int + + capabilitiesMu sync.RWMutex + capabilities map[string]struct{} + capabilitiesCh chan struct{} + capabilitiesOk sync.Once + + sessionAcks chan SessionAck +} + +// SessionAck is a parsed hidden acknowledgement for SESSION START/UPDATE requests. +type SessionAck struct { + Action string + Generation uint64 + Error string } func (h *baseHandler) String() string { @@ -40,6 +65,37 @@ func (h *baseHandler) Status() int { return h.status } +func (h *baseHandler) Capabilities() []string { + h.capabilitiesMu.RLock() + defer h.capabilitiesMu.RUnlock() + + capabilities := make([]string, 0, len(h.capabilities)) + for capability := range h.capabilities { + capabilities = append(capabilities, capability) + } + sort.Strings(capabilities) + return capabilities +} + +func (h *baseHandler) HasCapability(name string) bool { + h.capabilitiesMu.RLock() + defer h.capabilitiesMu.RUnlock() + + _, ok := h.capabilities[name] + return ok +} + +func (h *baseHandler) ReportServerError(message string) { + h.status = 1 + // Route through the DIAGNOSTIC (Log) sink, not Raw: a server-error report is + // an audit line, not bulk payload. Via Raw it would be gated out of the + // client log file whenever Client.LogPayload is false (the default), silently + // dropping the error from the on-disk audit trail. RawLog keeps it in the + // file like other diagnostics while still printing it to stdout. The message + // carries no trailing newline; the Log sink appends one. + dlog.Client.RawLog(formatServerErrorMessage(h.server, message)) +} + // SendMessage to the server. func (h *baseHandler) SendMessage(command string) error { encoded := base64.StdEncoding.EncodeToString([]byte(command)) @@ -61,10 +117,8 @@ func (h *baseHandler) Write(p []byte) (n int, err error) { for _, b := range p { switch b { case '\n': - // Backwards compatible with DTail 3 (e.g. get error message from server - // about protocol missmatch. + // Just add the newline to the buffer, don't treat as message delimiter h.receiveBuf.WriteByte(b) - fallthrough case protocol.MessageDelimiter: message := h.receiveBuf.String() h.handleMessage(message) @@ -77,14 +131,58 @@ func (h *baseHandler) Write(p []byte) (n int, err error) { } // Send data to the dtail server via Reader interface. +// +// Priority select: when Done() is closed we must still drain any pending +// commands before returning io.EOF, because closing the connection requires +// the '.ack close connection' message to be flushed to the server first. +// Without this drain the server waits up to 5 seconds for the ack. +// +// Command frames larger than p are delivered across multiple Read calls via +// pendingCommand (see consumeCommand); the leftover is always drained before +// a new command is dequeued so frames are never truncated or interleaved. func (h *baseHandler) Read(p []byte) (n int, err error) { + if len(h.pendingCommand) > 0 { + n = copy(p, h.pendingCommand) + h.pendingCommand = h.pendingCommand[n:] + if len(h.pendingCommand) == 0 { + // Release the backing array once the frame is fully delivered. + // Keeping it would retain the largest-ever frame for the handler + // lifetime and let the slice creep forward through its backing + // array on reuse; oversized frames are rare, so consumeCommand + // simply allocates a fresh buffer next time instead. + h.pendingCommand = nil + } + return + } + + // Check for a pending command first (non-blocking), giving it priority + // over the Done signal so that queued acks are always delivered. select { case command := <-h.commands: - n = copy(p, []byte(command)) + return h.consumeCommand(p, command), nil + default: + } + + // No command is immediately ready; block on whichever arrives first. + select { + case command := <-h.commands: + return h.consumeCommand(p, command), nil case <-h.Done(): return 0, io.EOF } - return +} + +// consumeCommand copies as much of command as fits into p and stashes the +// remainder in pendingCommand so the next Read calls can deliver the rest of +// the frame. pendingCommand is always empty here (Read drains and releases it +// before dequeuing a new command), so a fresh buffer is allocated for the +// remainder rather than reusing a long-lived one. +func (h *baseHandler) consumeCommand(p []byte, command string) int { + n := copy(p, command) + if n < len(command) { + h.pendingCommand = []byte(command[n:]) + } + return n } func (h *baseHandler) handleMessage(message string) { @@ -92,24 +190,229 @@ func (h *baseHandler) handleMessage(message string) { h.handleHiddenMessage(message) return } + if h.handleAuthKeyMessage(message) { + return + } + + // Add newline only if the message doesn't already end with one + if len(message) > 0 && message[len(message)-1] == '\n' { + dlog.Client.Raw(message) + } else { + dlog.Client.Raw(message + "\n") + } +} + +func (h *baseHandler) handleAuthKeyMessage(message string) bool { + isAuthKeyMessage, authKeyOK, authKeyDetail := parseAuthKeyMessage(message) + if !isAuthKeyMessage { + return false + } + + if authKeyOK { + dlog.Client.Debug(h.server, "AUTHKEY registration accepted by server") + return true + } + + if authKeyDetail == "" { + dlog.Client.Warn(h.server, "AUTHKEY registration failed") + return true + } + + dlog.Client.Warn(h.server, "AUTHKEY registration failed", authKeyDetail) + return true +} + +func parseAuthKeyMessage(message string) (isAuthKeyMessage bool, ok bool, detail string) { + if message == "" { + return false, false, "" + } + + payload := strings.TrimSpace(message) + parts := strings.Split(payload, protocol.FieldDelimiter) + if len(parts) > 0 { + payload = strings.TrimSpace(parts[len(parts)-1]) + } - dlog.Client.Raw(message) + switch { + case payload == "AUTHKEY OK": + return true, true, "" + case strings.HasPrefix(payload, "AUTHKEY ERR"): + detail := strings.TrimSpace(strings.TrimPrefix(payload, "AUTHKEY ERR")) + return true, false, detail + default: + return false, false, "" + } } // Handle messages received from server which are not meant to be displayed // to the end user. func (h *baseHandler) handleHiddenMessage(message string) { switch { + case strings.HasPrefix(message, protocol.HiddenCapabilitiesPrefix): + h.handleCapabilitiesMessage(message) + case strings.HasPrefix(message, protocol.HiddenSessionStartOKPrefix), + strings.HasPrefix(message, protocol.HiddenSessionUpdateOKPrefix), + strings.HasPrefix(message, protocol.HiddenSessionErrorPrefix): + h.handleSessionAckMessage(message) case strings.HasPrefix(message, ".syn close connection"): - go h.SendMessage(".ack close connection") + if err := h.SendMessage(".ack close connection"); err != nil { + dlog.Client.Debug(h.server, "Unable to acknowledge close connection", err) + } h.Shutdown() } } +func (h *baseHandler) handleCapabilitiesMessage(message string) { + capabilities := strings.Fields(strings.TrimPrefix(message, protocol.HiddenCapabilitiesPrefix)) + + h.capabilitiesMu.Lock() + defer h.capabilitiesMu.Unlock() + + if h.capabilities == nil { + h.capabilities = make(map[string]struct{}) + } + for _, capability := range capabilities { + if capability == "" { + continue + } + h.capabilities[capability] = struct{}{} + } + + h.capabilitiesOk.Do(func() { + if h.capabilitiesCh != nil { + close(h.capabilitiesCh) + } + }) +} + func (h *baseHandler) Done() <-chan struct{} { return h.done.Done() } +func (h *baseHandler) WaitForCapabilities(timeout time.Duration) bool { + if h.capabilitiesCh == nil { + return false + } + + if timeout <= 0 { + select { + case <-h.capabilitiesCh: + return true + default: + return false + } + } + + timer := time.NewTimer(timeout) + defer timer.Stop() + + select { + case <-h.capabilitiesCh: + return true + case <-h.Done(): + return false + case <-timer.C: + return false + } +} + +func (h *baseHandler) WaitForSessionAck(timeout time.Duration) (SessionAck, bool) { + if h.sessionAcks == nil { + return SessionAck{}, false + } + + if timeout <= 0 { + select { + case ack := <-h.sessionAcks: + return ack, true + default: + return SessionAck{}, false + } + } + + timer := time.NewTimer(timeout) + defer timer.Stop() + + select { + case ack := <-h.sessionAcks: + return ack, true + case <-h.Done(): + return SessionAck{}, false + case <-timer.C: + return SessionAck{}, false + } +} + func (h *baseHandler) Shutdown() { h.done.Shutdown() } + +func (h *baseHandler) handleSessionAckMessage(message string) { + ack, ok := parseSessionAckMessage(message) + if !ok { + dlog.Client.Warn(h.server, "Unable to parse session acknowledgement", message) + return + } + if h.sessionAcks == nil { + return + } + + select { + case h.sessionAcks <- ack: + case <-h.Done(): + default: + dlog.Client.Warn(h.server, "Dropping session acknowledgement because the queue is full", message) + } +} + +func parseSessionAckMessage(message string) (SessionAck, bool) { + payload := strings.TrimSpace(message) + if payload == "" { + return SessionAck{}, false + } + + switch { + case strings.HasPrefix(payload, protocol.HiddenSessionStartOKPrefix): + return parseSessionOKAck(strings.TrimPrefix(payload, protocol.HiddenSessionStartOKPrefix), "start") + case strings.HasPrefix(payload, protocol.HiddenSessionUpdateOKPrefix): + return parseSessionOKAck(strings.TrimPrefix(payload, protocol.HiddenSessionUpdateOKPrefix), "update") + case strings.HasPrefix(payload, protocol.HiddenSessionErrorPrefix): + return SessionAck{ + Action: "error", + Error: strings.TrimSpace(strings.TrimPrefix(payload, protocol.HiddenSessionErrorPrefix)), + }, true + default: + return SessionAck{}, false + } +} + +func parseSessionOKAck(payload string, action string) (SessionAck, bool) { + generationStr := strings.TrimSpace(payload) + if generationStr == "" { + return SessionAck{}, false + } + + generation, err := strconv.ParseUint(generationStr, 10, 64) + if err != nil { + return SessionAck{}, false + } + + return SessionAck{ + Action: action, + Generation: generation, + }, true +} + +// formatServerErrorMessage builds the "SERVER|<server>|ERROR|<message>" audit +// line shown to the user and written to the client log file. It carries NO +// trailing newline: it is emitted via the diagnostic (Log) sink, which appends +// the newline itself (adding one here would produce a blank line). +func formatServerErrorMessage(server string, message string) string { + return fmt.Sprintf("SERVER%s%s%sERROR%s%s", + protocol.FieldDelimiter, + server, + protocol.FieldDelimiter, + protocol.FieldDelimiter, + message, + ) +} |
