From 4961e83a9c50c55dbd78610564af5ebab5335807 Mon Sep 17 00:00:00 2001 From: Aarnav Tale Date: Mon, 16 Jun 2025 11:38:13 -0400 Subject: [PATCH] fix: cleanup agent and add to ci --- cmd/hp_agent/hp_agent.go | 10 +- go.mod | 2 +- internal/config/config.go | 45 +++++++ internal/config/preflight.go | 48 ++++++++ internal/hpagent/handler.go | 175 ++++++++------------------- internal/sshutil/connect.go | 201 ------------------------------- internal/sshutil/encoder.go | 125 ------------------- internal/sshutil/framebatcher.go | 61 ---------- internal/sshutil/handler.go | 152 ----------------------- internal/sshutil/sessions.go | 83 ------------- internal/tsnet/peers.go | 2 +- internal/tsnet/server.go | 8 +- mise.toml | 6 + 13 files changed, 160 insertions(+), 758 deletions(-) create mode 100644 internal/config/config.go create mode 100644 internal/config/preflight.go delete mode 100644 internal/sshutil/connect.go delete mode 100644 internal/sshutil/encoder.go delete mode 100644 internal/sshutil/framebatcher.go delete mode 100644 internal/sshutil/handler.go delete mode 100644 internal/sshutil/sessions.go diff --git a/cmd/hp_agent/hp_agent.go b/cmd/hp_agent/hp_agent.go index 03d90bb..b20c34d 100644 --- a/cmd/hp_agent/hp_agent.go +++ b/cmd/hp_agent/hp_agent.go @@ -2,11 +2,10 @@ package main import ( _ "github.com/joho/godotenv/autoload" - "github.com/tale/headplane/agent/internal/config" - "github.com/tale/headplane/agent/internal/hpagent" - "github.com/tale/headplane/agent/internal/sshutil" - "github.com/tale/headplane/agent/internal/tsnet" - "github.com/tale/headplane/agent/internal/util" + "github.com/tale/headplane/internal/config" + "github.com/tale/headplane/internal/hpagent" + "github.com/tale/headplane/internal/tsnet" + "github.com/tale/headplane/internal/util" ) type Register struct { @@ -32,6 +31,5 @@ func main() { ID: agent.ID, }) - sshutil.DispatchSSHStdin() hpagent.FollowMaster(agent) } diff --git a/go.mod b/go.mod index 098494d..b06e64f 100644 --- a/go.mod +++ b/go.mod @@ -5,7 +5,6 @@ go 1.24.0 toolchain go1.24.4 require ( - github.com/fxamacker/cbor/v2 v2.8.0 github.com/joho/godotenv v1.5.1 go4.org/mem v0.0.0-20240501181205-ae6ca9944745 golang.org/x/crypto v0.38.0 @@ -34,6 +33,7 @@ require ( github.com/coreos/go-iptables v0.7.1-0.20240112124308-65c67c9f46e6 // indirect github.com/dblohm7/wingoes v0.0.0-20240119213807-a09d6be7affa // indirect github.com/digitalocean/go-smbios v0.0.0-20180907143718-390a4f403a8e // indirect + github.com/fxamacker/cbor/v2 v2.8.0 // indirect github.com/gaissmai/bart v0.18.0 // indirect github.com/go-json-experiment/json v0.0.0-20250223041408-d3c622f1b874 // indirect github.com/go-ole/go-ole v1.3.0 // indirect diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..d33c014 --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,45 @@ +package config + +import "os" + +// Config represents the configuration for the agent. +type Config struct { + Debug bool + Hostname string + TSControlURL string + TSAuthKey string + WorkDir string +} + +const ( + DebugEnv = "HEADPLANE_AGENT_DEBUG" + HostnameEnv = "HEADPLANE_AGENT_HOSTNAME" + TSControlURLEnv = "HEADPLANE_AGENT_TS_SERVER" + TSAuthKeyEnv = "HEADPLANE_AGENT_TS_AUTHKEY" + WorkDirEnv = "HEADPLANE_AGENT_WORK_DIR" +) + +// Load reads the agent configuration from environment variables. +func Load() (*Config, error) { + c := &Config{ + Debug: false, + Hostname: os.Getenv(HostnameEnv), + TSControlURL: os.Getenv(TSControlURLEnv), + TSAuthKey: os.Getenv(TSAuthKeyEnv), + WorkDir: os.Getenv(WorkDirEnv), + } + + if os.Getenv(DebugEnv) == "true" { + c.Debug = true + } + + if err := validateRequired(c); err != nil { + return nil, err + } + + if err := validateTSReady(c); err != nil { + return nil, err + } + + return c, nil +} diff --git a/internal/config/preflight.go b/internal/config/preflight.go new file mode 100644 index 0000000..ee06d9f --- /dev/null +++ b/internal/config/preflight.go @@ -0,0 +1,48 @@ +package config + +import ( + "fmt" + "net/http" + "strings" +) + +// Checks to make sure all required environment variables are set +func validateRequired(config *Config) error { + if config.Hostname == "" { + return fmt.Errorf("%s is required", HostnameEnv) + } + + if config.TSControlURL == "" { + return fmt.Errorf("%s is required", TSControlURLEnv) + } + + if config.TSAuthKey == "" { + return fmt.Errorf("%s is required", TSAuthKeyEnv) + } + + if config.WorkDir == "" { + return fmt.Errorf("%s is required", WorkDirEnv) + } + + return nil +} + +// Pings the Tailscale control server to make sure it's up and running +func validateTSReady(config *Config) error { + testURL := config.TSControlURL + if strings.HasSuffix(testURL, "/") { + testURL = testURL[:len(testURL)-1] + } + + testURL = fmt.Sprintf("%s/key?v=116", testURL) + resp, err := http.Get(testURL) + if err != nil { + return fmt.Errorf("Failed to connect to TS control server: %s", err) + } + + if resp.StatusCode != 200 { + return fmt.Errorf("Failed to connect to TS control server: %s", resp.Status) + } + + return nil +} diff --git a/internal/hpagent/handler.go b/internal/hpagent/handler.go index 7bb48e7..826b74f 100644 --- a/internal/hpagent/handler.go +++ b/internal/hpagent/handler.go @@ -2,16 +2,14 @@ package hpagent import ( "bufio" - "bytes" - // "encoding/json" - "os" - // "sync" - "github.com/fxamacker/cbor/v2" - "github.com/tale/headplane/agent/internal/sshutil" - "github.com/tale/headplane/agent/internal/tsnet" - "github.com/tale/headplane/agent/internal/util" - // "tailscale.com/tailcfg" + "encoding/json" + "os" + "sync" + + "github.com/tale/headplane/internal/tsnet" + "github.com/tale/headplane/internal/util" + "tailscale.com/tailcfg" ) // Represents messages from the Headplane master @@ -19,11 +17,6 @@ type RecvMessage struct { NodeIDs []string } -type CborMessage struct { - Op string `cbor:"op"` - Payload cbor.RawMessage `cbor:"payload"` -} - type SendMessage struct { Type string Data any @@ -37,124 +30,58 @@ func FollowMaster(agent *tsnet.TSAgent) { for scanner.Scan() { line := scanner.Bytes() - - var msg CborMessage - decoder := cbor.NewDecoder(bytes.NewReader(line)) - err := decoder.Decode(&msg) - + var msg RecvMessage + err := json.Unmarshal(line, &msg) if err != nil { log.Error("Unable to decode message from master: %s", err) continue } - log.Debug("Received message from master: %s", msg) - switch msg.Op { - case "ssh_conn": - var sshPayload sshutil.SSHConnectPayload - err = cbor.Unmarshal(msg.Payload, &sshPayload) - if err != nil { - log.Error("Unable to unmarshal SSH connect payload: %s", err) - continue - } + log.Debug("Recieved message from master: %v", line) - sshutil.StartWebSSH(agent, sshPayload) - continue - - case "ssh_term": - var sshPayload sshutil.SSHClosePayload - err = cbor.Unmarshal(msg.Payload, &sshPayload) - if err != nil { - log.Error("Unable to unmarshal SSH close payload: %s", err) - continue - } - - sshutil.CloseWebSSH(agent, sshPayload) - continue - - case "ssh_resize": - var sshPayload sshutil.SSHResizePayload - err = cbor.Unmarshal(msg.Payload, &sshPayload) - if err != nil { - log.Error("Unable to unmarshal SSH resize payload: %s", err) - continue - } - - sshutil.ResizeWebSSH(agent, sshPayload) + if len(msg.NodeIDs) == 0 { + log.Debug("Message recieved had no node IDs") + log.Debug("Full message: %s", line) continue } + + // Accumulate the results since we invoke via gofunc + results := make(map[string]*tailcfg.HostinfoView) + mu := sync.Mutex{} + wg := sync.WaitGroup{} + + for _, nodeID := range msg.NodeIDs { + wg.Add(1) + go func(nodeID string) { + defer wg.Done() + result, err := agent.GetStatusForPeer(nodeID) + if err != nil { + log.Error("Unable to get status for node %s: %s", nodeID, err) + return + } + + if result == nil { + log.Debug("No status for node %s", nodeID) + return + } + + mu.Lock() + results[nodeID] = result + mu.Unlock() + }(nodeID) + } + + wg.Wait() + + // Send the results back to the Headplane master + log.Debug("Sending status back to master: %v", results) + log.Msg(&SendMessage{ + Type: "status", + Data: results, + }) } - // var msg RecvMessage - // err := json.Unmarshal(line, &msg) - // if err != nil { - // var cborMsg CborMessage - // dec := cbor.NewDecoder(bytes.NewReader(line)) - // err := dec.Decode(&cborMsg) - - // if err == nil { - // log.Info("Unmarshalled CBOR message: %s", cborMsg) - // var sshPayload SSHConnect - // err = cbor.Unmarshal(cborMsg.Payload, &sshPayload) - // sshutil.OpenSshPty(agent, sshutil.SshConnectParams{ - // Hostname: sshPayload.Hostname, - // Port: sshPayload.Port, - // Username: sshPayload.Username, - // Id: sshPayload.SessionId, - // }) - - // return; - // } - - // log.Error("Unable to unmarshal message: %s", err) - // log.Debug("Full Error: %v", err) - // continue - // } - - // log.Debug("Recieved message from master: %v", line) - - // if len(msg.NodeIDs) == 0 { - // log.Debug("Message recieved had no node IDs") - // log.Debug("Full message: %s", line) - // continue - // } - - // // Accumulate the results since we invoke via gofunc - // results := make(map[string]*tailcfg.HostinfoView) - // mu := sync.Mutex{} - // wg := sync.WaitGroup{} - - // for _, nodeID := range msg.NodeIDs { - // wg.Add(1) - // go func(nodeID string) { - // defer wg.Done() - // result, err := agent.GetStatusForPeer(nodeID) - // if err != nil { - // log.Error("Unable to get status for node %s: %s", nodeID, err) - // return - // } - - // if result == nil { - // log.Debug("No status for node %s", nodeID) - // return - // } - - // mu.Lock() - // results[nodeID] = result - // mu.Unlock() - // }(nodeID) - // } - - // wg.Wait() - - // // Send the results back to the Headplane master - // log.Debug("Sending status back to master: %v", results) - // log.Msg(&SendMessage{ - // Type: "status", - // Data: results, - // }) - // } - - // if err := scanner.Err(); err != nil { - // log.Fatal("Error reading from stdin: %s", err) - // } + if err := scanner.Err(); err != nil { + log.Fatal("Error reading from stdin: %s", err) + } } diff --git a/internal/sshutil/connect.go b/internal/sshutil/connect.go deleted file mode 100644 index c4e23cf..0000000 --- a/internal/sshutil/connect.go +++ /dev/null @@ -1,201 +0,0 @@ -package sshutil - -import ( - "context" - "errors" - "strconv" - "strings" - - "github.com/tale/headplane/agent/internal/tsnet" - "github.com/tale/headplane/agent/internal/util" - "golang.org/x/crypto/ssh" -) - -type SSHConnectPayload struct { - SessionId string `cbor:"sessionId"` - Username string `cbor:"username"` - Hostname string `cbor:"hostname"` - Port int `cbor:"port"` -} - -type SSHClosePayload struct { - SessionId string `cbor:"sessionId"` -} - -type SSHResizePayload struct { - SessionId string `cbor:"sessionId"` - Width int `cbor:"width"` - Height int `cbor:"height"` -} - -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("Initiating Tailscale SSH connection to %s@%s", params.Username, addr) - tailnetConn, err := agent.Dial(context.Background(), "tcp", addr) - if err != nil { - return nil, err - } - - 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(), - } - - conn, chans, reqs, err := ssh.NewClientConn(tailnetConn, addr, config) - if err != nil { - conn.Close() - return nil, err - } - - // At this point we have successfully connected to the node - sshClient := ssh.NewClient(conn, chans, reqs) - version := string(sshClient.ServerVersion()) - - if !strings.Contains(version, "Tailscale") { - conn.Close() - return nil, errors.New("server is not running Tailscale SSH") - } - - log.Info("Connected to %s@%s:%d via Tailscale SSH (%s)", params.Username, params.Hostname, params.Port, version) - return sshClient, nil -} - -func StartWebSSH(agent *tsnet.TSAgent, params SSHConnectPayload) { - log := util.GetLogger() - - if agent == nil { - log.Error("tsnet.TSAgent is not initialized correctly") - return - } - - if params.Hostname == "" || params.Port <= 0 || params.Username == "" || params.SessionId == "" { - log.Error("Invalid SSH connection parameters: %v", params) - return - } - - client, err := connectToTailscaleSSH(agent, params) - if err != nil { - 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() { - 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 - } - - 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", 24, 80, 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 - } - } - }() - - // 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 - } - - // This spawns 2 goroutins for stdout and stderr - dispatchSSHStdout(params.SessionId, ctx.Stdout, ctx.Stderr) - - log.Info("Opened an SSH PTY for %s", params.SessionId) - sess.Wait() - sess.Close() - client.Close() - - log.Info("SSH session for %s closed", params.SessionId) - RemoveSession(params.SessionId) - }() -} - -func CloseWebSSH(agent *tsnet.TSAgent, params SSHClosePayload) { - log := util.GetLogger() - - if agent == nil { - log.Error("tsnet.TSAgent is not initialized correctly") - return - } - - if params.SessionId == "" { - log.Error("Invalid SSH close parameters: %v", params) - return - } - - log.Debug("Closing SSH session for session ID: %s", params.SessionId) - ctx, ok := lookupSession(params.SessionId) - if !ok { - log.Info("No active SSH session found for session ID: %s", params.SessionId) - return - } - - RemoveSession(ctx.ID) - log.Info("SSH session for %s closed", params.SessionId) -} - -func ResizeWebSSH(agent *tsnet.TSAgent, params SSHResizePayload) { - log := util.GetLogger() - - if agent == nil { - log.Error("tsnet.TSAgent is not initialized correctly") - return - } - - if params.SessionId == "" || params.Width <= 0 || params.Height <= 0 { - log.Error("Invalid SSH resize parameters: %v", params) - return - } - - log.Debug("Resizing SSH session for session ID: %s to %dx%d", params.SessionId, params.Width, params.Height) - ctx, ok := lookupSession(params.SessionId) - if !ok { - log.Info("No active SSH session found for session ID: %s", params.SessionId) - return - } - - err := ctx.Session.WindowChange(params.Height, params.Width) - if err != nil { - log.Error("Failed to resize SSH session for (%s): %s", params.SessionId, err) - return - } - - log.Info("Resized SSH session for %s to %dx%d", params.SessionId, params.Width, params.Height) -} diff --git a/internal/sshutil/encoder.go b/internal/sshutil/encoder.go deleted file mode 100644 index f90e400..0000000 --- a/internal/sshutil/encoder.go +++ /dev/null @@ -1,125 +0,0 @@ -package sshutil - -import ( - "encoding/binary" - "fmt" -) - -// 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) - } - - if len(frame.SessionID) == 0 { - return nil, fmt.Errorf("session ID cannot be empty") - } - - 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 (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") - } - - frame.ChannelType = ChannelType(buf[5]) - if frame.ChannelType < ChannelTypeStdin || frame.ChannelType > ChannelTypeStderr { - return frame, fmt.Errorf("invalid channel type: %d", frame.ChannelType) - } - - 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/internal/sshutil/framebatcher.go b/internal/sshutil/framebatcher.go deleted file mode 100644 index 6653452..0000000 --- a/internal/sshutil/framebatcher.go +++ /dev/null @@ -1,61 +0,0 @@ -package sshutil - -import ( - "io" - "sync" - "time" - - "github.com/tale/headplane/agent/internal/util" -) - -type FrameBatcher struct { - mu sync.Mutex - buffer []byte - writer io.Writer - timer *time.Timer - interval time.Duration - done chan struct{} -} - -func NewFrameBatcher(writer io.Writer, interval time.Duration) *FrameBatcher { - return &FrameBatcher{ - writer: writer, - interval: interval, - done: make(chan struct{}), - } -} - -func (b *FrameBatcher) QueueMsg(msg []byte) { - b.mu.Lock() - defer b.mu.Unlock() - - b.buffer = append(b.buffer, msg...) - if b.timer != nil { - b.timer.Stop() - } - - b.timer = time.AfterFunc(b.interval, b.flush) -} - -func (b *FrameBatcher) flush() { - log := util.GetLogger() - - b.mu.Lock() - defer b.mu.Unlock() - - if len(b.buffer) == 0 { - return - } - - _, err := b.writer.Write(b.buffer) - if err != nil { - log.Error("Failed to write batched message: %v", err) - } - - b.buffer = nil -} - -func (b *FrameBatcher) Close() { - close(b.done) - b.flush() -} diff --git a/internal/sshutil/handler.go b/internal/sshutil/handler.go deleted file mode 100644 index 63b90c9..0000000 --- a/internal/sshutil/handler.go +++ /dev/null @@ -1,152 +0,0 @@ -package sshutil - -import ( - "io" - "os" - "time" - - "github.com/tale/headplane/agent/internal/util" -) - -// The file descriptors attached by the parent node process -const ( - InputFd = 3 - OutputFd = 4 -) - -// 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 { - // 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("Failed to read from SSH stdin: %v", err) - return - } - - // 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 { - // 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 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 - } - - batcher := NewFrameBatcher(fd, 10*time.Millisecond) // Roughly 60fps - - go readerStreamRoutine(StreamRoutine{ - SessionID: id, - Reader: stdout, - Writer: batcher, - ChannelType: ChannelTypeStdout, - }) - - go readerStreamRoutine(StreamRoutine{ - SessionID: id, - Reader: stderr, - Writer: batcher, - ChannelType: ChannelTypeStderr, - }) -} - -type StreamRoutine struct { - SessionID string - Reader io.Reader - Writer *FrameBatcher - ChannelType ChannelType -} - -func readerStreamRoutine(routine StreamRoutine) { - hpls1 := HPLSFrame1{} - buf := make([]byte, 16384) // 16 KiB buffer - for { - 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 - // } - // - - routine.Writer.QueueMsg(frame) - } -} diff --git a/internal/sshutil/sessions.go b/internal/sshutil/sessions.go deleted file mode 100644 index 3bfb04d..0000000 --- a/internal/sshutil/sessions.go +++ /dev/null @@ -1,83 +0,0 @@ -package sshutil - -import ( - "errors" - "io" - "sync" - - "github.com/tale/headplane/agent/internal/util" - "golang.org/x/crypto/ssh" -) - -type SessionContext struct { - ID string - Session *ssh.Session - Stdin io.WriteCloser - Stdout io.Reader - Stderr io.Reader - InputCh chan []byte -} - -var sessions = make(map[string]*SessionContext) -var sessionsLock sync.RWMutex - -func registerSessionChans(id string, session *ssh.Session) (*SessionContext, error) { - log := util.GetLogger() - - sessionsLock.Lock() - defer sessionsLock.Unlock() - - if _, exists := sessions[id]; exists { - return sessions[id], nil - } - - stdin, err := session.StdinPipe() - if err != nil { - return nil, errors.New("failed to create stdin pipe: " + err.Error()) - } - - stdout, err := session.StdoutPipe() - if err != nil { - stdin.Close() - return nil, errors.New("failed to create stdout pipe: " + err.Error()) - } - - 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] = ctx - log.Debug("Registered session %s with stdin, stdout, and stderr pipes", id) - return ctx, nil -} - -func lookupSession(id string) (*SessionContext, bool) { - sessionsLock.RLock() - defer sessionsLock.RUnlock() - - sessionContext, exists := sessions[id] - return sessionContext, exists -} - -func RemoveSession(id string) { - sessionsLock.Lock() - defer sessionsLock.Unlock() - - if sessionContext, exists := sessions[id]; exists { - sessionContext.Stdin.Close() // Close the stdin pipe - sessionContext.Session.Close() // Close the SSH session - delete(sessions, id) // Remove from the map - } -} diff --git a/internal/tsnet/peers.go b/internal/tsnet/peers.go index 9628dca..1ce2c2b 100644 --- a/internal/tsnet/peers.go +++ b/internal/tsnet/peers.go @@ -6,7 +6,7 @@ import ( "fmt" "strings" - "github.com/tale/headplane/agent/internal/util" + "github.com/tale/headplane/internal/util" "tailscale.com/tailcfg" "tailscale.com/types/key" diff --git a/internal/tsnet/server.go b/internal/tsnet/server.go index eb4ffe8..2909911 100644 --- a/internal/tsnet/server.go +++ b/internal/tsnet/server.go @@ -5,16 +5,16 @@ import ( "os" "path/filepath" - "github.com/tale/headplane/agent/internal/config" - "github.com/tale/headplane/agent/internal/util" - "tailscale.com/client/tailscale" + "github.com/tale/headplane/internal/config" + "github.com/tale/headplane/internal/util" + "tailscale.com/client/local" "tailscale.com/tsnet" ) // Wrapper type so we can add methods to the server. type TSAgent struct { *tsnet.Server - Lc *tailscale.LocalClient + Lc *local.Client ID string } diff --git a/mise.toml b/mise.toml index bc3e84b..cc54c04 100644 --- a/mise.toml +++ b/mise.toml @@ -17,6 +17,11 @@ description = "Builds the Go WebAssembly module for Tailscale SSH" env = { GOOS = "js", GOARCH = "wasm" } run = "go build -o app/hp_ssh.wasm ./cmd/hp_ssh" +[tasks.build-go-agent] +alias = ["agent"] +description = "Builds the Go agent for HostInfo" +run = "go build -o build/hp_agent ./cmd/hp_agent" + [tasks.generate-caddy-certs] alias = ["mkcert"] dir = "{{cwd}}/test/caddy/certs" @@ -29,6 +34,7 @@ run = [ [tasks.ci] description = "Runs the CI pipeline" depends = ["build-go-wasm"] +depends_post = ["build-go-agent"] run = [ "pnpm install", "pnpm run build"