diff --git a/commands/query_context.go b/commands/query_context.go index 70c21ad..99f3395 100644 --- a/commands/query_context.go +++ b/commands/query_context.go @@ -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), diff --git a/commands/track.go b/commands/track.go index 124a0ab..69a1a56 100644 --- a/commands/track.go +++ b/commands/track.go @@ -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, @@ -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 diff --git a/model/api.go b/model/api.go index bd8bd55..8e0d595 100644 --- a/model/api.go +++ b/model/api.go @@ -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 { diff --git a/model/command.go b/model/command.go index 8340e79..80deca6 100644 --- a/model/command.go +++ b/model/command.go @@ -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:"-"` diff --git a/model/sys.go b/model/sys.go index 493ce12..16f690a 100644 --- a/model/sys.go +++ b/model/sys.go @@ -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") != "" +} diff --git a/model/sys_ssh_test.go b/model/sys_ssh_test.go new file mode 100644 index 0000000..eeff01c --- /dev/null +++ b/model/sys_ssh_test.go @@ -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] })) + }) + } +} diff --git a/model/tracking_build.go b/model/tracking_build.go index 204f883..4847466 100644 --- a/model/tracking_build.go +++ b/model/tracking_build.go @@ -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 { @@ -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) diff --git a/model/tracking_build_test.go b/model/tracking_build_test.go index 7e12c00..f4d3d34 100644 --- a/model/tracking_build_test.go +++ b/model/tracking_build_test.go @@ -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