export more api

This commit is contained in:
tgic 2014-12-16 17:27:54 +08:00 committed by Boshi Lian
parent 926bd2438a
commit c44d089252
2 changed files with 68 additions and 25 deletions

View file

@ -6,12 +6,12 @@ import (
"net"
)
type SSHPiper struct {
DownstreamConfig ServerConfig
type SSHPiperConfig struct {
AdditionalChallenge func(conn ConnMetadata, client KeyboardInteractiveChallenge) (bool, error)
FindUpstream func(conn ConnMetadata) (net.Conn, *ClientConfig, error)
MapPublicKey func(conn ConnMetadata, key PublicKey) (Signer, error)
downstreamConfig ServerConfig
}
type upstream struct{ *connection }
@ -22,20 +22,51 @@ type pipedConn struct {
downstream *downstream
processAuthMsg func(msg *userAuthRequestMsg) (*userAuthRequestMsg, error)
clientConfig *ClientConfig
}
func (piper *SSHPiper) Serve(conn net.Conn) error {
type SSHPipe struct{ *pipedConn }
d, err := newDownstream(conn, &piper.DownstreamConfig)
if err != nil {
return err
func (p *SSHPipe) Wait() error {
return p.pipedConn.loop()
}
func (p *SSHPipe) Close() {
p.pipedConn.Close()
}
func (p *SSHPipe) GetUpstreamClientConfig() *ClientConfig {
return p.pipedConn.clientConfig
}
func (piper *SSHPiperConfig) GetDownstreamServerConfig() *ServerConfig {
return &piper.downstreamConfig
}
func (piper *SSHPiperConfig) AddHostKey(key Signer) {
piper.downstreamConfig.AddHostKey(key)
}
func (piper *SSHPiperConfig) Serve(conn net.Conn) (pipe *SSHPipe, err error) {
if piper.FindUpstream == nil {
return nil, fmt.Errorf("FindUpstream func not found")
}
defer d.Close()
d, err := newDownstream(conn, &piper.downstreamConfig)
if err != nil {
return nil, err
}
defer func() {
if pipe == nil {
d.Close()
}
}()
userAuthReq, err := d.nextAuthMsg()
if err != nil {
return err
return nil, err
}
d.user = userAuthReq.User
@ -49,13 +80,13 @@ func (piper *SSHPiper) Serve(conn net.Conn) error {
}))
if err != nil {
return err
return nil, err
}
userAuthReq, err := d.nextAuthMsg()
if err != nil {
return err
return nil, err
}
if userAuthReq.Method == "keyboard-interactive" {
@ -67,36 +98,41 @@ func (piper *SSHPiper) Serve(conn net.Conn) error {
ok, err := piper.AdditionalChallenge(d, prompter.Challenge)
if err != nil {
return err
return nil, err
}
if !ok {
return fmt.Errorf("additional challenge failed")
return nil, fmt.Errorf("additional challenge failed")
}
}
upconn, upconfig, err := piper.FindUpstream(d)
if err != nil {
return err
return nil, err
}
addr := upconn.RemoteAddr().String()
u, err := newUpstream(upconn, addr, upconfig)
if err != nil {
return err
return nil, err
}
defer u.Close()
defer func() {
if pipe == nil {
u.Close()
}
}()
p := &pipedConn{
upstream: u,
downstream: d,
upstream: u,
downstream: d,
clientConfig: upconfig,
}
p.processAuthMsg = func(msg *userAuthRequestMsg) (*userAuthRequestMsg, error) {
// only public msg need
if msg.Method != "publickey" {
if msg.Method != "publickey" || piper.MapPublicKey == nil {
return msg, nil
}
@ -143,11 +179,11 @@ func (piper *SSHPiper) Serve(conn net.Conn) error {
err = p.pipeAuth(userAuthReq)
if err != nil {
return err
return nil, err
}
// block until connection closed or errors occur
return p.loop()
return &SSHPipe{p}, nil
}
func (pipe *pipedConn) validAndAck(upKey, downKey PublicKey) (*userAuthRequestMsg, error) {

View file

@ -168,7 +168,7 @@ func main() {
return
}
piper := &ssh.SSHPiper{
piper := &ssh.SSHPiperConfig{
FindUpstream: findUpstreamFromUserfile,
MapPublicKey: mapPublicKeyFromUserfile,
}
@ -193,7 +193,7 @@ func main() {
logger.Fatalln(err)
}
piper.DownstreamConfig.AddHostKey(private)
piper.AddHostKey(private)
listener, err := net.Listen("tcp", fmt.Sprintf("%s:%d", ListenAddr, Port))
if err != nil {
@ -212,8 +212,15 @@ func main() {
logger.Printf("connection accepted: %v", c.RemoteAddr())
go func() {
err := piper.Serve(c)
logger.Printf("connection %v closed reason: %v", c.RemoteAddr(), err)
p, err := piper.Serve(c)
if err != nil {
logger.Printf("connection from %v establishing failed reason: %v", c.RemoteAddr(), err)
return
}
err = p.Wait()
logger.Printf("connection from %v closed reason: %v", c.RemoteAddr(), err)
}()
}
}