mirror of
https://github.com/tale/headplane.git
synced 2026-07-26 07:48:14 +00:00
fix: cleanup agent and add to ci
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+51
-124
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user