summaryrefslogtreecommitdiff
path: root/internal/clients/connectors/serverconnection.go
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2021-09-18 18:43:19 +0300
committerPaul Buetow <paul@buetow.org>2021-10-02 12:26:29 +0300
commit69b88a1cae0a61bd22530c384f40166b37b9f1ea (patch)
tree3a0bbe5a25c3035e765ed40133f5a41f4f8dfedd /internal/clients/connectors/serverconnection.go
parent6506e20f6c80f4acb7434eb9dd14f784a67189cd (diff)
remote connector is now an interface
Diffstat (limited to 'internal/clients/connectors/serverconnection.go')
-rw-r--r--internal/clients/connectors/serverconnection.go225
1 files changed, 225 insertions, 0 deletions
diff --git a/internal/clients/connectors/serverconnection.go b/internal/clients/connectors/serverconnection.go
new file mode 100644
index 0000000..fab2f87
--- /dev/null
+++ b/internal/clients/connectors/serverconnection.go
@@ -0,0 +1,225 @@
+package connectors
+
+import (
+ "context"
+ "fmt"
+ "io"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/mimecast/dtail/internal/clients/handlers"
+ "github.com/mimecast/dtail/internal/config"
+ "github.com/mimecast/dtail/internal/io/logger"
+ "github.com/mimecast/dtail/internal/ssh/client"
+
+ "golang.org/x/crypto/ssh"
+)
+
+// ServerConnection represents a client connection connection to a single server.
+type ServerConnection struct {
+ // The remote server's hostname connected to.
+ server string
+ // The remote server's port connected to.
+ port int
+ // The SSH client configuration used.
+ config *ssh.ClientConfig
+ // The SSH client handler to use.
+ handler handlers.Handler
+ // DTail commands sent from client to server. When client loses
+ // connection to the server it re-connects automatically and sends the
+ // same commands again.
+ commands []string
+ // Is it a persistent connection or a one-off?
+ isOneOff bool
+ // To deal with SSH server host keys
+ hostKeyCallback client.HostKeyCallback
+ // To determine if connection throttling has finished or not
+ throttlingDone bool
+}
+
+// NewServerConnection returns a new connection.
+func NewServerConnection(server string, userName string, authMethods []ssh.AuthMethod, hostKeyCallback client.HostKeyCallback, handler handlers.Handler, commands []string) *ServerConnection {
+ logger.Debug(server, "Creating new connection")
+
+ c := ServerConnection{
+ hostKeyCallback: hostKeyCallback,
+ server: server,
+ handler: handler,
+ commands: commands,
+ config: &ssh.ClientConfig{
+ User: userName,
+ Auth: authMethods,
+ HostKeyCallback: hostKeyCallback.Wrap(),
+ Timeout: time.Second * 2,
+ },
+ }
+
+ c.initServerPort()
+ return &c
+}
+
+// NewOneOffServerConnection creates new one-off connection (only for sending a series of commands and then quit).
+func NewOneOffServerConnection(server string, userName string, authMethods []ssh.AuthMethod, handler handlers.Handler, commands []string) *ServerConnection {
+ c := ServerConnection{
+ server: server,
+ handler: handler,
+ commands: commands,
+ config: &ssh.ClientConfig{
+ User: userName,
+ Auth: authMethods,
+ HostKeyCallback: ssh.InsecureIgnoreHostKey(),
+ },
+ isOneOff: true,
+ }
+
+ c.initServerPort()
+ return &c
+}
+
+// Server hostname
+func (c *ServerConnection) Server() string {
+ return c.server
+}
+
+// Handler for the connection
+func (c *ServerConnection) Handler() handlers.Handler {
+ return c.handler
+}
+
+// Attempt to parse the server port address from the provided server FQDN.
+func (c *ServerConnection) initServerPort() {
+ c.port = config.Common.SSHPort
+ parts := strings.Split(c.server, ":")
+
+ if len(parts) == 2 {
+ logger.Debug("Parsing port from hostname", parts)
+ port, err := strconv.Atoi(parts[1])
+ if err != nil {
+ logger.FatalExit("Unable to parse client port", c.server, parts, err)
+ }
+ c.server = parts[0]
+ c.port = port
+ }
+}
+
+// Start the server connection. Build up SSH session and send some DTail commands.
+func (c *ServerConnection) Start(ctx context.Context, cancel context.CancelFunc, throttleCh, statsCh chan struct{}) {
+ // Throttle how many connections can be established concurrently (based on ch length)
+ logger.Debug(c.server, "Throttling connection", len(throttleCh), cap(throttleCh))
+
+ select {
+ case throttleCh <- struct{}{}:
+ case <-ctx.Done():
+ logger.Debug(c.server, "Not establishing connection as context is done", len(throttleCh), cap(throttleCh))
+ return
+ }
+
+ logger.Debug(c.server, "Throttling says that the connection can be established", len(throttleCh), cap(throttleCh))
+
+ go func() {
+ defer func() {
+ if !c.throttlingDone {
+ logger.Debug(c.server, "Unthrottling connection (1)", len(throttleCh), cap(throttleCh))
+ c.throttlingDone = true
+ <-throttleCh
+ }
+ cancel()
+ }()
+
+ if err := c.dial(ctx, cancel, throttleCh, statsCh); err != nil {
+ logger.Warn(c.server, c.port, err)
+ if c.hostKeyCallback.Untrusted(fmt.Sprintf("%s:%d", c.server, c.port)) {
+ logger.Debug(c.server, "Not trusting host")
+ }
+ }
+ }()
+
+ <-ctx.Done()
+}
+
+// Dail into a new SSH connection. Close connection in case of an error.
+func (c *ServerConnection) dial(ctx context.Context, cancel context.CancelFunc, throttleCh, statsCh chan struct{}) error {
+ logger.Debug(c.server, "Incrementing connection stats")
+ statsCh <- struct{}{}
+ defer func() {
+ logger.Debug(c.server, "Decrementing connection stats")
+ <-statsCh
+ }()
+
+ logger.Debug(c.server, "Dialing into the connection")
+ address := fmt.Sprintf("%s:%d", c.server, c.port)
+
+ client, err := ssh.Dial("tcp", address, c.config)
+ if err != nil {
+ return err
+ }
+ defer client.Close()
+
+ return c.session(ctx, cancel, client, throttleCh)
+}
+
+// Create the SSH session. Close the session in case of an error.
+func (c *ServerConnection) session(ctx context.Context, cancel context.CancelFunc, client *ssh.Client, throttleCh chan struct{}) error {
+ logger.Debug(c.server, "session")
+
+ session, err := client.NewSession()
+ if err != nil {
+ return err
+ }
+ defer session.Close()
+
+ return c.handle(ctx, cancel, session, throttleCh)
+}
+
+func (c *ServerConnection) handle(ctx context.Context, cancel context.CancelFunc, session *ssh.Session, throttleCh chan struct{}) error {
+ logger.Debug(c.server, "handle")
+
+ stdinPipe, err := session.StdinPipe()
+ if err != nil {
+ return err
+ }
+
+ stdoutPipe, err := session.StdoutPipe()
+ if err != nil {
+ return err
+ }
+
+ if err := session.Shell(); err != nil {
+ return err
+ }
+
+ go func() {
+ io.Copy(stdinPipe, c.handler)
+ cancel()
+ }()
+
+ go func() {
+ io.Copy(c.handler, stdoutPipe)
+ cancel()
+ }()
+
+ go func() {
+ select {
+ case <-c.handler.Done():
+ case <-ctx.Done():
+ }
+ cancel()
+ }()
+
+ // Send all commands to client.
+ for _, command := range c.commands {
+ logger.Debug(command)
+ c.handler.SendMessage(command)
+ }
+
+ if !c.throttlingDone {
+ logger.Debug(c.server, "Unthrottling connection (2)", len(throttleCh), cap(throttleCh))
+ c.throttlingDone = true
+ <-throttleCh
+ }
+
+ <-ctx.Done()
+ c.handler.Shutdown()
+ return nil
+}