more test case for sshpiperd

This commit is contained in:
tgic 2014-12-21 03:07:56 +08:00 committed by Boshi Lian
parent 91ffa77be9
commit 3fa58e66ca
5 changed files with 321 additions and 129 deletions

View file

@ -6,3 +6,4 @@ go:
script:
- go test ./ssh
- go test ./sshpiperd

View file

@ -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 (

View file

@ -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
View 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
}

View 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")
}
}