support handle ping@openssh packet (#582)

* chore: Update dependencies in go.mod and go.sum

* chore: Update submodule commit for crypto

* refactor: Update hook functions to return ssh.PipePackageHookMethod for better integration

* fix: Increase sleep duration in TestFixed to ensure proper execution

* fix: Increase sleep duration in TestFixed to ensure proper file flush

* fix: Update docker-compose service images and modify test to use host-password-old

* refactor: Simplify hook functions to return error instead of PipePackageHookMethod for better error handling

* feat: Add reply-ping flag to enable ping response for compatibility with older sshd

* fix: Update subproject commit reference in crypto
This commit is contained in:
Boshi Lian 2025-05-08 11:53:40 -07:00 committed by GitHub
parent c93e3c2f4a
commit 2603cdfa8b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 196 additions and 48 deletions

View file

@ -59,7 +59,7 @@ func newAsciicastLogger(recorddir string, prefix string) *asciicastLogger {
}
}
func (l *asciicastLogger) uphook(msg []byte) ([]byte, error) {
func (l *asciicastLogger) uphook(msg []byte) error {
if msg[0] == msgChannelData {
clientChannelID := binary.BigEndian.Uint32(msg[1:5])
@ -71,7 +71,7 @@ func (l *asciicastLogger) uphook(msg []byte) ([]byte, error) {
_, err := fmt.Fprintf(f, "[%v,\"o\",\"%s\"]\n", t, jsonEscape(string(buf)))
if err != nil {
return msg, err
return err
}
}
} else if msg[0] == msgChannelOpenConfirm {
@ -79,10 +79,10 @@ func (l *asciicastLogger) uphook(msg []byte) ([]byte, error) {
serverChannelID := binary.BigEndian.Uint32(msg[5:9])
l.channelIDMap[serverChannelID] = clientChannelID
}
return msg, nil
return nil
}
func (l *asciicastLogger) downhook(msg []byte) ([]byte, error) {
func (l *asciicastLogger) downhook(msg []byte) error {
if msg[0] == msgChannelRequest {
t := time.Since(l.starttime).Seconds()
serverChannelID := binary.BigEndian.Uint32(msg[1:5])
@ -112,14 +112,14 @@ func (l *asciicastLogger) downhook(msg []byte) ([]byte, error) {
_, err := fmt.Fprintf(f, "[%v,\"r\", \"%vx%v\"]\n", t, width, height)
if err != nil {
return msg, err
return err
}
}
case "shell", "exec":
jsonEnvs, err := json.Marshal(l.envs)
if err != nil {
return msg, err
return err
}
f, err := os.OpenFile(
@ -129,7 +129,7 @@ func (l *asciicastLogger) downhook(msg []byte) ([]byte, error) {
)
if err != nil {
return msg, err
return err
}
l.channels[clientChannelID] = f
@ -146,11 +146,11 @@ func (l *asciicastLogger) downhook(msg []byte) ([]byte, error) {
)
if err != nil {
return msg, err
return err
}
}
}
return msg, nil
return nil
}
func (l *asciicastLogger) Close() (err error) {

View file

@ -27,6 +27,7 @@ type daemon struct {
recordfmt string
usernameAsRecorddir bool
filterHostkeysReqeust bool
replyPing bool
}
func generateSshKey(keyfile string) error {
@ -227,8 +228,8 @@ func (d *daemon) run() error {
log.Infof("ssh connection pipe created %v (username [%v]) -> %v (username [%v])", p.DownstreamConnMeta().RemoteAddr(), p.DownstreamConnMeta().User(), p.UpstreamConnMeta().RemoteAddr(), p.UpstreamConnMeta().User())
var uphook func([]byte) ([]byte, error)
var downhook func([]byte) ([]byte, error)
uphookchain := &hookChain{}
downhookchain := &hookChain{}
if d.recorddir != "" {
var recorddir string
@ -252,8 +253,8 @@ func (d *daemon) run() error {
recorder := newAsciicastLogger(recorddir, prefix)
defer recorder.Close()
uphook = recorder.uphook
downhook = recorder.downhook
uphookchain.append(ssh.InspectPacketHook(recorder.uphook))
downhookchain.append(ssh.InspectPacketHook(recorder.downhook))
} else if d.recordfmt == "typescript" {
recorder, err := newFilePtyLogger(recorddir)
if err != nil {
@ -262,34 +263,35 @@ func (d *daemon) run() error {
}
defer recorder.Close()
uphook = recorder.loggingTty
uphookchain.append(ssh.InspectPacketHook(recorder.loggingTty))
}
}
if d.filterHostkeysReqeust {
nextUpHook := uphook
uphook = func(b []byte) ([]byte, error) {
uphookchain.append(func(b []byte) (ssh.PipePacketHookMethod, []byte, error) {
if b[0] == 80 {
var x struct {
RequestName string `sshtype:"80"`
}
_ = ssh.Unmarshal(b, &x)
if x.RequestName == "hostkeys-prove-00@openssh.com" || x.RequestName == "hostkeys-00@openssh.com" {
return nil, nil
return ssh.PipePacketHookTransform, nil, nil
}
}
if nextUpHook != nil {
return nextUpHook(b)
}
return b, nil
}
return ssh.PipePacketHookTransform, b, nil
})
}
if d.replyPing {
downhookchain.append(ssh.PingPacketReply)
}
if d.config.PipeStartCallback != nil {
d.config.PipeStartCallback(p.DownstreamConnMeta(), p.ChallengeContext())
}
err = p.WaitWithHook(uphook, downhook)
err = p.WaitWithHook(uphookchain.hook(), downhookchain.hook())
if d.config.PipeErrorCallback != nil {
d.config.PipeErrorCallback(p.DownstreamConnMeta(), p.ChallengeContext(), err)

37
cmd/sshpiperd/hook.go Normal file
View file

@ -0,0 +1,37 @@
package main
import "golang.org/x/crypto/ssh"
type hookChain struct {
hooks []ssh.PipePacketHook
}
// chain stops if any of the hooks return ssh.PipePacketHookReply
func (h *hookChain) append(hook ssh.PipePacketHook) {
if hook != nil {
h.hooks = append(h.hooks, hook)
}
}
func (h *hookChain) hook() ssh.PipePacketHook {
if len(h.hooks) == 0 {
return nil
}
return func(packet []byte) (method ssh.PipePacketHookMethod, packetOut []byte, err error) {
packetOut = packet
for _, hk := range h.hooks {
method, packetOut, err = hk(packetOut)
if err != nil {
return
}
if method == ssh.PipePacketHookReply {
return
}
}
return
}
}

View file

@ -0,0 +1,91 @@
package main
import (
"errors"
"testing"
"golang.org/x/crypto/ssh"
)
func TestHookChain_Hook(t *testing.T) {
hc := &hookChain{}
// Mock hooks
hook1 := func(packet []byte) (ssh.PipePacketHookMethod, []byte, error) {
return ssh.PipePacketHookTransform, append(packet, '1'), nil
}
hook2 := func(packet []byte) (ssh.PipePacketHookMethod, []byte, error) {
return ssh.PipePacketHookTransform, append(packet, '2'), nil
}
hook3 := func(packet []byte) (ssh.PipePacketHookMethod, []byte, error) {
return ssh.PipePacketHookTransform, append(packet, '3'), nil
}
hook4 := func(packet []byte) (ssh.PipePacketHookMethod, []byte, error) {
return ssh.PipePacketHookReply, append(packet, '4'), nil
}
hook5 := func(packet []byte) (ssh.PipePacketHookMethod, []byte, error) {
return ssh.PipePacketHookTransform, append(packet, '5'), nil
}
hc.append(hook1)
hc.append(hook2)
hc.append(hook3)
hc.append(hook4)
hc.append(hook5)
finalHook := hc.hook()
if finalHook == nil {
t.Fatal("expected a non-nil hook")
}
packet := []byte("test")
method, packetOut, err := finalHook(packet)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if method != ssh.PipePacketHookReply {
t.Errorf("expected method to be PipePacketHookReply, got %v", method)
}
expectedPacket := "test1234"
if string(packetOut) != expectedPacket {
t.Errorf("expected packetOut to be %q, got %q", expectedPacket, string(packetOut))
}
}
func TestHookChain_HookWithError(t *testing.T) {
hc := &hookChain{}
// Mock hooks
hook1 := func(packet []byte) (ssh.PipePacketHookMethod, []byte, error) {
return ssh.PipePacketHookTransform, append(packet, '1'), nil
}
hookWithError := func(packet []byte) (ssh.PipePacketHookMethod, []byte, error) {
return ssh.PipePacketHookTransform, nil, errors.New("mock error")
}
hc.append(hook1)
hc.append(hookWithError)
finalHook := hc.hook()
if finalHook == nil {
t.Fatal("expected a non-nil hook")
}
packet := []byte("test")
_, _, err := finalHook(packet)
if err == nil {
t.Fatal("expected an error, got nil")
}
expectedError := "mock error"
if err.Error() != expectedError {
t.Errorf("expected error %q, got %q", expectedError, err.Error())
}
}

View file

@ -172,6 +172,12 @@ func main() {
Usage: "filter out hostkeys-00@openssh.com which cause client side warnings",
EnvVars: []string{"SSHPIPERD_DROP_HOSTKEYS_MESSAGE"},
},
&cli.BoolFlag{
Name: "reply-ping",
Value: true,
Usage: "reply to ping@openssh instead of passing it to upstream, this is useful for old sshd which doesn't support ping@openssh",
EnvVars: []string{"SSHPIPERD_REPLY_PING"},
},
&cli.StringSliceFlag{
Name: "allowed-proxy-addresses",
Value: cli.NewStringSlice(),
@ -297,6 +303,7 @@ func main() {
d.recordfmt = ctx.String("screen-recording-format")
d.usernameAsRecorddir = ctx.Bool("username-as-recorddir")
d.filterHostkeysReqeust = ctx.Bool("drop-hostkeys-message")
d.replyPing = ctx.Bool("reply-ping")
if d.recordfmt != "typescript" && d.recordfmt != "asciicast" {
return fmt.Errorf("invalid screen recording format: %v", d.recordfmt)

View file

@ -49,7 +49,7 @@ func newFilePtyLogger(outputdir string) (*filePtyLogger, error) {
}, nil
}
func (l *filePtyLogger) loggingTty(msg []byte) ([]byte, error) {
func (l *filePtyLogger) loggingTty(msg []byte) error {
if msg[0] == msgChannelData {
@ -67,12 +67,12 @@ func (l *filePtyLogger) loggingTty(msg []byte) ([]byte, error) {
_, err := l.typescript.Write(buf)
if err != nil {
return msg, err
return err
}
}
return msg, nil
return nil
}
func (l *filePtyLogger) Close() (err error) {

2
crypto

@ -1 +1 @@
Subproject commit 8682cc0c5ad6edd98b8992b8b1a932738b3d1bb3
Subproject commit 2a4b9c2448bc0257714a950e5d61a55d3633eb65

View file

@ -1,8 +1,6 @@
version: '3.4'
services:
host-password:
image: lscr.io/linuxserver/openssh-server:9.9_p1-r2-ls190
image: linuxserver/openssh-server:9.9_p1-r2-ls190
environment:
- PASSWORD_ACCESS=true
- USER_PASSWORD=pass
@ -22,8 +20,21 @@ services:
- default
- netdistract
host-password-old:
image: linuxserver/openssh-server:8.1_p1-r0-ls19
environment:
- PASSWORD_ACCESS=true
- USER_PASSWORD=pass
- USER_NAME=user
- LOG_STDOUT=true
volumes:
- shared:/shared
networks:
- default
- netdistract
host-publickey:
image: lscr.io/linuxserver/openssh-server:9.9_p1-r2-ls190
image: linuxserver/openssh-server:9.9_p1-r2-ls190
environment:
- USER_NAME=user
- LOG_STDOUT=true
@ -38,7 +49,7 @@ services:
- ./sshdconfig/no_penalties.conf:/config/sshd/sshd_config.d/no_penalties.conf:ro
host-capublickey:
image: lscr.io/linuxserver/openssh-server:9.9_p1-r2-ls190
image: linuxserver/openssh-server:9.9_p1-r2-ls190
environment:
- USER_NAME=ca_user
- LOG_STDOUT=true

View file

@ -10,7 +10,7 @@ import (
"github.com/google/uuid"
)
func TestFixed(t *testing.T) {
func TestOldSshd(t *testing.T) {
piperaddr, piperport := nextAvailablePiperAddress()
@ -19,7 +19,7 @@ func TestFixed(t *testing.T) {
piperport,
"/sshpiperd/plugins/fixed",
"--target",
"host-password:2222",
"host-password-old:2222",
)
if err != nil {
@ -61,7 +61,7 @@ func TestFixed(t *testing.T) {
"-l",
"user",
"127.0.0.1",
fmt.Sprintf(`sh -c "echo SSHREADY && sleep 1 && echo -n %v > /shared/%v"`, randtext, targetfie), // sleep 1 to cover https://github.com/tg123/sshpiper/issues/323
fmt.Sprintf(`sh -c "echo SSHREADY && sleep 1 && echo -n %v > /shared/%v"`, randtext, targetfie), // sleep 5 to cover https://github.com/tg123/sshpiper/issues/323
)
if err != nil {

10
go.mod
View file

@ -16,7 +16,7 @@ require (
github.com/tg123/remotesigner v0.0.3
github.com/urfave/cli/v2 v2.27.6
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/crypto v0.37.0
golang.org/x/crypto v0.38.0
google.golang.org/grpc v1.72.0
google.golang.org/protobuf v1.36.6
gopkg.in/yaml.v3 v3.0.1
@ -44,7 +44,7 @@ require (
go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracehttp v1.31.0 // indirect
go.opentelemetry.io/otel/metric v1.34.0 // indirect
go.opentelemetry.io/otel/trace v1.34.0 // indirect
golang.org/x/sync v0.12.0 // indirect
golang.org/x/sync v0.14.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20250218202821-56aae31c358a // indirect
gopkg.in/evanphx/json-patch.v4 v4.12.0 // indirect
k8s.io/gengo/v2 v2.0.0-20240911193312-2b36238f13e9 // indirect
@ -83,9 +83,9 @@ require (
golang.org/x/mod v0.21.0 // indirect
golang.org/x/net v0.35.0 // indirect
golang.org/x/oauth2 v0.26.0 // indirect
golang.org/x/sys v0.31.0 // indirect
golang.org/x/term v0.30.0
golang.org/x/text v0.23.0 // indirect
golang.org/x/sys v0.33.0 // indirect
golang.org/x/term v0.32.0
golang.org/x/text v0.25.0 // indirect
golang.org/x/time v0.7.0 // indirect
golang.org/x/tools v0.26.0 // indirect
gopkg.in/inf.v0 v0.9.1 // indirect

16
go.sum
View file

@ -210,8 +210,8 @@ golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.3.0/go.mod h1:FU7BRWz2tNW+3quACPkgCx/L+uEAv1htQ0V83Z9Rj+Y=
golang.org/x/sync v0.6.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.12.0 h1:MHc5BpPuC30uJk597Ri8TV3CNZcTLu6B6z4lJy+g6Jw=
golang.org/x/sync v0.12.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
golang.org/x/sync v0.14.0 h1:woo0S4Yywslg6hp4eUFjTVOyKt0RookbpAHG4c1HmhQ=
golang.org/x/sync v0.14.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20191026070338-33540a1f6037/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@ -227,16 +227,16 @@ golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.17.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/sys v0.31.0 h1:ioabZlmFYtWhL+TRYpcnNlLwhyxaM9kWTDEmfnprqik=
golang.org/x/sys v0.31.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/sys v0.33.0 h1:q3i8TbbEz+JRD9ywIRlyRAQbM0qF7hu24q3teo2hbuw=
golang.org/x/sys v0.33.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
golang.org/x/term v0.5.0/go.mod h1:jMB1sMXY+tzblOD4FWmEbocvup2/aLOaQEp7JmGp78k=
golang.org/x/term v0.12.0/go.mod h1:owVbMEjm3cBLCHdkQu9b1opXd4ETQWc3BhuQGKgXgvU=
golang.org/x/term v0.17.0/go.mod h1:lLRBjIVuehSbZlaOtGMbcMncT+aqLLLmKrsjNrUguwk=
golang.org/x/term v0.20.0/go.mod h1:8UkIAJTvZgivsXaD6/pH6U9ecQzZ45awqEOzuCvwpFY=
golang.org/x/term v0.30.0 h1:PQ39fJZ+mfadBm0y5WlL4vlM7Sx1Hgf13sMIY2+QS9Y=
golang.org/x/term v0.30.0/go.mod h1:NYYFdzHoI5wRh/h5tDMdMqCqPJZEuNqVR5xJLd/n67g=
golang.org/x/term v0.32.0 h1:DR4lr0TjUs3epypdhTOkMmuF5CDFJ/8pOnbzMZPQ7bg=
golang.org/x/term v0.32.0/go.mod h1:uZG1FhGx848Sqfsq4/DlJr3xGGsYMu/L5GW4abiaEPQ=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
@ -244,8 +244,8 @@ golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8=
golang.org/x/text v0.13.0/go.mod h1:TvPlkZtksWOMsz7fbANvkp4WM8x/WCo/om8BMLbz+aE=
golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.15.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU=
golang.org/x/text v0.23.0 h1:D71I7dUrlY+VX0gQShAThNGHFxZ13dGLBHQLVl1mJlY=
golang.org/x/text v0.23.0/go.mod h1:/BLNzu4aZCJ1+kcD0DNRotWKage4q2rGVAg4o22unh4=
golang.org/x/text v0.25.0 h1:qVyWApTSYLk/drJRO5mDlNYskwQznZmkpV2c8q9zls4=
golang.org/x/text v0.25.0/go.mod h1:WEdwpYrmk1qmdHvhkSTNPm3app7v4rsT8F2UD6+VHIA=
golang.org/x/time v0.7.0 h1:ntUhktv3OPE6TgYxXWv9vKvUSJyIFJlyohwbkEwPrKQ=
golang.org/x/time v0.7.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=