This commit is contained in:
qaq
2025-06-17 12:17:33 +08:00
parent 4a6b2c8e0f
commit 9fa968d07c
648 changed files with 64968 additions and 1621 deletions
+48
View File
@@ -0,0 +1,48 @@
package ssh
import (
"net"
"golang.org/x/crypto/ssh"
)
// a dummy ssh client conn used by client connector
type ClientConn struct {
net.Conn
client *ssh.Client
}
func NewClientConn(session *Session) net.Conn {
return &ClientConn{
Conn: session.Conn,
client: session.client,
}
}
func (c *ClientConn) Client() *ssh.Client {
return c.client
}
type sshConn struct {
channel ssh.Channel
net.Conn
}
func NewConn(conn net.Conn, channel ssh.Channel) net.Conn {
return &sshConn{
Conn: conn,
channel: channel,
}
}
func (c *sshConn) Read(b []byte) (n int, err error) {
return c.channel.Read(b)
}
func (c *sshConn) Write(b []byte) (n int, err error) {
return c.channel.Write(b)
}
func (c *sshConn) Close() error {
return c.channel.Close()
}
+126
View File
@@ -0,0 +1,126 @@
package ssh
import (
"context"
"net"
"time"
"github.com/go-gost/core/logger"
"golang.org/x/crypto/ssh"
)
const (
defaultKeepaliveInterval = 30 * time.Second
defaultKeepaliveTimeout = 15 * time.Second
defaultkeepaliveRetries = 1
)
type Session struct {
net.Conn
client *ssh.Client
closed chan struct{}
dead chan struct{}
log logger.Logger
}
func NewSession(c net.Conn, client *ssh.Client, log logger.Logger) *Session {
return &Session{
Conn: c,
client: client,
closed: make(chan struct{}),
dead: make(chan struct{}),
log: log,
}
}
func (s *Session) OpenChannel(name string) (ssh.Channel, <-chan *ssh.Request, error) {
return s.client.OpenChannel(name, nil)
}
func (s *Session) IsClosed() bool {
select {
case <-s.dead:
return true
case <-s.closed:
return true
default:
}
return false
}
func (s *Session) Wait() error {
defer close(s.closed)
return s.client.Wait()
}
func (s *Session) WaitClose() {
defer s.client.Close()
select {
case <-s.dead:
s.log.Debugf("session is dead")
case <-s.closed:
s.log.Debugf("session is closed")
}
}
func (s *Session) Keepalive(interval, timeout time.Duration, retries int) {
if interval <= 0 {
interval = defaultKeepaliveInterval
}
if timeout <= 0 {
timeout = defaultKeepaliveTimeout
}
if retries <= 0 {
retries = defaultkeepaliveRetries
}
s.log.Debugf("keepalive is enabled, interval: %v, timeout: %v, retries: %d", interval, timeout, retries)
defer close(s.dead)
t := time.NewTicker(interval)
defer t.Stop()
count := retries
for {
select {
case <-t.C:
start := time.Now()
err := func() error {
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
select {
case err := <-s.ping():
return err
case <-ctx.Done():
return ctx.Err()
}
}()
if err != nil {
s.log.Debugf("ssh ping: %v", err)
count--
if count == 0 {
return
}
continue
}
s.log.Debugf("ssh ping OK, RTT: %v", time.Since(start))
count = retries
case <-s.closed:
return
}
}
}
func (s *Session) ping() <-chan error {
ch := make(chan error, 1)
go func() {
defer close(ch)
if _, _, err := s.client.SendRequest("ping", true, nil); err != nil {
ch <- err
}
}()
return ch
}
+76
View File
@@ -0,0 +1,76 @@
package ssh
import (
"context"
"errors"
"fmt"
"os"
"github.com/go-gost/core/auth"
"golang.org/x/crypto/ssh"
)
const (
GostSSHTunnelRequest = "gost-tunnel" // extended request type for ssh tunnel
)
var (
ErrSessionDead = errors.New("session is dead")
)
// PasswordCallbackFunc is a callback function used by SSH server.
// It authenticates user using a password.
type PasswordCallbackFunc func(conn ssh.ConnMetadata, password []byte) (*ssh.Permissions, error)
func PasswordCallback(au auth.Authenticator) PasswordCallbackFunc {
if au == nil {
return nil
}
return func(conn ssh.ConnMetadata, password []byte) (*ssh.Permissions, error) {
if _, ok := au.Authenticate(context.Background(), conn.User(), string(password)); ok {
return nil, nil
}
return nil, fmt.Errorf("password rejected for %s", conn.User())
}
}
// PublicKeyCallbackFunc is a callback function used by SSH server.
// It offers a public key for authentication.
type PublicKeyCallbackFunc func(c ssh.ConnMetadata, pubKey ssh.PublicKey) (*ssh.Permissions, error)
func PublicKeyCallback(keys map[string]bool) PublicKeyCallbackFunc {
if len(keys) == 0 {
return nil
}
return func(c ssh.ConnMetadata, pubKey ssh.PublicKey) (*ssh.Permissions, error) {
if keys[string(pubKey.Marshal())] {
return &ssh.Permissions{
// Record the public key used for authentication.
Extensions: map[string]string{
"pubkey-fp": ssh.FingerprintSHA256(pubKey),
},
}, nil
}
return nil, fmt.Errorf("unknown public key for %q", c.User())
}
}
// ParseSSHAuthorizedKeysFile parses ssh authorized keys file.
func ParseAuthorizedKeysFile(name string) (map[string]bool, error) {
authorizedKeysBytes, err := os.ReadFile(name)
if err != nil {
return nil, err
}
authorizedKeysMap := make(map[string]bool)
for len(authorizedKeysBytes) > 0 {
pubKey, _, _, rest, err := ssh.ParseAuthorizedKey(authorizedKeysBytes)
if err != nil {
return nil, err
}
authorizedKeysMap[string(pubKey.Marshal())] = true
authorizedKeysBytes = rest
}
return authorizedKeysMap, nil
}