fix: cleanup agent and add to ci

This commit is contained in:
Aarnav Tale
2025-06-16 11:38:13 -04:00
parent ab7587e7a9
commit 4961e83a9c
13 changed files with 160 additions and 758 deletions
+4 -6
View File
@@ -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)
}
+1 -1
View File
@@ -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
+45
View File
@@ -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
}
+48
View File
@@ -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
View File
@@ -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)
}
}
-201
View File
@@ -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)
}
-125
View File
@@ -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
}
-61
View File
@@ -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()
}
-152
View File
@@ -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)
}
}
-83
View File
@@ -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
}
}
+1 -1
View File
@@ -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"
+4 -4
View File
@@ -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
}
+6
View File
@@ -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"