diff --git a/.gitignore b/.gitignore index 365d692f..538ca41f 100644 --- a/.gitignore +++ b/.gitignore @@ -23,5 +23,5 @@ _testmain.go *.test *.prof sshpiperd/example/sshpiperd_key* -snap -sshpiperd +sshpiperd/snap +sshpiperd/sshpiperd \ No newline at end of file diff --git a/sshpiperd/auditor/provider.go b/sshpiperd/auditor/provider.go index d93bf9c6..c53f5a6b 100644 --- a/sshpiperd/auditor/provider.go +++ b/sshpiperd/auditor/provider.go @@ -43,7 +43,7 @@ func Register(name string, driver Provider) { drivers.Register(name, driver) } -// All return all registerd auditors +// All return all registered auditors func All() []string { return drivers.Drivers() } diff --git a/sshpiperd/auditor/typescriptlogger/audit.go b/sshpiperd/auditor/typescriptlogger/audit.go new file mode 100644 index 00000000..b45a0845 --- /dev/null +++ b/sshpiperd/auditor/typescriptlogger/audit.go @@ -0,0 +1,86 @@ +package typescriptlogger + +import ( + "fmt" + "os" + "path" + "time" + + "golang.org/x/crypto/ssh" +) + +const ( + msgChannelData = 94 +) + +type filePtyLogger struct { + typescript *os.File + timing *os.File + + oldtime time.Time +} + +func newFilePtyLogger(outputdir string) (*filePtyLogger, error) { + + now := time.Now() + + filename := fmt.Sprintf("%d", now.Unix()) + + typescript, err := os.OpenFile(path.Join(outputdir, fmt.Sprintf("%v.typescript", filename)), os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0600) + + if err != nil { + return nil, err + } + + _, err = typescript.Write([]byte(fmt.Sprintf("Script started on %v\n", now.Format(time.ANSIC)))) + + if err != nil { + return nil, err + } + + timing, err := os.OpenFile(path.Join(outputdir, fmt.Sprintf("%v.timing", filename)), os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0600) + + if err != nil { + return nil, err + } + + return &filePtyLogger{ + typescript: typescript, + timing: timing, + oldtime: time.Now(), + }, nil +} + +func (l *filePtyLogger) loggingTty(conn ssh.ConnMetadata, msg []byte) ([]byte, error) { + + if msg[0] == msgChannelData { + + buf := msg[9:] + + now := time.Now() + + delta := now.Sub(l.oldtime) + + // see term-utils/script.c + fmt.Fprintf(l.timing, "%v.%06v %v\n", int64(delta/time.Second), int64(delta/time.Microsecond), len(buf)) + + l.oldtime = now + + _, err := l.typescript.Write(buf) + + if err != nil { + return msg, err + } + + } + + return msg, nil +} + +func (l *filePtyLogger) Close() (err error) { + _, err = l.typescript.Write([]byte(fmt.Sprintf("Script done on %v\n", time.Now().Format(time.ANSIC)))) + l.typescript.Close() + l.timing.Close() + + return nil // TODO +} diff --git a/sshpiperd/auditor/typescriptlogger/plugin.go b/sshpiperd/auditor/typescriptlogger/plugin.go new file mode 100644 index 00000000..e593759e --- /dev/null +++ b/sshpiperd/auditor/typescriptlogger/plugin.go @@ -0,0 +1,52 @@ +package typescriptlogger + +import ( + "log" + "os" + "path" + + "golang.org/x/crypto/ssh" + + "github.com/tg123/sshpiper/sshpiperd/auditor" +) + +type plugin struct { + Config struct { + OutputDir string `long:"auditor-typescriptlogger-outputdir" default:"/var/sshpiper" description:"Place where logged typescript files were saved" env:"SSHPIPERD_AUDITOR_TYPESCRIPTLOGGER_OUTPUTDIR" ini-name:"auditor-typescriptlogger-outputdir"` + } +} + +func (p *plugin) GetName() string { + return "typescript-logger" +} + +func (p *plugin) GetOpts() interface{} { + return &p.Config +} + +func (p *plugin) Create(conn ssh.ConnMetadata) (auditor.Auditor, error) { + dir := path.Join(p.Config.OutputDir, conn.User()) + err := os.MkdirAll(dir, 0700) + if err != nil { + return nil, err + } + + return newFilePtyLogger(dir) +} + +func (p *plugin) Init(logger *log.Logger) error { + + return nil +} + +func (l *filePtyLogger) GetUpstreamHook() auditor.Hook { + return l.loggingTty +} + +func (l *filePtyLogger) GetDownstreamHook() auditor.Hook { + return nil +} + +func init() { + auditor.Register("typescript-logger", new(plugin)) +} diff --git a/sshpiperd/challenger/provider.go b/sshpiperd/challenger/provider.go index 485145be..60a0323a 100644 --- a/sshpiperd/challenger/provider.go +++ b/sshpiperd/challenger/provider.go @@ -27,7 +27,7 @@ func Register(name string, driver Provider) { drivers.Register(name, driver) } -// All return all registerd challenger +// All return all registered challenger func All() []string { return drivers.Drivers() } diff --git a/sshpiperd/registry/plugin.go b/sshpiperd/registry/plugin.go index 61278e2f..0a07bf2c 100644 --- a/sshpiperd/registry/plugin.go +++ b/sshpiperd/registry/plugin.go @@ -4,10 +4,16 @@ import ( "log" ) +// Plugin is to be registered with sshpiper to provide additional functions type Plugin interface { + + // The name of the Plugin GetName() string + // A ref to a struct which holds the options for the plugins + // will be populated by cmd or other plugin runners GetOpts() interface{} + // Will be called before the Plugin is used to ensure the Plugin is ready Init(logger *log.Logger) error } diff --git a/sshpiperd/registry/registry.go b/sshpiperd/registry/registry.go index a46d1267..28e6024a 100644 --- a/sshpiperd/registry/registry.go +++ b/sshpiperd/registry/registry.go @@ -5,17 +5,20 @@ import ( "sync" ) +// Registry is a place to hold all plugins type Registry struct { driversMu sync.RWMutex drivers map[string]interface{} } +// NewRegistry creates a new Registry func NewRegistry() *Registry { return &Registry{drivers: make(map[string]interface{})} } -// copy from database/sql +// Register adds a Plugin with given name to Registry func (r *Registry) Register(name string, driver interface{}) { + // copy from database/sql r.driversMu.Lock() defer r.driversMu.Unlock() if driver == nil { @@ -27,6 +30,7 @@ func (r *Registry) Register(name string, driver interface{}) { r.drivers[name] = driver } +// Drivers return all registered Plugins func (r *Registry) Drivers() []string { r.driversMu.RLock() defer r.driversMu.RUnlock() @@ -38,6 +42,7 @@ func (r *Registry) Drivers() []string { return list } +// Get returns an Plugins by name, return nil if not found func (r *Registry) Get(name string) interface{} { r.driversMu.RLock() defer r.driversMu.RUnlock() diff --git a/sshpiperd/upstream/provider.go b/sshpiperd/upstream/provider.go index 0787a75c..6d8d7e32 100644 --- a/sshpiperd/upstream/provider.go +++ b/sshpiperd/upstream/provider.go @@ -10,7 +10,7 @@ import ( // Handler will be installed into sshpiper and help to establish the connection to upstream // the returned auth pipe is to map/convert downstream auth method to another auth for -// connecting to upstrem. +// connecting to upstream. // e.g. map downstream public key to another upstream private key type Handler func(conn ssh.ConnMetadata) (net.Conn, *ssh.SSHPiperAuthPipe, error) @@ -30,7 +30,7 @@ func Register(name string, driver Provider) { drivers.Register(name, driver) } -// All return all registerd upstream providers +// All return all registered upstream providers func All() []string { return drivers.Drivers() } diff --git a/sshpiperd/upstream/workingdir/plugin.go b/sshpiperd/upstream/workingdir/plugin.go new file mode 100644 index 00000000..bd996417 --- /dev/null +++ b/sshpiperd/upstream/workingdir/plugin.go @@ -0,0 +1,37 @@ +package workingdir + +import ( + "log" + + "github.com/tg123/sshpiper/sshpiperd/upstream" +) + +var logger *log.Logger + +type plugin struct { +} + +func (p *plugin) GetName() string { + return "workingdir" +} + +func (p *plugin) GetOpts() interface{} { + return &config +} + +func (p *plugin) GetHandler() upstream.Handler { + return findUpstreamFromUserfile +} + +func (p *plugin) Init(glogger *log.Logger) error { + + logger = glogger + + logger.Printf("upstream provider: workingdir %v init", config.WorkingDir) + + return nil +} + +func init() { + upstream.Register("workingdir", &plugin{}) +} diff --git a/sshpiperd/upstream/workingdir/workingdir.go b/sshpiperd/upstream/workingdir/workingdir.go index 36684cdf..dcad81d2 100644 --- a/sshpiperd/upstream/workingdir/workingdir.go +++ b/sshpiperd/upstream/workingdir/workingdir.go @@ -22,9 +22,9 @@ import ( type userFile string var ( - UserAuthorizedKeysFile userFile = "authorized_keys" - UserKeyFile userFile = "id_rsa" - UserUpstreamFile userFile = "sshpiper_upstream" + userAuthorizedKeysFile userFile = "authorized_keys" + userKeyFile userFile = "id_rsa" + userUpstreamFile userFile = "sshpiper_upstream" usernameRule *regexp.Regexp ) @@ -129,12 +129,12 @@ func findUpstreamFromUserfile(conn ssh.ConnMetadata) (net.Conn, *ssh.SSHPiperAut return nil, nil, fmt.Errorf("downstream is not using a valid username") } - err := UserUpstreamFile.checkPerm(user) + err := userUpstreamFile.checkPerm(user) if err != nil { return nil, nil, err } - data, err := UserUpstreamFile.read(user) + data, err := userUpstreamFile.read(user) if err != nil { return nil, nil, err } @@ -183,7 +183,7 @@ func mapPublicKeyFromUserfile(conn ssh.ConnMetadata, key ssh.PublicKey) (signer } }() - err = UserAuthorizedKeysFile.checkPerm(user) + err = userAuthorizedKeysFile.checkPerm(user) if err != nil { return nil, err } @@ -191,7 +191,7 @@ func mapPublicKeyFromUserfile(conn ssh.ConnMetadata, key ssh.PublicKey) (signer keydata := key.Marshal() var rest []byte - rest, err = UserAuthorizedKeysFile.read(user) + rest, err = userAuthorizedKeysFile.read(user) if err != nil { return nil, err } @@ -206,13 +206,13 @@ func mapPublicKeyFromUserfile(conn ssh.ConnMetadata, key ssh.PublicKey) (signer } if bytes.Equal(authedPubkey.Marshal(), keydata) { - err = UserKeyFile.checkPerm(user) + err = userKeyFile.checkPerm(user) if err != nil { return nil, err } var privateBytes []byte - privateBytes, err = UserKeyFile.read(user) + privateBytes, err = userKeyFile.read(user) if err != nil { return nil, err } @@ -224,7 +224,7 @@ func mapPublicKeyFromUserfile(conn ssh.ConnMetadata, key ssh.PublicKey) (signer } // 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()) + logger.Printf("auth succ, using mapped private key [%v] for user [%v] from [%v]", userKeyFile.realPath(user), user, conn.RemoteAddr()) return private, nil } } diff --git a/sshpiperd/upstream/workingdir/workingdir_test.go b/sshpiperd/upstream/workingdir/workingdir_test.go index 63e8991f..e869f279 100644 --- a/sshpiperd/upstream/workingdir/workingdir_test.go +++ b/sshpiperd/upstream/workingdir/workingdir_test.go @@ -205,7 +205,7 @@ func TestFindUpstreamFromUserfile(t *testing.T) { addr := listener.Addr().String() t.Logf("fake server at %v", addr) - err = ioutil.WriteFile(UserUpstreamFile.realPath(user), []byte(addr), 0777) + err = ioutil.WriteFile(userUpstreamFile.realPath(user), []byte(addr), 0777) if err != nil { t.Fatalf("cant create file: %v", err) } @@ -216,7 +216,7 @@ func TestFindUpstreamFromUserfile(t *testing.T) { t.Fatalf("should return err when file too open") } - err = os.Chmod(UserUpstreamFile.realPath(user), 0400) + err = os.Chmod(userUpstreamFile.realPath(user), 0400) if err != nil { t.Fatalf("cant change file mode %v", err) } @@ -260,13 +260,13 @@ func TestMapPublicKeyFromUserfile(t *testing.T) { _ = privateKey2 - err := ioutil.WriteFile(UserKeyFile.realPath(user), testdata.PEMBytes["rsa"], 0777) + err := ioutil.WriteFile(userKeyFile.realPath(user), testdata.PEMBytes["rsa"], 0777) if err != nil { t.Fatalf("cant create file: %v", err) } authKeys := ssh.MarshalAuthorizedKey(publicKey) - err = ioutil.WriteFile(UserAuthorizedKeysFile.realPath(user), authKeys, 0777) + err = ioutil.WriteFile(userAuthorizedKeysFile.realPath(user), authKeys, 0777) if err != nil { t.Fatalf("cant create file: %v", err) } @@ -279,7 +279,7 @@ func TestMapPublicKeyFromUserfile(t *testing.T) { t.Fatalf("should return err when file too open") } - err = os.Chmod(UserAuthorizedKeysFile.realPath(user), 0600) + err = os.Chmod(userAuthorizedKeysFile.realPath(user), 0600) if err != nil { t.Fatalf("cant change file mode %v", err) } @@ -290,7 +290,7 @@ func TestMapPublicKeyFromUserfile(t *testing.T) { t.Fatalf("should return err when file too open") } - err = os.Chmod(UserKeyFile.realPath(user), 0600) + err = os.Chmod(userKeyFile.realPath(user), 0600) if err != nil { t.Fatalf("cant change file mode %v", err) } @@ -314,7 +314,7 @@ func TestMapPublicKeyFromUserfile(t *testing.T) { t.Logf("testing not in UserAuthorizedKeysFile") authKeys = ssh.MarshalAuthorizedKey(privateKey2.PublicKey()) - err = ioutil.WriteFile(UserAuthorizedKeysFile.realPath(user), authKeys, 0600) + err = ioutil.WriteFile(userAuthorizedKeysFile.realPath(user), authKeys, 0600) if err != nil { t.Fatalf("cant create file: %v", err) }