From 55eacb59e9afc0eca21ff5383df9e88567cb57a1 Mon Sep 17 00:00:00 2001 From: Aarnav Tale Date: Sat, 31 May 2025 09:45:47 -0400 Subject: [PATCH] feat: expand frame type to support stdout/stderr chan --- agent/cmd/hp_agent/hp_agent.go | 4 +- agent/internal/hpagent/handler.go | 34 +----- agent/internal/sshutil/connect.go | 189 +++++++++++++---------------- agent/internal/sshutil/encoder.go | 130 ++++++++++++++++---- agent/internal/sshutil/handler.go | 140 +++++++++++++++++---- agent/internal/sshutil/sessions.go | 54 ++++++--- app/server/agent/dispatcher.ts | 33 ++++- app/server/agent/encoder.ts | 99 +++++++++++++++ app/server/agent/ssh.ts | 136 +++++++++++---------- app/server/index.ts | 22 +++- package.json | 2 +- pnpm-lock.yaml | 22 ++-- 12 files changed, 580 insertions(+), 285 deletions(-) create mode 100644 app/server/agent/encoder.ts diff --git a/agent/cmd/hp_agent/hp_agent.go b/agent/cmd/hp_agent/hp_agent.go index d4dcfe5..03d90bb 100644 --- a/agent/cmd/hp_agent/hp_agent.go +++ b/agent/cmd/hp_agent/hp_agent.go @@ -1,8 +1,6 @@ package main import ( - "os" - _ "github.com/joho/godotenv/autoload" "github.com/tale/headplane/agent/internal/config" "github.com/tale/headplane/agent/internal/hpagent" @@ -34,6 +32,6 @@ func main() { ID: agent.ID, }) - sshutil.StartInputReader(os.NewFile(3, "sshin")) + sshutil.DispatchSSHStdin() hpagent.FollowMaster(agent) } diff --git a/agent/internal/hpagent/handler.go b/agent/internal/hpagent/handler.go index 76868d6..7968d3d 100644 --- a/agent/internal/hpagent/handler.go +++ b/agent/internal/hpagent/handler.go @@ -20,24 +20,10 @@ type RecvMessage struct { } type CborMessage struct { - Op string `cbor:"op"` + Op string `cbor:"op"` Payload cbor.RawMessage `cbor:"payload"` } -type SSHConnect struct { - SessionId string `cbor:"sessionId"` - Username string `cbor:"username"` - Hostname string `cbor:"hostname"` - Port int `cbor:"port"` -} - -type SSHMessage struct { - op string - username string - hostname string - Id string -} - type SendMessage struct { Type string Data any @@ -51,7 +37,7 @@ func FollowMaster(agent *tsnet.TSAgent) { for scanner.Scan() { line := scanner.Bytes() - log.Info("Got bytes delimited by newline"); + log.Info("Got bytes delimited by newline") var msg CborMessage decoder := cbor.NewDecoder(bytes.NewReader(line)) @@ -59,28 +45,20 @@ func FollowMaster(agent *tsnet.TSAgent) { if err != nil { log.Error("Unable to decode message from master: %s", err) - continue; + continue } log.Debug("Received message from master: %s", msg) - var sshPayload SSHConnect + var sshPayload sshutil.SSHConnectPayload err = cbor.Unmarshal(msg.Payload, &sshPayload) if err != nil { log.Error("Unable to unmarshal SSH connect payload: %s", err) - continue; + continue } - log.Info("Opening SSH PTY for session %s to %s@%s:%d", sshPayload.SessionId, sshPayload.Username, sshPayload.Hostname, sshPayload.Port) - sshutil.OpenSshPty(agent, sshutil.SshConnectParams{ - Hostname: sshPayload.Hostname, - Port: sshPayload.Port, - Username: sshPayload.Username, - Id: sshPayload.SessionId, - }) + sshutil.StartWebSSH(agent, sshPayload) } - - // var msg RecvMessage // err := json.Unmarshal(line, &msg) // if err != nil { diff --git a/agent/internal/sshutil/connect.go b/agent/internal/sshutil/connect.go index 3e35aa6..0fee7c7 100644 --- a/agent/internal/sshutil/connect.go +++ b/agent/internal/sshutil/connect.go @@ -3,8 +3,6 @@ package sshutil import ( "context" "errors" - "io" - "os" "strconv" "strings" @@ -13,149 +11,128 @@ import ( "golang.org/x/crypto/ssh" ) -type SshConnectParams struct { - Hostname string - Port int - Username string - Id string +type SSHConnectPayload struct { + SessionId string `cbor:"sessionId"` + Username string `cbor:"username"` + Hostname string `cbor:"hostname"` + Port int `cbor:"port"` } -func dialAndValidateTailscaleSSH(agent *tsnet.TSAgent, params SshConnectParams) (*ssh.Client, error) { +func connectToTailscaleSSH(agent *tsnet.TSAgent, params SSHConnectPayload) (*ssh.Client, error) { log := util.GetLogger() - addr := strings.Join([]string{params.Hostname, ":", strconv.Itoa(params.Port)}, "") - log.Debug("Attempting to dial %s via Tailscale SSH", addr) - conn, err := agent.Dial(context.Background(), "tcp", addr) + log.Debug("Initiating Tailscale SSH connection to %s@%s", params.Username, addr) + tailnetConn, err := agent.Dial(context.Background(), "tcp", addr) if err != nil { - log.Error("Failed to connect to Tailscale SSH: %s", err) return nil, err } - log.Debug("Connected to Tailscale SSH at %s", addr) + log.Debug("Routed connection via tsnet to %s", addr) config := &ssh.ClientConfig{ User: params.Username, + // This isn't a concern because we are only dialing within the Tailnet + // and every device is trusted and *should* be ACL accessible. HostKeyCallback: ssh.InsecureIgnoreHostKey(), } - clientConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config) + conn, chans, reqs, err := ssh.NewClientConn(tailnetConn, addr, config) if err != nil { - log.Error("Failed to create SSH client connection: %s", err) conn.Close() return nil, err } - client := ssh.NewClient(clientConn, chans, reqs) - sVer := string(client.ServerVersion()) + // At this point we have successfully connected to the node + sshClient := ssh.NewClient(conn, chans, reqs) + version := string(sshClient.ServerVersion()) - if !strings.Contains(sVer, "Tailscale") { - log.Error("Connected to non-Tailscale SSH server: %s", sVer) + if !strings.Contains(version, "Tailscale") { conn.Close() - return nil, errors.New("not a Tailscale SSH server") + return nil, errors.New("server is not running Tailscale SSH") } - log.Info("Connected to SSH server running %s at %s", client.ServerVersion(), addr) - return client, nil + log.Info("Connected to %s@%s:%d via Tailscale SSH (%s)", params.Username, params.Hostname, params.Port, version) + return sshClient, nil } -func bindStdinToFd(sess *ssh.Session, fd int) error { - log := util.GetLogger() - - sshIn := os.NewFile(uintptr(fd), "sshInput") - if sshIn == nil { - log.Error("Failed to create file from stdin fd %d", fd) - return errors.New("failed to create file from stdin fd") - } - - stdin, err := sess.StdinPipe() - if err != nil { - log.Error("Failed to get stdin pipe: %s", err) - return err - } - - go io.Copy(stdin, sshIn) // From Node → SSH session - return nil -} - -func bindStdoutToFd(sess *ssh.Session, fd int) error { - log := util.GetLogger() - - sshOut := os.NewFile(uintptr(fd), "sshOutput") - if sshOut == nil { - log.Error("Failed to create file from stdout fd %d", fd) - return errors.New("failed to create file from stdout fd") - } - - stdout, err := sess.StdoutPipe() - if err != nil { - log.Error("Failed to get stdout pipe: %s", err) - return err - } - - go io.Copy(sshOut, stdout) // From SSH → Node - return nil -} - -func OpenSshPty(agent *tsnet.TSAgent, params SshConnectParams) (*ssh.Client, error) { +func StartWebSSH(agent *tsnet.TSAgent, params SSHConnectPayload) { log := util.GetLogger() if agent == nil { - log.Error("Tailscale agent is nil") - return nil, errors.New("tailscale agent is nil") + log.Error("tsnet.TSAgent is not initialized correctly") + return } - if params.Hostname == "" || params.Port <= 0 || params.Username == "" { - log.Error("Invalid SSH connection parameters: %+v", params) - return nil, errors.New("invalid SSH connection parameters") + if params.Hostname == "" || params.Port <= 0 || params.Username == "" || params.SessionId == "" { + log.Error("Invalid SSH connection parameters: %v", params) + return } - client, err := dialAndValidateTailscaleSSH(agent, params) + client, err := connectToTailscaleSSH(agent, params) if err != nil { - log.Error("Failed to open SSH pty: %s", err) - return nil, err + log.Error("Failed to connect to Tailscale SSH for (%s): %s", params.SessionId, err) + return } + // Everything in the func is related to the SSH session. + // Each session runs in its own goroutine, allowing concurrency. go func() { - sess, err := client.NewSession() - if err != nil { - log.Error("Failed to create new SSH session: %s", err) - client.Close() - } - - modes := ssh.TerminalModes{ - ssh.ECHO: 1, // enable echoing - ssh.TTY_OP_ISPEED: 14400, - ssh.TTY_OP_OSPEED: 14400, - } - - if err := sess.RequestPty("xterm-256color", 80, 40, modes); err != nil { - log.Error("Failed to request PTY: %s", err) - client.Close() - } - - ctx := addSession(params.Id, sess) - go func() { - for data := range ctx.InputCh { - if _, err := ctx.Stdin.Write(data); err != nil { - log.Error("Failed to write to SSH stdin: %s", err) - return - } + log.Debug("Creating SSH session for session ID: %s", params.SessionId) + sess, err := client.NewSession() + if err != nil { + log.Error("Failed to create new SSH session: %s", err) + client.Close() + return } - }(); - if err := sess.Shell(); err != nil { - log.Error("Failed to start shell: %s", err) + modes := ssh.TerminalModes{ + ssh.ECHO: 1, + ssh.TTY_OP_ISPEED: 14400, + ssh.TTY_OP_OSPEED: 14400, + } + + // Resize event is possible via the control channel later + err = sess.RequestPty("xterm-256color", 80, 40, modes) + if err != nil { + log.Error("Failed to request PTY for (%s): %s", params.SessionId, err) + return + } + + ctx, err := registerSessionChans(params.SessionId, sess) + if err != nil { + log.Error("Failed to register session channels for (%s): %s", params.SessionId, err) + client.Close() + return + } + + // Input buffer handler + go func() { + for data := range ctx.InputCh { + _, err := ctx.Stdin.Write(data) + if err != nil { + log.Error("Failed to write to SSH stdin: %s", err) + return + } + } + }() + + // This spawns 2 goroutins for stdout and stderr + dispatchSSHStdout(params.SessionId, ctx.Stdout, ctx.Stderr) + + // Spin up a shell and wait for the pty to terminate + err = sess.Shell() + if err != nil { + log.Error("Failed to start shell for (%s): %s", params.SessionId, err) + client.Close() + return + } + + log.Info("Opened an SSH PTY for %s", params.SessionId) + sess.Wait() + sess.Close() client.Close() - } - - log.Info("Successfully opened SSH pty for %s@%s:%d", params.Username, params.Hostname, params.Port) - go streamSSHOutput(params.Id, ctx.Stdout, os.NewFile(4, "sshOutput")) - sess.Wait(); - sess.Close(); - log.Info("SSH session %s closed (goSide)", params.Id) - RemoveSession(params.Id) + log.Info("SSH session for %s closed", params.SessionId) + RemoveSession(params.SessionId) }() - - return client, nil } diff --git a/agent/internal/sshutil/encoder.go b/agent/internal/sshutil/encoder.go index 66d307a..f90e400 100644 --- a/agent/internal/sshutil/encoder.go +++ b/agent/internal/sshutil/encoder.go @@ -5,37 +5,121 @@ import ( "fmt" ) -func encodeFrame(id string, data []byte) ([]byte, error) { - sid := []byte(id) - if len(sid) > 255 { - return nil, fmt.Errorf("session ID too long") +// An SSH frame is used to wrap raw binary data to and from an SSH session +// in order to allow multiplexing connections over a single file descriptor. +// +// In practice, this is how we can easily support multiple SSH connections +// through the 2 file descriptors created by the parent node process. +// +// This is the format of an SSH frame: +// - Magic: The first 4 bytes are HPLS (0x48504C53) to identify the frame. +// - Version Byte: The first byte is the version of the frame format. +// - Channel Type: The second byte indicates the type of channel. +// - Session ID: The next bytes are the length and actual session ID. +// - Payload: The remaining bytes are the payload length and actual data. +// +// +---------+----------+--------------+-------------+----------+ +// | Magic | Version | Channel Type | SID Length | SID | +// | 4 bytes | 1 byte | 1 byte | 1 byte (S) | S bytes | +// +---------+----------+--------------+------------------------+ +// | Payload Length | Payload | +// | 4 bytes (u32, P) | P bytes | +// +--------------------+---------------------------------------+ + +const ( + MagicString = "HPLS" + VersionByte = 1 +) + +type ChannelType int + +const ( + ChannelTypeStdin ChannelType = iota + ChannelTypeStdout + ChannelTypeStderr +) + +type SSHFrame struct { + ChannelType ChannelType + SessionID string + Payload []byte + + Length func() int +} + +type HPLSFrame1 struct{} + +func (t HPLSFrame1) Encode(frame SSHFrame) ([]byte, error) { + frameChan := frame.ChannelType + switch frameChan { + case ChannelTypeStdin, ChannelTypeStdout, ChannelTypeStderr: + default: + return nil, fmt.Errorf("invalid channel type: %d", frameChan) } - payloadLen := len(data) - buf := make([]byte, 1+len(sid)+4+payloadLen) + if len(frame.SessionID) == 0 { + return nil, fmt.Errorf("session ID cannot be empty") + } - buf[0] = byte(len(sid)) // SID length - copy(buf[1:], sid) // SID - binary.BigEndian.PutUint32(buf[1+len(sid):], uint32(payloadLen)) // Payload length - copy(buf[1+len(sid)+4:], data) // Payload + if len(frame.Payload) == 0 { + return nil, fmt.Errorf("payload cannot be empty") + } + sid := []byte(frame.SessionID) + if len(sid) > 255 { + return nil, fmt.Errorf("session ID exceeds 255 byte limit") + } + + if len(frame.Payload) > 0xFFFFFFFF { + return nil, fmt.Errorf("payload exceeds 4GB limit") + } + + frameLen := 4 // Magic + frameLen += 1 // Version byte + frameLen += 1 // Channel type + frameLen += 1 + len(sid) // Session ID length + SID + frameLen += 4 + len(frame.Payload) // Payload length + Payload + + buf := make([]byte, frameLen) + copy(buf[0:4], []byte(MagicString)) + buf[4] = VersionByte + buf[5] = byte(frameChan) + buf[6] = byte(len(sid)) + + offset := 7 + len(sid) + copy(buf[7:offset], sid) + + binary.BigEndian.PutUint32(buf[offset:offset+4], uint32(len(frame.Payload))) + copy(buf[offset+4:], frame.Payload) return buf, nil } -func decodeFrame(buf []byte) (id string, payload []byte, ok bool) { - if len(buf) < 5 { - return "", nil, false +func (t HPLSFrame1) Decode(buf []byte) (SSHFrame, error) { + frame := SSHFrame{} + if len(buf) < 5 || string(buf[0:4]) != MagicString || buf[4] != VersionByte { + return frame, fmt.Errorf("illegal HPLS1 frame format") } - sidLen := int(buf[0]) - if len(buf) < 1+sidLen+4 { - return "", nil, false - } - id = string(buf[1 : 1+sidLen]) - payloadLen := int(binary.BigEndian.Uint32(buf[1+sidLen:])) - if len(buf) < 1+sidLen+4+payloadLen { - return "", nil, false + frame.ChannelType = ChannelType(buf[5]) + if frame.ChannelType < ChannelTypeStdin || frame.ChannelType > ChannelTypeStderr { + return frame, fmt.Errorf("invalid channel type: %d", frame.ChannelType) } - payload = buf[1+sidLen+4 : 1+sidLen+4+payloadLen] - return id, payload, true + + sidLen := int(buf[6]) + if len(buf) < 7+sidLen+4 { + return frame, fmt.Errorf("buffer too short for session ID and payload length") + } + + frame.SessionID = string(buf[7 : 7+sidLen]) + payloadLen := int(binary.BigEndian.Uint32(buf[7+sidLen:])) + if len(buf) < 7+sidLen+4+payloadLen { + return frame, fmt.Errorf("buffer too short for payload") + } + + frame.Payload = buf[7+sidLen+4 : 7+sidLen+4+payloadLen] + frame.Length = func() int { + return 7 + sidLen + 4 + payloadLen + } + + return frame, nil } diff --git a/agent/internal/sshutil/handler.go b/agent/internal/sshutil/handler.go index db2722a..e9b64ff 100644 --- a/agent/internal/sshutil/handler.go +++ b/agent/internal/sshutil/handler.go @@ -2,55 +2,145 @@ package sshutil import ( "io" + "os" "github.com/tale/headplane/agent/internal/util" ) -func StartInputReader(fd3 io.Reader) { - log := util.GetLogger(); +// The file descriptors attached by the parent node process +const ( + InputFd = 3 + OutputFd = 4 +) - log.Info("Starting SSH fd3 input reader") +// DispatchSSHStdin listens for SSH stdin frames on the InputFd file descriptor +// and writes the payload to the appropriate session's stdin. +// +// This function runs in a goroutine in main and is responsible for dispatching +// to ALL connections, not just its own like dispatchSSHStdout does. +func DispatchSSHStdin() { + log := util.GetLogger() + + log.Debug("Opening file descriptor: %d for SSH stdin", InputFd) + fd := os.NewFile(InputFd, "ssh_stdin") + if fd == nil { + log.Error("Failed to open file descriptor %d for SSH stdin", InputFd) + return + } + + log.Info("Listening for SSH stdin on fd %d", InputFd) go func() { buffer := make([]byte, 8192) + hpls1 := HPLSFrame1{} + for { - n, err := fd3.Read(buffer) + // This is the only check where we can detect if the descriptor + // was closed so we can return and exit the goroutine. + bufCount, err := fd.Read(buffer) if err != nil { - log.Error("fd3 read error: %v", err) + log.Error("Failed to read from SSH stdin: %v", err) return } - offset := 0 - for offset < n { - id, payload, ok := decodeFrame(buffer[offset:n]) - if !ok { - break // Wait for more data - } - offset += 1 + len(id) + 4 + len(payload) - sessionsMu.RLock() - sess, ok := sessions[id] - sessionsMu.RUnlock() - if !ok { - log.Error("invalid session id: %s", id) - continue - } - _, err := sess.Stdin.Write(payload) + // Check if we have 0 EOF, which means the descriptor was closed. + if bufCount == 0 { + log.Info("SSH stdin closed, stopping listener") + return + } + + offset := 0 + for offset < bufCount { + frame, err := hpls1.Decode(buffer[offset:bufCount]) if err != nil { - log.Error("failed to write to session stdin: %v", err) + // We need to wait for more data to decode the frame + break + } + + if frame.ChannelType != ChannelTypeStdin { + log.Error("Received invalid channel type: %d, expected %d", frame.ChannelType, ChannelTypeStdin) continue } + + offset += frame.Length() + log.Debug("Received SSH stdin frame: %s", frame.SessionID) + sess, ok := lookupSession(frame.SessionID) + if !ok { + log.Error("Invalid session ID: %s", frame.SessionID) + continue + } + + // Write the payload to the session's stdin + writeCount, err := sess.Stdin.Write(frame.Payload) + if err != nil { + log.Error("Failed to write to session stdin: %v", err) + continue + } + + log.Debug("Wrote %d bytes to session %s stdin", writeCount, frame.SessionID) } } }() } -func streamSSHOutput(id string, r io.Reader, fd4 io.Writer) { +func dispatchSSHStdout(id string, stdout io.Reader, stderr io.Reader) { + log := util.GetLogger() + + log.Debug("Opening file descriptor: %d for SSH stdout", OutputFd) + fd := os.NewFile(OutputFd, "ssh_stdout") + if fd == nil { + log.Error("Failed to open file descriptor %d for SSH stdout", OutputFd) + return + } + + go readerStreamRoutine(StreamRoutine{ + SessionID: id, + Reader: stdout, + Writer: fd, + ChannelType: ChannelTypeStdout, + }) + + go readerStreamRoutine(StreamRoutine{ + SessionID: id, + Reader: stderr, + Writer: fd, + ChannelType: ChannelTypeStderr, + }) +} + +type StreamRoutine struct { + SessionID string + Reader io.Reader + Writer io.Writer + ChannelType ChannelType +} + +func readerStreamRoutine(routine StreamRoutine) { + hpls1 := HPLSFrame1{} buf := make([]byte, 1024) for { - n, err := r.Read(buf) + byteCount, err := routine.Reader.Read(buf) if err != nil { + if err != io.EOF { + util.GetLogger().Error("Failed to read from reader: %v", err) + } + + break + } + + frame, err := hpls1.Encode(SSHFrame{ + ChannelType: routine.ChannelType, + SessionID: routine.SessionID, + Payload: buf[:byteCount], + }) + + if err != nil { + util.GetLogger().Error("Failed to encode frame: %v", err) + continue + } + + if _, err := routine.Writer.Write(frame); err != nil { + util.GetLogger().Error("Failed to write frame to writer: %v", err) break } - frame, _ := encodeFrame(id, buf[:n]) - fd4.Write(frame) } } diff --git a/agent/internal/sshutil/sessions.go b/agent/internal/sshutil/sessions.go index 7c07c39..3bfb04d 100644 --- a/agent/internal/sshutil/sessions.go +++ b/agent/internal/sshutil/sessions.go @@ -1,9 +1,11 @@ package sshutil import ( + "errors" "io" "sync" + "github.com/tale/headplane/agent/internal/util" "golang.org/x/crypto/ssh" ) @@ -12,56 +14,70 @@ type SessionContext struct { Session *ssh.Session Stdin io.WriteCloser Stdout io.Reader + Stderr io.Reader InputCh chan []byte } var sessions = make(map[string]*SessionContext) -var sessionsMu sync.RWMutex +var sessionsLock sync.RWMutex -func addSession(id string, session *ssh.Session) *SessionContext { - sessionsMu.Lock() - defer sessionsMu.Unlock() +func registerSessionChans(id string, session *ssh.Session) (*SessionContext, error) { + log := util.GetLogger() + + sessionsLock.Lock() + defer sessionsLock.Unlock() if _, exists := sessions[id]; exists { - return nil // Session with this ID already exists + return sessions[id], nil } stdin, err := session.StdinPipe() if err != nil { - return nil // Handle error appropriately in production code + return nil, errors.New("failed to create stdin pipe: " + err.Error()) } stdout, err := session.StdoutPipe() if err != nil { - stdin.Close() // Close stdin if stdout pipe creation fails - return nil // Handle error appropriately in production code + stdin.Close() + return nil, errors.New("failed to create stdout pipe: " + err.Error()) } - sessionContext := &SessionContext{ + stderr, err := session.StderrPipe() + if err != nil { + stdin.Close() + return nil, errors.New("failed to create stderr pipe: " + err.Error()) + } + + ctx := &SessionContext{ ID: id, Session: session, Stdin: stdin, Stdout: stdout, + Stderr: stderr, + // Buffered channel to queue input data + InputCh: make(chan []byte, 256), } - sessions[id] = sessionContext - return sessionContext + sessions[id] = ctx + log.Debug("Registered session %s with stdin, stdout, and stderr pipes", id) + return ctx, nil } -func GetSession(id string) *SessionContext { - sessionsMu.RLock() - defer sessionsMu.RUnlock() +func lookupSession(id string) (*SessionContext, bool) { + sessionsLock.RLock() + defer sessionsLock.RUnlock() - return sessions[id] // Returns nil if session does not exist + sessionContext, exists := sessions[id] + return sessionContext, exists } func RemoveSession(id string) { - sessionsMu.Lock() - defer sessionsMu.Unlock() + sessionsLock.Lock() + defer sessionsLock.Unlock() if sessionContext, exists := sessions[id]; exists { - sessionContext.Stdin.Close() // Close the stdin pipe + sessionContext.Stdin.Close() // Close the stdin pipe sessionContext.Session.Close() // Close the SSH session - delete(sessions, id) // Remove from the map + delete(sessions, id) // Remove from the map } } diff --git a/app/server/agent/dispatcher.ts b/app/server/agent/dispatcher.ts index c02d25c..42290bd 100644 --- a/app/server/agent/dispatcher.ts +++ b/app/server/agent/dispatcher.ts @@ -1,5 +1,7 @@ import type { Writable } from 'node:stream'; -import { decode, encode } from 'cbor2'; +import { encode } from 'cborg'; +import { WSContext } from 'hono/ws'; +import { ChannelType } from './encoder'; interface Command { op: string; @@ -25,7 +27,6 @@ export async function dispatchCommand( return new Promise((resolve, reject) => { const encodedCommand = Buffer.concat([encode(command), Buffer.from('\n')]); dispatcher.write(encodedCommand, (err) => { - console.log('Command dispatched:', command, err); if (err) { reject(err); } else { @@ -34,3 +35,31 @@ export async function dispatchCommand( }); }); } + +interface SSHConnectData extends Command { + op: 'ssh_conn_successful'; + payload: { + sessionId: string; + }; +} + +interface SSHConnectFailedData extends Command { + op: 'ssh_conn_failed'; + payload: { + reason: string; + }; +} + +interface SSHFrameData extends Command { + op: 'ssh_frame'; + payload: { + channel: ChannelType; + frame: Buffer; + }; +} + +type WebData = SSHConnectData | SSHConnectFailedData | SSHFrameData; + +export function dispatchWeb(dispatcher: WSContext, data: WebData) { + return dispatcher.send(encode(data)); +} diff --git a/app/server/agent/encoder.ts b/app/server/agent/encoder.ts new file mode 100644 index 0000000..ab7e43b --- /dev/null +++ b/app/server/agent/encoder.ts @@ -0,0 +1,99 @@ +// Refer to agent/internal/sshutil/encoder.go for more details +// This is the Node.js implementation of the SSH encoder +import log from '~/utils/log'; + +const MAGIC = 'HPLS'; +const VERSION = 1; + +// 0 -> Stdin +// 1 -> Stdout +// 2 -> Stderr +export type ChannelType = 0 | 1 | 2; + +interface SSHFrame { + sessionId: string; + channel: ChannelType; + payload: Blob | ArrayBufferLike | string; +} + +export async function encodeSSHFrame(frame: SSHFrame) { + const sid = Buffer.from(frame.sessionId, 'utf8'); + if (sid.length > 255) { + log.error('agent', 'SSH session ID too long: %s', frame.sessionId); + return; + } + + const payload = Buffer.isBuffer(frame.payload) + ? frame.payload + : typeof frame.payload === 'string' + ? Buffer.from(frame.payload, 'utf8') + : frame.payload instanceof Blob + ? Buffer.from(await frame.payload.arrayBuffer()) + : Buffer.from(frame.payload); + + // Size can only hold 4 bytes + if (payload.length > 0xffffffff) { + log.error('agent', 'SSH payload too large: %d bytes', payload.length); + return; + } + + const frameSize = + 4 + // Magic + 1 + // Version + 1 + // Channel Type + (1 + sid.length) + // Session ID length + SID + (4 + payload.length); // Payload length + Payload + + const buf = Buffer.alloc(frameSize); + buf.write(MAGIC, 0, 'utf8'); + buf.writeUInt8(VERSION, 4); + buf.writeUInt8(frame.channel, 5); + buf.writeUInt8(sid.length, 6); + + const offset = 7 + sid.length; + sid.copy(buf, 7); + + buf.writeUInt32BE(payload.length, offset); + payload.copy(buf, offset + 4); + return buf; +} + +export function decodeSSHFrame(data: Buffer) { + if (data.length < 5) { + log.error('agent', 'SSH frame too short: %d bytes', data.length); + return; + } + + const magic = data.toString('utf8', 0, 4); + const version = data.readUInt8(4); + + if (magic !== MAGIC || version !== VERSION) { + log.error('agent', 'Invalid SSH frame magic or version'); + return; + } + + const channel = data.readUInt8(5) as ChannelType; + const sidLength = data.readUInt8(6); + if (data.length < 7 + sidLength + 4) { + log.error('agent', 'SSH frame too short for session ID and payload'); + return; + } + + const sessionId = data.toString('utf8', 7, 7 + sidLength); + const payloadLength = data.readUInt32BE(7 + sidLength); + if (data.length < 7 + sidLength + 4 + payloadLength) { + log.error('agent', 'SSH frame too short for payload'); + return; + } + + const payload = data.subarray( + 7 + sidLength + 4, + 7 + sidLength + 4 + payloadLength, + ); + + return { + sessionId, + channel, + payload, + }; +} diff --git a/app/server/agent/ssh.ts b/app/server/agent/ssh.ts index 0656aa9..3cb1346 100644 --- a/app/server/agent/ssh.ts +++ b/app/server/agent/ssh.ts @@ -4,7 +4,8 @@ import type { Readable, Writable } from 'node:stream'; import { Context } from 'hono'; import { WSContext, WSEvents } from 'hono/ws'; import log from '~/utils/log'; -import { dispatchCommand } from './dispatcher'; +import { dispatchCommand, dispatchWeb } from './dispatcher'; +import { decodeSSHFrame, encodeSSHFrame } from './encoder'; interface SSHConnection { username: string; @@ -19,19 +20,44 @@ interface SSHSession { ws: WSContext; } +interface FrameDecodeSuccess { + id: string; + data: Buffer; +} + +interface FrameDecodeFailure { + id: undefined; + data: undefined; +} + export function createSSHMultiplexer(proc: ChildProcess): SSHMultiplexer { - return new SSHMultiplexer(proc); + const control = proc.stdin; + const sshInput = proc.stdio[3]; + const sshOutput = proc.stdio[4]; + + if (!control || !sshInput || !sshOutput) { + throw new Error('Invalid SSH multiplexer process: missing stdio streams'); + } + + return new SSHMultiplexer( + control, + sshInput as Writable, + sshOutput as Readable, + ); } export class SSHMultiplexer { private connections: Map; - private child: ChildProcess; + private control: Writable; + private sshInput: Writable; + private sshOutput: Readable; - constructor(proc: ChildProcess) { + constructor(control: Writable, sshInput: Writable, sshOutput: Readable) { this.connections = new Map(); - this.child = proc; - - this.handleStdout(); + this.control = control; + this.sshInput = sshInput; + this.sshOutput = sshOutput; + this.configureStdout(); } // TODO: Determine if we want to allow multiple connections for the same @@ -45,8 +71,8 @@ export class SSHMultiplexer { ws, }; - log.info('agent', 'Dispatching SSH connection for %s', sessionId); - await dispatchCommand(this.child.stdin, { + log.debug('agent', 'Dispatching SSH connection for %s', sessionId); + await dispatchCommand(this.control, { op: 'ssh_conn', payload: { sessionId, @@ -76,9 +102,22 @@ export class SSHMultiplexer { try { const sessionId = await this.connect(conn, ws); ws.raw = sessionId; - ws.send(JSON.stringify({ status: 'connected', sessionId })); + dispatchWeb(ws, { + op: 'ssh_conn_successful', + payload: { sessionId }, + }); } catch (error) { - ws.close(1011, `Connection failed: ${error.message}`); + const errorMessage = + error instanceof Error ? error.message : 'Unknown error'; + + dispatchWeb(ws, { + op: 'ssh_conn_failed', + payload: { + reason: errorMessage, + }, + }); + + ws.close(1011, 'Connection failed'); } }, @@ -95,8 +134,13 @@ export class SSHMultiplexer { return; } - const encodedFrame = this.encodeFrame(sessionId, event.data); - this.child.stdio[3]?.write(encodedFrame); + const encodedFrame = await encodeSSHFrame({ + sessionId, + channel: 0, // stdin + payload: event.data, + }); + + this.sshInput.write(encodedFrame); }, onClose: (_, ws) => { @@ -121,65 +165,35 @@ export class SSHMultiplexer { } log.error('agent', 'SSH WebSocket Error with %s', sessionId); - console.log(event); + log.debug('agent', 'Error details: %o', event); }, }; } - private encodeFrame(id: string, data: string | Buffer): Buffer { - const sid = Buffer.from(id, 'utf8'); - const payload = Buffer.isBuffer(data) ? data : Buffer.from(data, 'utf8'); - - // FIX: include +4 for the payload length - const frame = Buffer.alloc(1 + sid.length + 4 + payload.length); - frame.writeUint8(sid.length, 0); // 1 byte for sid length - sid.copy(frame, 1); // SID - frame.writeUint32BE(payload.length, 1 + sid.length); // 4 bytes for payload length - payload.copy(frame, 1 + sid.length + 4); // Payload - - return frame; - } - - private decodeFrame(frame: Buffer): { id: string; data: Buffer } | undefined { - if (frame.length < 5) return; - - const sidLength = frame.readUint8(0); - if (frame.length < 1 + sidLength + 4) return; - - const id = frame.slice(1, 1 + sidLength).toString('utf8'); - const payloadLength = frame.readUint32BE(1 + sidLength); - if (frame.length < 1 + sidLength + 4 + payloadLength) return; - - const data = frame.slice( - 1 + sidLength + 4, - 1 + sidLength + 4 + payloadLength, - ); - - return { id, data }; - } - - private handleStdout() { - const stdout = this.child.stdio[4]; - if (!stdout) { - return; - } - - stdout.on('data', (bytes) => { - console.log(Buffer.from(bytes).toString('utf8')); - const decoded = this.decodeFrame(bytes); - if (!decoded) { + private configureStdout() { + this.sshOutput.on('data', (bytes) => { + const frame = decodeSSHFrame(bytes); + if (!frame) { return; } - const { id, data } = decoded; - console.log(id, data); - const session = this.connections.get(id); + const session = this.connections.get(frame.sessionId); if (!session || !session.connected) { - log.warn('agent', 'Received data for disconnected session %s', id); + log.warn( + 'agent', + 'Received data for invalid SSH session %s', + frame.sessionId, + ); return; } - session.ws.send(data); + dispatchWeb(session.ws, { + op: 'ssh_frame', + payload: { + channel: frame.channel, + data: frame.payload, + }, + }); }); } } diff --git a/app/server/index.ts b/app/server/index.ts index 0bf31c4..f8c15ea 100644 --- a/app/server/index.ts +++ b/app/server/index.ts @@ -1,6 +1,6 @@ import { env, versions } from 'node:process'; +import type { WSEvents } from 'hono/ws'; import { createHonoServer } from 'react-router-hono-server/node'; - import log from '~/utils/log'; import { configureConfig, configureLogger, envVariables } from './config/env'; import { loadIntegration } from './config/integration'; @@ -97,7 +97,25 @@ export default createHonoServer({ app.get( '/_ssh_plexer', upgradeWebSocket((c) => { - return agentManager.multiplexer!.websocketHandler(c); + // MARK: This is a limitation of the hono NPM module we use + const wsHandler = agentManager.multiplexer?.websocketHandler( + c, + ) as WSEvents; + + return { + onOpen: wsHandler + ? wsHandler.onOpen + : (_, ws) => ws.close(1000, 'Multiplexer not available'), + onClose: wsHandler + ? wsHandler.onClose + : (_, ws) => ws.close(1000, 'Multiplexer not available'), + onMessage: wsHandler + ? wsHandler.onMessage + : (_, ws) => ws.close(1000, 'Multiplexer not available'), + onError: wsHandler + ? wsHandler.onError + : (_, ws) => ws.close(1000, 'Multiplexer error'), + }; }), ); }, diff --git a/package.json b/package.json index 7d9bd9c..d256b0d 100644 --- a/package.json +++ b/package.json @@ -29,7 +29,7 @@ "@uiw/codemirror-theme-xcode": "^4.23.12", "@uiw/react-codemirror": "^4.23.12", "arktype": "^2.1.20", - "cbor2": "^2.0.1", + "cborg": "^4.2.11", "clsx": "^2.1.1", "dotenv": "^16.5.0", "isbot": "^5.1.28", diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 3546ee5..dc8f2c0 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -64,9 +64,9 @@ importers: arktype: specifier: ^2.1.20 version: 2.1.20 - cbor2: - specifier: ^2.0.1 - version: 2.0.1 + cborg: + specifier: ^4.2.11 + version: 4.2.11 clsx: specifier: ^2.1.1 version: 2.1.1 @@ -412,10 +412,6 @@ packages: '@codemirror/view@6.36.8': resolution: {integrity: sha512-yoRo4f+FdnD01fFt4XpfpMCcCAo9QvZOtbrXExn4SqzH32YC6LgzqxfLZw/r6Ge65xyY03mK/UfUqrVw1gFiFg==} - '@cto.af/wtf8@0.0.2': - resolution: {integrity: sha512-ATm4UQiKrdm5GnU6BvIwUDN+LDEtt23zuzKFpnfDT59ULAd0aMYm/nSFzbSO02garLcXumRC13PzNfa7BsfvSg==} - engines: {node: '>=20'} - '@dnd-kit/accessibility@3.1.1': resolution: {integrity: sha512-2P+YgaXF+gRsIihwwY1gCsQSYnu9Zyj2py8kY5fFvUM1qm2WA2u639R6YNVfU4GWr+ZM5mqEsfHZZLoRONbemw==} peerDependencies: @@ -1702,9 +1698,9 @@ packages: caniuse-lite@1.0.30001718: resolution: {integrity: sha512-AflseV1ahcSunK53NfEs9gFWgOEmzr0f+kaMFA4xiLZlr9Hzt7HxcSpIFcnNCUkz6R6dWKa54rUz3HUmI3nVcw==} - cbor2@2.0.1: - resolution: {integrity: sha512-9bE8+tueGxONyxpttNKkAKKcGVtAPeoSJ64AjVTTjEuBOuRaeeP76EN9BbmQqkz1ZeTP0QPvksNBKwvEutIUzQ==} - engines: {node: '>=20'} + cborg@4.2.11: + resolution: {integrity: sha512-7gs3iaqtsD9OHowgqzc6ixQGwSBONqosVR2co0Bg0pARgrLap+LCcEIXJuuIz2jHy0WWQeDMFPEsU2r17I2XPQ==} + hasBin: true chokidar@4.0.3: resolution: {integrity: sha512-Qgzu8kfBvo+cA4962jnP1KkS6Dop5NS6g7R5LFYJr4b8Ub94PPQXUksCw9PvXoeXPRRddRNC5C1JQUR2SMGtnA==} @@ -3152,8 +3148,6 @@ snapshots: style-mod: 4.1.2 w3c-keyname: 2.2.8 - '@cto.af/wtf8@0.0.2': {} - '@dnd-kit/accessibility@3.1.1(react@19.1.0)': dependencies: react: 19.1.0 @@ -4817,9 +4811,7 @@ snapshots: caniuse-lite@1.0.30001718: {} - cbor2@2.0.1: - dependencies: - '@cto.af/wtf8': 0.0.2 + cborg@4.2.11: {} chokidar@4.0.3: dependencies: