diff --git a/.travis.yml b/.travis.yml index 3b050088..e3a4dfd0 100644 --- a/.travis.yml +++ b/.travis.yml @@ -6,3 +6,4 @@ go: script: - go test ./ssh + - go test ./sshpiperd diff --git a/ssh/sshpiper_test.go b/ssh/sshpiper_test.go index fb68b6e8..583ee2ab 100644 --- a/ssh/sshpiper_test.go +++ b/ssh/sshpiper_test.go @@ -1,3 +1,8 @@ +// Copyright 2014, 2015 tgic. All rights reserved. +// this file is governed by MIT-license +// +// https://github.com/tg123/sshpiper + package ssh import ( diff --git a/sshpiperd/sshpiperd.go b/sshpiperd/sshpiperd.go index 721a0854..0a5da2fe 100644 --- a/sshpiperd/sshpiperd.go +++ b/sshpiperd/sshpiperd.go @@ -1,7 +1,11 @@ +// Copyright 2014, 2015 tgic. All rights reserved. +// this file is governed by MIT-license +// +// https://github.com/tg123/sshpiper + package main import ( - "bytes" "flag" "fmt" "github.com/tg123/sshpiper/ssh" @@ -10,15 +14,6 @@ import ( "log" "net" "os" - "strings" -) - -type userFile string - -var ( - UserAuthorizedKeysFile userFile = "authorized_keys" - UserKeyFile userFile = "id_rsa" - UserUpstreamFile userFile = "sshpiper_upstream" ) var ( @@ -42,125 +37,6 @@ func init() { flag.Parse() } -func userSpecFile(user, file string) string { - return fmt.Sprintf("%s/%s/%s", WorkingDir, user, file) -} - -func (file userFile) read(user string) ([]byte, error) { - return ioutil.ReadFile(userSpecFile(user, string(file))) -} - -func (file userFile) realPath(user string) string { - return userSpecFile(user, string(file)) -} - -// return error if not 400, nil if 400 and no err occurs -func (file userFile) check400(user string) error { - filename := userSpecFile(user, string(file)) - f, err := os.Open(filename) - if err != nil { - return err - } - defer f.Close() - - fi, err := f.Stat() - if err != nil { - return err - } - - if fi.Mode().Perm() != 0400 { - return fmt.Errorf("%v's perm is too open, change it to 400", filename) - } - - return nil -} - -func findUpstreamFromUserfile(conn ssh.ConnMetadata) (net.Conn, error) { - user := conn.User() - - err := UserUpstreamFile.check400(user) - if err != nil { - return nil, err - } - - addr, err := UserUpstreamFile.read(user) - if err != nil { - return nil, err - } - - saddr := strings.TrimSpace(string(addr)) - - logger.Printf("mapping user [%s] to [%s]", user, saddr) - - c, err := net.Dial("tcp", saddr) - if err != nil { - return nil, err - } - - return c, nil -} - -func mapPublicKeyFromUserfile(conn ssh.ConnMetadata, key ssh.PublicKey) (ssh.Signer, error) { - user := conn.User() - - var err error - defer func() { // print error when func exit - if err != nil { - logger.Printf("mapping private key error: %v, public key auth denied for [%v] from [%v]", err, user, conn.RemoteAddr()) - } - }() - - err = UserAuthorizedKeysFile.check400(user) - if err != nil { - return nil, err - } - - keydata := key.Marshal() - - var rest []byte - rest, err = UserAuthorizedKeysFile.read(user) - if err != nil { - return nil, err - } - - var authedPubkey ssh.PublicKey - - for len(rest) > 0 { - authedPubkey, _, _, rest, err = ssh.ParseAuthorizedKey(rest) - - if err != nil { - return nil, err - } - - if bytes.Equal(authedPubkey.Marshal(), keydata) { - err = UserKeyFile.check400(user) - if err != nil { - return nil, err - } - - var privateBytes []byte - privateBytes, err = UserKeyFile.read(user) - if err != nil { - return nil, err - } - - var private ssh.Signer - private, err = ssh.ParsePrivateKey(privateBytes) - if err != nil { - return nil, err - } - - // in log may see this twice, one is for query the other is real sign again - logger.Printf("auth succ, using mapped private key [%v] for user [%v] from [%v]", UserKeyFile.realPath(user), user, conn.RemoteAddr()) - return private, nil - } - } - - logger.Printf("public key auth failed user [%v] from [%v]", conn.User(), conn.RemoteAddr()) - - return nil, nil -} - func main() { if ShowHelp { @@ -168,6 +44,7 @@ func main() { return } + // TODO make this pluggable piper := &ssh.SSHPiperConfig{ FindUpstream: findUpstreamFromUserfile, MapPublicKey: mapPublicKeyFromUserfile, diff --git a/sshpiperd/workingdir.go b/sshpiperd/workingdir.go new file mode 100644 index 00000000..22ae0a99 --- /dev/null +++ b/sshpiperd/workingdir.go @@ -0,0 +1,143 @@ +// Copyright 2014, 2015 tgic. All rights reserved. +// this file is governed by MIT-license +// +// https://github.com/tg123/sshpiper + +package main + +import ( + "bytes" + "fmt" + "github.com/tg123/sshpiper/ssh" + "io/ioutil" + "net" + "os" + "strings" +) + +type userFile string + +var ( + UserAuthorizedKeysFile userFile = "authorized_keys" + UserKeyFile userFile = "id_rsa" + UserUpstreamFile userFile = "sshpiper_upstream" +) + +func userSpecFile(user, file string) string { + return fmt.Sprintf("%s/%s/%s", WorkingDir, user, file) +} + +func (file userFile) read(user string) ([]byte, error) { + return ioutil.ReadFile(userSpecFile(user, string(file))) +} + +func (file userFile) realPath(user string) string { + return userSpecFile(user, string(file)) +} + +// return error if not 400, nil if 400 and no err occurs +func (file userFile) checkPerm(user string) error { + filename := userSpecFile(user, string(file)) + f, err := os.Open(filename) + if err != nil { + return err + } + defer f.Close() + + fi, err := f.Stat() + if err != nil { + return err + } + + if fi.Mode().Perm() != 0400 { + return fmt.Errorf("%v's perm is too open, change it to 400", filename) + } + + return nil +} + +func findUpstreamFromUserfile(conn ssh.ConnMetadata) (net.Conn, error) { + user := conn.User() + + err := UserUpstreamFile.checkPerm(user) + if err != nil { + return nil, err + } + + addr, err := UserUpstreamFile.read(user) + if err != nil { + return nil, err + } + + saddr := strings.TrimSpace(string(addr)) + + logger.Printf("mapping user [%s] to [%s]", user, saddr) + + c, err := net.Dial("tcp", saddr) + if err != nil { + return nil, err + } + + return c, nil +} + +func mapPublicKeyFromUserfile(conn ssh.ConnMetadata, key ssh.PublicKey) (ssh.Signer, error) { + user := conn.User() + + var err error + defer func() { // print error when func exit + if err != nil { + logger.Printf("mapping private key error: %v, public key auth denied for [%v] from [%v]", err, user, conn.RemoteAddr()) + } + }() + + err = UserAuthorizedKeysFile.checkPerm(user) + if err != nil { + return nil, err + } + + keydata := key.Marshal() + + var rest []byte + rest, err = UserAuthorizedKeysFile.read(user) + if err != nil { + return nil, err + } + + var authedPubkey ssh.PublicKey + + for len(rest) > 0 { + authedPubkey, _, _, rest, err = ssh.ParseAuthorizedKey(rest) + + if err != nil { + return nil, err + } + + if bytes.Equal(authedPubkey.Marshal(), keydata) { + err = UserKeyFile.checkPerm(user) + if err != nil { + return nil, err + } + + var privateBytes []byte + privateBytes, err = UserKeyFile.read(user) + if err != nil { + return nil, err + } + + var private ssh.Signer + private, err = ssh.ParsePrivateKey(privateBytes) + if err != nil { + return nil, err + } + + // in log may see this twice, one is for query the other is real sign again + logger.Printf("auth succ, using mapped private key [%v] for user [%v] from [%v]", UserKeyFile.realPath(user), user, conn.RemoteAddr()) + return private, nil + } + } + + logger.Printf("public key auth failed user [%v] from [%v]", conn.User(), conn.RemoteAddr()) + + return nil, nil +} diff --git a/sshpiperd/workingdir_test.go b/sshpiperd/workingdir_test.go new file mode 100644 index 00000000..28fb148d --- /dev/null +++ b/sshpiperd/workingdir_test.go @@ -0,0 +1,166 @@ +// Copyright 2014, 2015 tgic. All rights reserved. +// this file is governed by MIT-license +// +// https://github.com/tg123/sshpiper + +package main + +import ( + "bytes" + //"fmt" + "io" + "io/ioutil" + "net" + "os" + "testing" +) + +func buildWorkingDir(users []string, t *testing.T) { + WorkingDir = "" + dir, err := ioutil.TempDir(os.TempDir(), "sshpiperd_workingdir") + + if err != nil { + t.Fatalf("setup temp dir:%v", err) + } + + WorkingDir = dir + + for _, u := range users { + os.Mkdir(WorkingDir+"/"+u, os.ModePerm) + } + + t.Logf("switch workingdir to %v", WorkingDir) +} + +func cleanupWorkdir(t *testing.T) { + if WorkingDir == "" { + return + } + + t.Logf("cleaning workingdir %v", WorkingDir) + + os.RemoveAll(WorkingDir) +} + +func TestReadUserFile(t *testing.T) { + user1 := "testuser1" + user2 := "testuser2" + + buildWorkingDir([]string{user1, user2}, t) + defer cleanupWorkdir(t) + + data1 := []byte("byte[] := data1") + data2 := []byte("this is data2") + + f := userFile("f") + + err := ioutil.WriteFile(f.realPath(user1), data1, os.ModePerm) + if err != nil { + t.Fatalf("cant create file: %v", err) + } + + err = ioutil.WriteFile(f.realPath(user2), data2, os.ModePerm) + if err != nil { + t.Fatalf("cant create file: %v", err) + } + + d, err := f.read(user1) + if err != nil || !bytes.Equal(d, data1) { + t.Fatalf("read faild") + } + + d, err = f.read(user2) + if err != nil || bytes.Equal(d, data1) { + t.Fatalf("reading wrong user file") + } +} + +func TestCheckPerm(t *testing.T) { + user := "testuser" + buildWorkingDir([]string{user}, t) + defer cleanupWorkdir(t) + + f := userFile("perm") + + err := ioutil.WriteFile(f.realPath(user), nil, os.ModePerm) + if err != nil { + t.Fatalf("cant create file: %v", err) + } + + err = f.checkPerm(user) + if err == nil { + t.Fatalf("should fail when read 0777 user file") + } + + err = os.Chmod(f.realPath(user), 0400) + if err != nil { + t.Fatalf("cant change file mode %v", err) + } + + err = f.checkPerm(user) + if err != nil { + t.Fatalf("fail when read 0400 user file", err) + } +} + +type stubConnMetadata struct{ user string } + +func (s stubConnMetadata) User() string { + return s.user +} + +func (s stubConnMetadata) SessionID() []byte { return nil } +func (s stubConnMetadata) ClientVersion() []byte { return nil } +func (s stubConnMetadata) ServerVersion() []byte { return nil } +func (s stubConnMetadata) RemoteAddr() net.Addr { return nil } +func (s stubConnMetadata) LocalAddr() net.Addr { return nil } + +func TestFindUpstreamFromUserfile(t *testing.T) { + user := "testuser" + buildWorkingDir([]string{user}, t) + defer cleanupWorkdir(t) + + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("cant create fake server: %v", err) + } + defer listener.Close() + + go func() { + c, _ := listener.Accept() + io.Copy(c, c) + c.Close() + }() + + addr := listener.Addr().String() + t.Logf("fake server at %v", addr) + + err = ioutil.WriteFile(UserUpstreamFile.realPath(user), []byte(addr), 0400) + if err != nil { + t.Fatalf("cant create file: %v", err) + } + + t.Logf("testing conn dial to %v", addr) + conn, err := findUpstreamFromUserfile(stubConnMetadata{user}) + + d := []byte("hello") + + _, err = conn.Write(d) + if err != nil { + t.Fatalf("cant write to conn: %v", err) + } + + b := make([]byte, len(d)) + _, err = conn.Read(b) + + if err != nil || !bytes.Equal(b, d) { + t.Fatalf("conn to upstream does not work") + } + + t.Logf("testing user not found") + _, err = findUpstreamFromUserfile(stubConnMetadata{"nosuchuser"}) + if err == nil { + t.Fatalf("should return err when finding nosuchuser") + } + +}