Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion commands/query_context.go
Original file line number Diff line number Diff line change
Expand Up @@ -129,7 +129,7 @@ func collectSystemContext(getenv func(string) string) *model.QuerySystemContext
UptimeSec: st.UptimeSec,
LoadAvg: st.LoadAvg,
IsRoot: os.Geteuid() == 0,
SSH: getenv("SSH_CONNECTION") != "" || getenv("SSH_TTY") != "",
SSH: model.IsSSHSession(getenv),
Container: st.Container,
Multiplexer: detectMultiplexer(getenv),
TermProgram: model.SanitizeContextString(getenv("TERM_PROGRAM"), 64),
Expand Down
2 changes: 2 additions & 0 deletions commands/track.go
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ func commandTrack(c *cli.Context) error {
cmdPhase := c.String("phase")
result := c.Int("result")
ppid := c.Int("ppid")
viaSSH := model.IsSSHSession(os.Getenv)

instance := &model.Command{
Shell: shell,
Expand All @@ -88,6 +89,7 @@ func commandTrack(c *cli.Context) error {
Time: time.Now(),
Phase: model.CommandPhasePre,
PPID: ppid,
ViaSSH: &viaSSH,
}

// Fast path: `track` runs inside the shell hook on every command, so it must
Expand Down
1 change: 1 addition & 0 deletions model/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ type TrackingData struct {
EndTimeNano int64 `json:"endTimeNano"`
Result int `json:"result"`
PPID int `json:"ppid,omitempty"`
ViaSSH *bool `json:"viaSsh,omitempty"`
}

type TrackingMetaData struct {
Expand Down
3 changes: 3 additions & 0 deletions model/command.go
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,9 @@ type Command struct {
Result int `json:"result"`
Phase CommandPhase `json:"phase"`
PPID int `json:"ppid,omitempty"`
// ViaSSH is set when the shell is an SSH login. Nil on records written
// before it was captured, which the server stores as unknown.
ViaSSH *bool `json:"ssh,omitempty"`

// Only work in file
RecordingTime time.Time `json:"-"`
Expand Down
7 changes: 7 additions & 0 deletions model/sys.go
Original file line number Diff line number Diff line change
Expand Up @@ -62,3 +62,10 @@ func GetOSAndVersion() (*SysInfo, error) {
Version: "unknown",
}, nil
}

// IsSSHSession reports whether the current shell is an SSH login. sshd exports
// these variables into the login shell, and every child (like `shelltime track`)
// inherits them.
func IsSSHSession(getenv func(string) string) bool {
return getenv("SSH_CONNECTION") != "" || getenv("SSH_CLIENT") != "" || getenv("SSH_TTY") != ""
}
25 changes: 25 additions & 0 deletions model/sys_ssh_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
package model

import (
"testing"

"github.com/stretchr/testify/require"
)

func TestIsSSHSession(t *testing.T) {
tests := []struct {
name string
env map[string]string
want bool
}{
{name: "local shell", env: map[string]string{"TERM_PROGRAM": "ghostty"}, want: false},
{name: "ssh connection", env: map[string]string{"SSH_CONNECTION": "10.0.0.2 51234 10.0.0.9 22"}, want: true},
{name: "ssh client only", env: map[string]string{"SSH_CLIENT": "10.0.0.2 51234 22"}, want: true},
{name: "ssh tty only", env: map[string]string{"SSH_TTY": "/dev/pts/1"}, want: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.want, IsSSHSession(func(k string) string { return tt.env[k] }))
})
}
}
4 changes: 4 additions & 0 deletions model/tracking_build.go
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,7 @@ func BuildTrackingData(ctx context.Context, store CommandStore, config ShellTime
EndTimeNano: postCommand.Time.UnixNano(),
Result: postCommand.Result,
PPID: postCommand.PPID,
ViaSSH: postCommand.ViaSSH,
}

if config.DataMasking != nil && *config.DataMasking {
Expand All @@ -113,6 +114,9 @@ func BuildTrackingData(ctx context.Context, store CommandStore, config ShellTime
if closestPreCommand != nil {
td.StartTime = closestPreCommand.Time.Unix()
td.StartTimeNano = closestPreCommand.Time.UnixNano()
if pre := closestPreCommand.ViaSSH; pre != nil && (td.ViaSSH == nil || *pre) {
td.ViaSSH = pre
}
}

trackingData = append(trackingData, td)
Expand Down
37 changes: 37 additions & 0 deletions model/tracking_build_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,43 @@ func TestBuildTrackingData(t *testing.T) {
require.Equal(t, StorageEngineBolt, res.Meta.CliEngine)
}

func TestBuildTrackingDataViaSSH(t *testing.T) {
yes, no := true, false
tests := []struct {
name string
pre *bool
post *bool
want *bool
}{
{name: "unknown on records from older CLIs", want: nil},
{name: "taken from post", pre: &no, post: &yes, want: &yes},
{name: "falls back to pre", pre: &yes, want: &yes},
{name: "true on either side wins", pre: &yes, post: &no, want: &yes},
{name: "local shell", pre: &no, post: &no, want: &no},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
store, err := newBoltStore(filepath.Join(t.TempDir(), "commands.db"))
require.NoError(t, err)
defer store.Close()

ctx := context.Background()
start := time.Now()
pre := Command{Shell: "zsh", SessionID: 7, Command: "make", Username: "u", Hostname: "devbox", Time: start, ViaSSH: tt.pre}
require.NoError(t, store.SavePre(ctx, pre, start))
post := pre
post.Time = start.Add(time.Second)
post.ViaSSH = tt.post
require.NoError(t, store.SavePost(ctx, post, 0, post.Time))

res, err := BuildTrackingData(ctx, store, ShellTimeConfig{})
require.NoError(t, err)
require.Len(t, res.Data, 1)
require.Equal(t, tt.want, res.Data[0].ViaSSH)
})
}
}

func TestBuildTrackingDataFileEngine(t *testing.T) {
t.Setenv("HOME", t.TempDir())
InitFolder("") // reset globals to the default .shelltime under the temp HOME
Expand Down
Loading