Files
forge-tools-ssh/internal/ssh/manager.go
T

711 lines
18 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package ssh
import (
"context"
"errors"
"fmt"
"io"
"log"
"net"
"os"
"path/filepath"
"strings"
"sync"
"time"
"git.totmin.ru/en2zmax/forge-toolkit/configreload"
"golang.org/x/crypto/ssh"
)
// Manager manages multiple SSH connections for a single session.
type Manager struct {
connections map[string]*Client
primary string
keyManager *KeyManager
// policy - текущая (закэшированная) per-agent политика подключений.
// nil = политика не настроена: ad-hoc режим как раньше (без ограничений).
policy *Policy
// policyLoader - live-reload политики из <FORGE_TENANT_CONFIG>/ssh.json
// (см. configreload). nil, если FORGE_TENANT_CONFIG не задан (легаси) или
// задана статическая политика через SetPolicy (тесты).
policyLoader *configreload.Loader[*Policy]
mu sync.RWMutex
aliasLocks map[string]*sync.Mutex
reconnectFails map[string]time.Time // alias -> last failed reconnect time
dockerCache map[string]*bool // alias -> docker available (nil = unknown)
}
// NewManager creates a new SSH connection manager. Если задан
// FORGE_TENANT_CONFIG/ssh.json - политика подхватывается через live-reload
// (configreload по контент-хэшу): правки ssh.json применяются к следующему
// Connect без рестарта.
func NewManager(keyPath string) *Manager {
mgr := &Manager{
connections: make(map[string]*Client),
keyManager: NewKeyManager(keyPath),
aliasLocks: make(map[string]*sync.Mutex),
reconnectFails: make(map[string]time.Time),
dockerCache: make(map[string]*bool),
}
if err := mgr.keyManager.EnsureKey(); err != nil {
log.Printf("manager: key setup warning: %v", err)
}
if path, ok := PolicyPathFromEnv(); ok {
mgr.policyLoader = configreload.New(path, parsePolicy)
}
return mgr
}
// SetPolicy задаёт per-agent политику подключений вручную (тесты).
// Отключает live-reload (loader == nil), чтобы тесты были детерминированы.
func (m *Manager) SetPolicy(p *Policy) {
_ = p.normalize()
m.policy = p
m.policyLoader = nil
}
// currentPolicy возвращает актуальную политику, перечитав её при изменении
// содержимого ssh.json (live-reload). Нет loader'а → статическая policy.
// Ошибки: ErrNotFound → (nil, err) (нет политики = ad-hoc без ограничений);
// парсинг-ошибка → (lastGood, err) (сохраняем рабочую, но сообщаем).
func (m *Manager) currentPolicy() (*Policy, error) {
if m.policyLoader == nil {
return m.policy, nil
}
return m.policyLoader.Get()
}
// getAliasLock returns a per-alias lock.
func (m *Manager) getAliasLock(alias string) *sync.Mutex {
m.mu.Lock()
defer m.mu.Unlock()
if lock, ok := m.aliasLocks[alias]; ok {
return lock
}
lock := &sync.Mutex{}
m.aliasLocks[alias] = lock
return lock
}
// generateAlias creates a unique alias and reserves it.
func (m *Manager) generateAlias(username, host string) string {
base := fmt.Sprintf("%s@%s", username, host)
m.mu.Lock()
defer m.mu.Unlock()
if _, exists := m.connections[base]; !exists {
m.connections[base] = nil // Reserve
return base
}
for i := 2; i < 100; i++ {
candidate := fmt.Sprintf("%s-%d", base, i)
if _, exists := m.connections[candidate]; !exists {
m.connections[candidate] = nil // Reserve
return candidate
}
}
// Fallback (might collide but highly unlikely to fill 100 slots)
final := fmt.Sprintf("%s-%d", base, 100)
m.connections[final] = nil
return final
}
// ConnectOptions contains options for SSH connection.
type ConnectOptions struct {
Host string
Port int
Username string
Password string
PrivateKeyPath string
Alias string
Via string
// Target is an optional destination host routed through a PAM gateway
// (SafeInspect). When set, the gateway receives "Username@Target" and
// TargetPassword is used as the second authentication stage.
Target string
TargetPassword string
// Profile - имя профиля из ssh.json (см. policy.go). Если задан, креды
// (Host/Username/Port/Key/Via/Target) берутся из политики, а не от
// модели. Модель выбирает только имя; адрес и ключ задаёт оператор.
Profile string
}
// Connect establishes an SSH connection and returns the alias. Применяет
// per-agent политику (см. policy.go): профиль резолвит креды, а host
// обязан пройти allowlist. Модель не может подключиться вне разрешённого.
func (m *Manager) Connect(ctx context.Context, opts ConnectOptions) (alias string, err error) {
if opts.Port == 0 {
opts.Port = 22
}
// Live-reload политики: перечитываем ssh.json по контент-хэшу.
// ErrNotFound → нет политики (ad-hoc без ограничений). Парсинг-ошибка →
// last-good (если был) + предупреждение; без last-good — fail-closed.
pol, perr := m.currentPolicy()
if perr != nil {
if !errors.Is(perr, configreload.ErrNotFound) && pol == nil {
return "", fmt.Errorf("ssh: per-agent policy unavailable: %w", perr)
}
log.Printf("ssh: policy warning: %v", perr)
}
// Резолв профиля: платформа объявляет креды, модель выбирает по имени.
if opts.Profile != "" {
if pol == nil {
return "", errors.New("ssh: profile requested but no per-agent policy configured")
}
pr, ok := pol.Resolve(opts.Profile)
if !ok {
return "", fmt.Errorf("ssh: profile %q not found in policy", opts.Profile)
}
opts.Host = pr.Host
opts.Username = pr.Username
if pr.Port != 0 {
opts.Port = pr.Port
}
if opts.PrivateKeyPath == "" {
opts.PrivateKeyPath = pr.KeyPath
}
if opts.Via == "" {
opts.Via = pr.Via
}
if opts.Target == "" {
opts.Target = pr.Target
if opts.TargetPassword == "" {
opts.TargetPassword = pr.TargetPass
}
}
}
// Креды обязательны (профиль дал их, либо модель прислала host+username).
if opts.Host == "" || opts.Username == "" {
return "", errors.New("ssh: host and username are required")
}
// Allowlist: вне списка - deny (fail-closed). Пустой список = allow all.
if pol != nil && !pol.Allow(opts.Host) {
return "", fmt.Errorf("ssh: host %q is not in allowed_hosts of per-agent policy", opts.Host)
}
// Default key: если ни профиль, ни вызов не задали ключ - берём
// DefaultKeyPath политики, иначе "" (системный ключ KeyManager).
if pol != nil && opts.PrivateKeyPath == "" {
opts.PrivateKeyPath = pol.DefaultKeyPath
}
var reserved bool
if opts.Alias == "" {
// Include the PAM target (if any) so distinct targets routed through
// the same gateway get distinct auto-generated aliases.
aliasBase := opts.Username
if opts.Target != "" {
aliasBase = opts.Username + "@" + opts.Target
}
opts.Alias = m.generateAlias(aliasBase, opts.Host)
reserved = true
}
if opts.Via == opts.Alias {
return "", errors.New("'via' cannot be the same as 'alias'")
}
m.mu.Lock()
existing, exists := m.connections[opts.Alias]
if exists {
if existing != nil {
m.mu.Unlock()
if existing.creds.Host == opts.Host && existing.creds.Username == opts.Username && existing.creds.Target == opts.Target {
return opts.Alias, nil
}
return "", fmt.Errorf("alias '%s' already exists for %s@%s", opts.Alias, existing.creds.Username, existing.creds.Host)
}
// Existing is nil (reserved)
if !reserved {
m.mu.Unlock()
return "", fmt.Errorf("alias '%s' is currently connecting/reserved", opts.Alias)
}
// It's our reservation, proceed
} else {
// New explicit alias
m.connections[opts.Alias] = nil // Reserve
}
m.mu.Unlock()
// Defer cleanup of reservation on error
defer func() {
if err != nil {
m.mu.Lock()
// Only remove if it's still nil (failed to connect)
if c, ok := m.connections[opts.Alias]; ok && c == nil {
delete(m.connections, opts.Alias)
}
m.mu.Unlock()
}
}()
creds := Credentials{
Host: opts.Host,
Port: opts.Port,
Username: opts.Username,
Password: opts.Password,
Via: opts.Via,
Target: opts.Target,
TargetPassword: opts.TargetPassword,
}
if opts.PrivateKeyPath != "" {
keyBytes, err := os.ReadFile(opts.PrivateKeyPath)
if err != nil {
return "", fmt.Errorf("failed to read private key: %w", err)
}
signer, err := ssh.ParsePrivateKey(keyBytes)
if err != nil {
return "", fmt.Errorf("failed to parse private key: %w", err)
}
creds.PrivateKey = signer
} else if opts.Password == "" {
signer, err := m.keyManager.LoadPrivateKey()
if err != nil {
return "", fmt.Errorf("no auth provided and system key unavailable: %w", err)
}
creds.PrivateKey = signer
log.Printf("ssh: using system key for %s", opts.Alias)
}
var jumpClient *Client
if opts.Via != "" {
m.mu.RLock()
jumpClient = m.connections[opts.Via]
m.mu.RUnlock()
if jumpClient == nil {
return "", fmt.Errorf("jump host '%s' not connected", opts.Via)
}
}
client, err := NewClient(ctx, opts.Alias, creds, jumpClient)
if err != nil {
return "", err
}
m.mu.Lock()
m.connections[opts.Alias] = client
if m.primary == "" {
m.primary = opts.Alias
}
m.mu.Unlock()
return opts.Alias, nil
}
// Disconnect closes one or all connections.
func (m *Manager) Disconnect(alias string) (string, error) {
m.mu.Lock()
defer m.mu.Unlock()
if alias == "" {
count := 0
for a, client := range m.connections {
if client != nil {
client.Close()
count++
}
delete(m.connections, a)
delete(m.aliasLocks, a)
delete(m.reconnectFails, a)
delete(m.dockerCache, a)
}
m.primary = ""
return fmt.Sprintf("Disconnected all (%d) connections", count), nil
}
client, ok := m.connections[alias]
if !ok {
return "", fmt.Errorf("no connection with alias '%s'", alias)
}
if client != nil {
client.Close()
}
delete(m.connections, alias)
delete(m.aliasLocks, alias)
delete(m.reconnectFails, alias)
delete(m.dockerCache, alias)
if m.primary == alias {
m.primary = ""
for a, c := range m.connections {
if c != nil {
m.primary = a
break
}
}
}
return fmt.Sprintf("Disconnected '%s'", alias), nil
}
// resolveTarget returns the target alias.
func (m *Manager) resolveTarget(target string) (string, error) {
if target != "" && target != "primary" {
m.mu.RLock()
_, ok := m.connections[target]
m.mu.RUnlock()
if !ok {
return "", fmt.Errorf("no connection with alias '%s'", target)
}
return target, nil
}
m.mu.RLock()
primary := m.primary
m.mu.RUnlock()
if primary == "" {
return "", errors.New("no active connection")
}
return primary, nil
}
// Run executes a command on the target connection.
func (m *Manager) Run(ctx context.Context, cmd, target string) (*RunResult, error) {
alias, err := m.resolveTarget(target)
if err != nil {
return nil, err
}
lock := m.getAliasLock(alias)
lock.Lock()
defer lock.Unlock()
m.mu.RLock()
client := m.connections[alias]
m.mu.RUnlock()
if client == nil {
return nil, fmt.Errorf("connection '%s' not found", alias)
}
result, err := client.Run(ctx, cmd)
if err != nil {
if isConnectionError(err) {
// Check reconnect backoff
m.mu.RLock()
lastFail := m.reconnectFails[alias]
m.mu.RUnlock()
if time.Since(lastFail) < 5*time.Second {
return nil, fmt.Errorf("connection lost (reconnect backoff): %w", err)
}
log.Printf("ssh: connection lost for %s, reconnecting", alias)
if reconnErr := client.Reconnect(ctx, m.getJumpClient(client.creds.Via)); reconnErr != nil {
m.mu.Lock()
m.reconnectFails[alias] = time.Now()
m.mu.Unlock()
return nil, fmt.Errorf("reconnect failed: %w", reconnErr)
}
// Clear backoff on success
m.mu.Lock()
delete(m.reconnectFails, alias)
m.mu.Unlock()
return client.Run(ctx, cmd)
}
return nil, err
}
return result, nil
}
// getJumpClient returns the jump client.
func (m *Manager) getJumpClient(via string) *Client {
if via == "" {
return nil
}
m.mu.RLock()
defer m.mu.RUnlock()
return m.connections[via]
}
// isConnectionError checks if error indicates lost connection.
func isConnectionError(err error) bool {
if err == nil {
return false
}
// Type-safe checks first
if errors.Is(err, io.EOF) {
return true
}
var netErr *net.OpError
if errors.As(err, &netErr) {
return true
}
// String matching for SSH-specific errors
errStr := err.Error()
return strings.Contains(errStr, "connection reset") ||
strings.Contains(errStr, "broken pipe") ||
strings.Contains(errStr, "connection refused") ||
strings.Contains(errStr, "use of closed network connection")
}
// Execute runs a command and returns formatted output.
func (m *Manager) Execute(ctx context.Context, cmd, target string) (string, error) {
result, err := m.Run(ctx, cmd, target)
if err != nil {
return "", err
}
var output strings.Builder
if result.Stdout != "" {
output.WriteString(result.Stdout)
}
if result.Stderr != "" {
if output.Len() > 0 {
output.WriteString("\n")
}
output.WriteString(result.Stderr)
}
if output.Len() == 0 {
return "(No output)", nil
}
if result.ExitCode != 0 {
fmt.Fprintf(&output, "\n[Exit Code: %d]", result.ExitCode)
}
// Truncate if too long
const maxBytes = 51200
outputStr := output.String()
if len(outputStr) > maxBytes {
outputStr = outputStr[:maxBytes] + "\n... [Output truncated]"
}
return outputStr, nil
}
// resolvePath resolves a path to an absolute path using the connection's CWD.
// No path restrictions — the connected user's OS permissions are the only boundary.
func (m *Manager) resolvePath(path, alias string) string {
m.mu.RLock()
client := m.connections[alias]
m.mu.RUnlock()
cwd := "/"
if client != nil {
cwd = client.CWD()
}
if !filepath.IsAbs(path) {
path = filepath.Join(cwd, path)
}
return filepath.Clean(path)
}
// ReadFile reads a file.
func (m *Manager) ReadFile(ctx context.Context, path, target string) (string, error) {
alias, err := m.resolveTarget(target)
if err != nil {
return "", err
}
resolved := m.resolvePath(path, alias)
lock := m.getAliasLock(alias)
lock.Lock()
defer lock.Unlock()
m.mu.RLock()
client := m.connections[alias]
m.mu.RUnlock()
if client == nil {
return "", fmt.Errorf("connection '%s' is no longer active", alias)
}
sftpClient, err := client.SFTP()
if err != nil {
return "", err
}
file, err := sftpClient.Open(resolved)
if err != nil {
return "", fmt.Errorf("failed to open file: %w", err)
}
defer file.Close()
// Check file size before reading to prevent OOM
const maxReadSize = 10 * 1024 * 1024 // 10 MB
stat, err := file.Stat()
if err != nil {
return "", fmt.Errorf("failed to stat file: %w", err)
}
if stat.Size() > maxReadSize {
return "", fmt.Errorf("file too large (%d bytes, max %d bytes); use 'run' with head/tail to read portions", stat.Size(), maxReadSize)
}
content, err := io.ReadAll(file)
if err != nil {
return "", fmt.Errorf("failed to read file: %w", err)
}
return string(content), nil
}
// WriteFile writes content to a file.
func (m *Manager) WriteFile(ctx context.Context, path, content, target string) error {
alias, err := m.resolveTarget(target)
if err != nil {
return err
}
resolved := m.resolvePath(path, alias)
lock := m.getAliasLock(alias)
lock.Lock()
defer lock.Unlock()
m.mu.RLock()
client := m.connections[alias]
m.mu.RUnlock()
if client == nil {
return fmt.Errorf("connection '%s' is no longer active", alias)
}
sftpClient, err := client.SFTP()
if err != nil {
return err
}
file, err := sftpClient.Create(resolved)
if err != nil {
return fmt.Errorf("failed to create file: %w", err)
}
defer file.Close()
_, err = file.Write([]byte(content))
if err != nil {
return fmt.Errorf("failed to write file: %w", err)
}
return nil
}
// ListDir lists directory contents.
func (m *Manager) ListDir(ctx context.Context, path, target string) ([]FileInfo, error) {
alias, err := m.resolveTarget(target)
if err != nil {
return nil, err
}
resolved := m.resolvePath(path, alias)
lock := m.getAliasLock(alias)
lock.Lock()
defer lock.Unlock()
m.mu.RLock()
client := m.connections[alias]
m.mu.RUnlock()
if client == nil {
return nil, fmt.Errorf("connection '%s' is no longer active", alias)
}
sftpClient, err := client.SFTP()
if err != nil {
return nil, err
}
entries, err := sftpClient.ReadDir(resolved)
if err != nil {
return nil, fmt.Errorf("failed to list directory: %w", err)
}
var files []FileInfo
for _, entry := range entries {
ftype := "file"
if entry.IsDir() {
ftype = "dir"
}
files = append(files, FileInfo{
Name: entry.Name(),
Type: ftype,
Size: entry.Size(),
Permissions: entry.Mode().String(),
})
}
return files, nil
}
// FileInfo represents file metadata.
type FileInfo struct {
Name string `json:"name"`
Type string `json:"type"`
Size int64 `json:"size"`
Permissions string `json:"permissions"`
}
// GetPublicKey returns the system's public SSH key.
func (m *Manager) GetPublicKey() (string, error) {
return m.keyManager.GetPublicKey()
}
// IsDockerAvailable checks if Docker is available on the target, with per-alias caching.
func (m *Manager) IsDockerAvailable(ctx context.Context, target string) (bool, error) {
alias, err := m.resolveTarget(target)
if err != nil {
return false, err
}
m.mu.RLock()
cached := m.dockerCache[alias]
m.mu.RUnlock()
if cached != nil {
return *cached, nil
}
output, err := m.Execute(ctx, "command -v docker >/dev/null 2>&1 && echo 'ok' || echo 'missing'", target)
if err != nil {
return false, err
}
available := strings.Contains(output, "ok")
m.mu.Lock()
m.dockerCache[alias] = &available
m.mu.Unlock()
return available, nil
}
// Close closes all connections.
func (m *Manager) Close() {
m.mu.Lock()
defer m.mu.Unlock()
for _, client := range m.connections {
if client != nil {
client.Close()
}
}
m.connections = make(map[string]*Client)
m.aliasLocks = make(map[string]*sync.Mutex)
m.reconnectFails = make(map[string]time.Time)
m.dockerCache = make(map[string]*bool)
m.primary = ""
}