more test case for sshpiperd
This commit is contained in:
parent
91ffa77be9
commit
3fa58e66ca
5 changed files with 321 additions and 129 deletions
|
|
@ -6,3 +6,4 @@ go:
|
|||
|
||||
script:
|
||||
- go test ./ssh
|
||||
- go test ./sshpiperd
|
||||
|
|
|
|||
|
|
@ -1,3 +1,8 @@
|
|||
// Copyright 2014, 2015 tgic<farmer1992@gmail.com>. All rights reserved.
|
||||
// this file is governed by MIT-license
|
||||
//
|
||||
// https://github.com/tg123/sshpiper
|
||||
|
||||
package ssh
|
||||
|
||||
import (
|
||||
|
|
|
|||
|
|
@ -1,7 +1,11 @@
|
|||
// Copyright 2014, 2015 tgic<farmer1992@gmail.com>. 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,
|
||||
|
|
|
|||
143
sshpiperd/workingdir.go
Normal file
143
sshpiperd/workingdir.go
Normal file
|
|
@ -0,0 +1,143 @@
|
|||
// Copyright 2014, 2015 tgic<farmer1992@gmail.com>. 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
|
||||
}
|
||||
166
sshpiperd/workingdir_test.go
Normal file
166
sshpiperd/workingdir_test.go
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
// Copyright 2014, 2015 tgic<farmer1992@gmail.com>. 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")
|
||||
}
|
||||
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue