Initial commit: forge-tools-ssh — MCP-сервер для администрирования по SSH

This commit is contained in:
Maksim Totmin
2026-10-01 10:13:48 +07:00
commit 1821fe7968
30 changed files with 5342 additions and 0 deletions
+396
View File
@@ -0,0 +1,396 @@
package ssh
import (
"context"
"errors"
"fmt"
"io"
"log"
"net"
"strings"
"sync"
"time"
"github.com/pkg/sftp"
"golang.org/x/crypto/ssh"
)
// Client represents a single SSH connection with state tracking.
type Client struct {
alias string
conn *ssh.Client
sftp *sftp.Client
cwd string
mu sync.Mutex
creds Credentials
}
// Credentials holds SSH connection parameters.
type Credentials struct {
Host string
Port int
Username string
Password string
PrivateKey ssh.Signer
Via string
// Target is an optional destination host routed through a PAM gateway.
// When set, the SSH username sent to the gateway becomes
// "Username@Target" and authentication switches to multi-stage
// keyboard-interactive.
Target string
// TargetPassword is the password of the Target server account,
// requested by the PAM gateway as a second authentication stage.
TargetPassword string
}
// NewClient creates a new SSH client.
func NewClient(ctx context.Context, alias string, creds Credentials, jumpClient *Client) (*Client, error) {
client := &Client{
alias: alias,
creds: creds,
cwd: "",
}
if err := client.connect(ctx, jumpClient); err != nil {
return nil, err
}
return client, nil
}
// connect establishes the SSH connection. ctx отменяем: уважается и при
// TCP-dial, и при пробе «pwd» (иначе полу-готовый хост, принявший TCP, но
// не отдающий шелл, мог бы подвесить вызов без срока — см. исходный
// зависонный кейс).
func (c *Client) connect(ctx context.Context, jumpClient *Client) error {
c.mu.Lock()
defer c.mu.Unlock()
if ctx == nil {
ctx = context.Background()
}
// Close stale SFTP client before closing the connection
if c.sftp != nil {
c.sftp.Close()
c.sftp = nil
}
if c.conn != nil {
c.conn.Close()
c.conn = nil
}
// PAM gateway routing: the gateway parses "user@target" from the SSH
// username and proxies the session to the target host. The plain
// username is kept in Credentials for reconnect and logging.
user := c.creds.Username
if c.creds.Target != "" {
user = fmt.Sprintf("%s@%s", c.creds.Username, c.creds.Target)
}
config := &ssh.ClientConfig{
User: user,
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
Timeout: 30 * time.Second,
}
var authMethods []ssh.AuthMethod
if c.creds.PrivateKey != nil {
authMethods = append(authMethods, ssh.PublicKeys(c.creds.PrivateKey))
}
if c.creds.Target != "" {
// SafeInspect PAM authenticates in two stages:
// 1. keyboard-interactive with the PAM gateway password (the server
// replies with partial success),
// 2. the standard "password" method with the target server password.
// The Go SSH client tries each AuthMethod in order on partial success,
// so we supply both methods in the order above.
authMethods = append(authMethods, pamKeyboardInteractive(c.creds))
if c.creds.TargetPassword != "" {
authMethods = append(authMethods, ssh.Password(c.creds.TargetPassword))
}
} else if c.creds.Password != "" {
authMethods = append(authMethods, ssh.Password(c.creds.Password))
}
if len(authMethods) == 0 {
return errors.New("no authentication method provided (key or password required)")
}
config.Auth = authMethods
addr := fmt.Sprintf("%s:%d", c.creds.Host, c.creds.Port)
var conn *ssh.Client
var err error
if jumpClient != nil {
jumpConn := jumpClient.conn
if jumpConn == nil {
return errors.New("jump host not connected")
}
netConn, err := dialWithCtx(ctx, func() (net.Conn, error) {
return jumpConn.Dial("tcp", addr)
})
if err != nil {
return fmt.Errorf("failed to dial through jump host: %w", err)
}
ncc, chans, reqs, err := ssh.NewClientConn(netConn, addr, config)
if err != nil {
netConn.Close()
return fmt.Errorf("failed to create client connection through jump: %w", err)
}
conn = ssh.NewClient(ncc, chans, reqs)
} else {
// net.Dialer.DialContext вместо ssh.Dial: уважает ctx (отмену
// задачи), сохраняя 30-секундный лимит на установку соединения.
raw, err := (&net.Dialer{Timeout: 30 * time.Second}).DialContext(ctx, "tcp", addr)
if err != nil {
return fmt.Errorf("failed to connect: %w", err)
}
ncc, chans, reqs, err := ssh.NewClientConn(raw, addr, config)
if err != nil {
raw.Close()
return fmt.Errorf("failed to create client connection: %w", err)
}
conn = ssh.NewClient(ncc, chans, reqs)
}
c.conn = conn
// Проба «pwd» с дедлайном: соединение установлено, но хост может не
// отдать шелл - ограничиваем ожидание, чтобы не подвесить вызов.
probeCtx, probeCancel := context.WithTimeout(ctx, 15*time.Second)
defer probeCancel()
output, err := c.runRaw(probeCtx, "pwd")
if err != nil {
c.cwd = "~"
} else {
c.cwd = strings.TrimSpace(output)
}
log.Printf("ssh: connected %s@%s (%s)", c.creds.Username, c.creds.Host, c.alias)
return nil
}
// dialWithCtx выполняет сетевой dial с уважением к отмене ctx: если ctx
// отменён раньше, чем dial вернул соединение, канал разрывается. jump-
// диал не имеет ctx-версии, поэтому оборачиваем в select.
func dialWithCtx(ctx context.Context, dial func() (net.Conn, error)) (net.Conn, error) {
type res struct {
conn net.Conn
err error
}
ch := make(chan res, 1)
go func() {
c, err := dial()
ch <- res{c, err}
}()
select {
case <-ctx.Done():
return nil, ctx.Err()
case r := <-ch:
return r.conn, r.err
}
}
// pamKeyboardInteractive returns an ssh.AuthMethod for the first stage of a
// SafeInspect PAM gateway. The gateway presents a keyboard-interactive
// challenge asking for the PAM user's password; after the server signals
// partial success, the target server password is supplied by the following
// "password" AuthMethod. Challenge rounds that carry instructions only (no
// questions) are answered with an empty response.
func pamKeyboardInteractive(creds Credentials) ssh.AuthMethod {
return ssh.KeyboardInteractive(func(user, instruction string, questions []string, echos []bool) ([]string, error) {
if len(questions) == 0 {
return []string{}, nil
}
answers := make([]string, len(questions))
for i := range answers {
answers[i] = creds.Password
}
return answers, nil
})
}
// runRaw executes a command without CWD handling. Команда ограничена ctx
// (дедлайном): уважается при отмене задачи/превышении лимита, иначе
// полу-ответивший хост мог бы подвесить соединение без срока.
func (c *Client) runRaw(ctx context.Context, cmd string) (string, error) {
session, err := c.conn.NewSession()
if err != nil {
return "", fmt.Errorf("failed to create session: %w", err)
}
defer session.Close()
type out struct {
s string
err error
}
ch := make(chan out, 1)
go func() {
o, e := session.CombinedOutput(cmd)
ch <- out{string(o), e}
}()
select {
case <-ctx.Done():
_ = session.Close() // разрывает блокирующий вызов в горутине
return "", ctx.Err()
case r := <-ch:
return r.s, r.err
}
}
// Run executes a command with CWD tracking.
func (c *Client) Run(ctx context.Context, cmd string) (*RunResult, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.conn == nil {
return nil, errors.New("not connected")
}
delimiter := fmt.Sprintf("___MCP_PWD_%d___", time.Now().UnixNano())
wrappedCmd := fmt.Sprintf(
`cd %q && %s; __EXIT__=$?; echo ""; echo "%s"; pwd; exit $__EXIT__`,
c.cwd, strings.TrimRight(cmd, " \t\r\n"), delimiter,
)
session, err := c.conn.NewSession()
if err != nil {
return nil, fmt.Errorf("failed to create session: %w", err)
}
defer session.Close()
stdout, _ := session.StdoutPipe()
stderr, _ := session.StderrPipe()
if err := session.Start(wrappedCmd); err != nil {
return nil, fmt.Errorf("failed to start command: %w", err)
}
type readResult struct {
stdout []byte
stderr []byte
}
resultChan := make(chan readResult, 1)
const maxStdout = 10 * 1024 * 1024 // 10 MB
const maxStderr = 1 * 1024 * 1024 // 1 MB
go func() {
stdoutBytes, _ := io.ReadAll(io.LimitReader(stdout, maxStdout))
stderrBytes, _ := io.ReadAll(io.LimitReader(stderr, maxStderr))
resultChan <- readResult{stdout: stdoutBytes, stderr: stderrBytes}
}()
var res readResult
select {
case <-ctx.Done():
_ = session.Signal(ssh.SIGKILL)
_ = session.Close() // Unblock io.ReadAll by closing pipes
// Wait briefly for the reader goroutine to finish
select {
case <-resultChan:
case <-time.After(2 * time.Second):
}
return nil, ctx.Err()
case res = <-resultChan:
}
var exitCode int
if err := session.Wait(); err != nil {
if exitErr, ok := err.(*ssh.ExitError); ok {
exitCode = exitErr.ExitStatus()
} else {
return nil, fmt.Errorf("command failed: %w", err)
}
}
stdoutStr := string(res.stdout)
cleanOutput := stdoutStr
if idx := strings.Index(stdoutStr, delimiter); idx != -1 {
cleanOutput = stdoutStr[:idx]
remaining := strings.TrimSpace(stdoutStr[idx+len(delimiter):])
if remaining != "" {
c.cwd = remaining
}
}
return &RunResult{
Stdout: strings.TrimSpace(cleanOutput),
Stderr: strings.TrimSpace(string(res.stderr)),
ExitCode: exitCode,
CWD: c.cwd,
}, nil
}
// RunResult contains command execution result.
type RunResult struct {
Stdout string
Stderr string
ExitCode int
CWD string
}
// SFTP returns the SFTP client.
func (c *Client) SFTP() (*sftp.Client, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.conn == nil {
return nil, errors.New("not connected")
}
if c.sftp != nil {
return c.sftp, nil
}
sftpClient, err := sftp.NewClient(c.conn)
if err != nil {
return nil, fmt.Errorf("failed to create SFTP client: %w", err)
}
c.sftp = sftpClient
return c.sftp, nil
}
// Close closes the connection.
func (c *Client) Close() error {
c.mu.Lock()
defer c.mu.Unlock()
if c.sftp != nil {
c.sftp.Close()
c.sftp = nil
}
if c.conn != nil {
err := c.conn.Close()
c.conn = nil
return err
}
return nil
}
// Alias returns the connection alias.
func (c *Client) Alias() string {
return c.alias
}
// CWD returns the current working directory.
func (c *Client) CWD() string {
c.mu.Lock()
defer c.mu.Unlock()
return c.cwd
}
// Reconnect attempts to reconnect.
func (c *Client) Reconnect(ctx context.Context, jumpClient *Client) error {
log.Printf("ssh: reconnecting %s", c.alias)
return c.connect(ctx, jumpClient)
}
+139
View File
@@ -0,0 +1,139 @@
// Package ssh provides SSH connection management for the MCP server.
package ssh
import (
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/pem"
"fmt"
"log"
"os"
"path/filepath"
"golang.org/x/crypto/ssh"
)
const (
// DevKeyPath - локальный путь по умолчанию для запуска вне forge
// (без SSH_MCP_KEY_PATH). Производственный запуск всегда передаёт
// SSH_MCP_KEY_PATH от forge (см. cmd/serve/wiring.go): либо системный
// ключ data/_system/ssh/id_ed25519, либо пер-агентный ssh_key_dir.
DevKeyPath = "./data/_system/ssh/id_ed25519"
)
// KeyManager handles SSH key generation and loading.
type KeyManager struct {
keyPath string
}
// NewKeyManager creates a new KeyManager. If keyPath is empty, the system
// default is used (SSH_MCP_KEY_PATH env, otherwise local ./data/_system/ssh).
func NewKeyManager(keyPath string) *KeyManager {
if keyPath == "" {
keyPath = getDefaultKeyPath()
}
return &KeyManager{keyPath: keyPath}
}
// getDefaultKeyPath returns the appropriate path. Honors SSH_MCP_KEY_PATH if
// set (forge всегда так делает); иначе локальный дефолт под единой схемой
// данных (data/_system/ssh). Автодетект окружения убран: никаких жёстко
// зашитых /data - путь всегда приходит от оператора через env.
func getDefaultKeyPath() string {
if p := os.Getenv("SSH_MCP_KEY_PATH"); p != "" {
return p
}
return DevKeyPath
}
// EnsureKey ensures the system key exists, generating if necessary.
func (km *KeyManager) EnsureKey() error {
keyDir := filepath.Dir(km.keyPath)
// Check if directory exists
stat, err := os.Stat(keyDir)
if os.IsNotExist(err) {
// Directory doesn't exist - create it (0700; единственные личные данные).
if err := os.MkdirAll(keyDir, 0700); err != nil {
return fmt.Errorf("failed to create key directory %s: %w", keyDir, err)
}
log.Printf("ssh-key: created directory %s", keyDir)
} else if err != nil {
return fmt.Errorf("failed to access key directory %s: %w", keyDir, err)
} else if !stat.IsDir() {
return fmt.Errorf("key path %s exists but is not a directory", keyDir)
}
// Test write permissions by attempting to create a temp file
testFile := filepath.Join(keyDir, ".write_test")
if err := os.WriteFile(testFile, []byte("test"), 0600); err != nil {
return fmt.Errorf("key directory %s is not writable: %w", keyDir, err)
}
os.Remove(testFile)
if _, err := os.Stat(km.keyPath); os.IsNotExist(err) {
log.Printf("ssh-key: generating new key at %s", km.keyPath)
return km.generateKey()
}
return nil
}
// generateKey creates a new Ed25519 key pair.
func (km *KeyManager) generateKey() error {
pubKey, privKey, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
return fmt.Errorf("failed to generate key: %w", err)
}
privKeyBytes, err := ssh.MarshalPrivateKey(privKey, "ssh-mcp")
if err != nil {
return fmt.Errorf("failed to marshal private key: %w", err)
}
if err := os.WriteFile(km.keyPath, pem.EncodeToMemory(privKeyBytes), 0600); err != nil {
return fmt.Errorf("failed to write private key: %w", err)
}
sshPubKey, err := ssh.NewPublicKey(pubKey)
if err != nil {
return fmt.Errorf("failed to create SSH public key: %w", err)
}
// Add "SSH-MCP" comment to public key for identification
pubKeyBytes := []byte(fmt.Sprintf("%s %s SSH-MCP\n",
sshPubKey.Type(),
base64.StdEncoding.EncodeToString(sshPubKey.Marshal())))
if err := os.WriteFile(km.keyPath+".pub", pubKeyBytes, 0644); err != nil {
return fmt.Errorf("failed to write public key: %w", err)
}
log.Println("ssh-key: generated successfully")
return nil
}
// LoadPrivateKey loads the private key from disk.
func (km *KeyManager) LoadPrivateKey() (ssh.Signer, error) {
keyBytes, err := os.ReadFile(km.keyPath)
if err != nil {
return nil, fmt.Errorf("failed to read private key: %w", err)
}
signer, err := ssh.ParsePrivateKey(keyBytes)
if err != nil {
return nil, fmt.Errorf("failed to parse private key: %w", err)
}
return signer, nil
}
// GetPublicKey returns the public key string.
func (km *KeyManager) GetPublicKey() (string, error) {
pubKeyBytes, err := os.ReadFile(km.keyPath + ".pub")
if err != nil {
return "", fmt.Errorf("failed to read public key: %w", err)
}
return string(pubKeyBytes), nil
}
+710
View File
@@ -0,0 +1,710 @@
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 = ""
}
+171
View File
@@ -0,0 +1,171 @@
// Package ssh: per-tenant политика подключений (контракт PLAN.md
// "Per-tenant forge-tools"). Файл <FORGE_TENANT_CONFIG>/ssh.json декларирует
// именованные профили, на которые агент ссылается по алиасу, и allowlist
// хостов, валидирующий ad-hoc-подключения. Цель - "платформа объявляет,
// модель ссылается": LLM не должен сам выбирать host/ключ; это делает
// оператор, а агент выбирает только из разрешённого.
package ssh
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
"git.totmin.ru/en2zmax/forge-toolkit"
)
// Profile - одно именованное подключение, на которое агент ссылается по
// Alias. Значения Host/Username/KeyPath задаёт оператор; модель их не вводит.
type Profile struct {
Alias string `json:"alias"`
Host string `json:"host"`
Username string `json:"username"`
Port int `json:"port"`
// KeyPath - private key для подключения; может быть ${VAR} или пустым
// (тогда берётся DefaultKeyPath либо системный ключ KeyManager).
KeyPath string `json:"key"`
// Via - jump-host алиас для туннелирования; Target - PAM-шлюз.
Via string `json:"via"`
Target string `json:"target"`
TargetPass string `json:"target_password"`
}
// Policy - декларативная per-agent политика подключений.
type Policy struct {
Profiles []Profile `json:"profiles"`
// AllowedHosts - glob-паттерны разрешённых хостов ("10.0.*",
// "*.internal"). Пустой список = allow all (обратная совместимость с
// ad-hoc режимом без конфига). Наличие списка включает жалостную
// проверку: host вне списка - deny.
AllowedHosts []string `json:"allowed_hosts"`
// DefaultKeyPath - ключ по умолчанию, если ни профиль, ни вызов не
// указали private_key_path. Пусто = системный ключ (SSH_MCP_KEY_PATH).
DefaultKeyPath string `json:"default_key_path"`
byAlias map[string]Profile
hostGlob []*regexp.Regexp
}
// LoadPolicy читает <path>/ssh.json, раскрывает ${VAR} и валидирует.
// Отсутствие файла - nil, nil (политика не настроена -> ad-hoc без
// ограничений). Битый файл - ошибка (fail-closed).
func LoadPolicy(path string) (*Policy, error) {
data, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
return nil, fmt.Errorf("read ssh policy %s: %w", path, err)
}
return parsePolicy(data)
}
// parsePolicy разбирает готовые байты ssh.json (${VAR} + валидация).
// Используется и LoadPolicy, и configreload.Loader (live-reload, см.
// manager.go): loader передаёт сырые байты файла на каждый Get.
func parsePolicy(data []byte) (*Policy, error) {
p := &Policy{}
if err := json.Unmarshal(toolkit.Expand(data), p); err != nil {
return nil, fmt.Errorf("parse ssh policy: %w", err)
}
if err := p.normalize(); err != nil {
return nil, fmt.Errorf("ssh policy: %w", err)
}
return p, nil
}
func (p *Policy) normalize() error {
p.byAlias = make(map[string]Profile, len(p.Profiles))
for _, pr := range p.Profiles {
if pr.Alias == "" || pr.Host == "" || pr.Username == "" {
return fmt.Errorf("profile must have alias, host and username")
}
if _, dup := p.byAlias[pr.Alias]; dup {
return fmt.Errorf("duplicate profile alias %q", pr.Alias)
}
p.byAlias[pr.Alias] = pr
}
for _, h := range p.AllowedHosts {
re, err := globToRegex(h)
if err != nil {
return fmt.Errorf("allowed_hosts pattern %q: %w", h, err)
}
p.hostGlob = append(p.hostGlob, re)
}
return nil
}
// Resolve возвращает профиль по алиасу.
func (p *Policy) Resolve(alias string) (Profile, bool) {
if p == nil {
return Profile{}, false
}
pr, ok := p.byAlias[alias]
return pr, ok
}
// Allow сообщает, разрешён ли хост. Без allowed_hosts - всегда true
// (ad-hoc сохранён как раньше); иначе host обязан матчить хоть один glob.
func (p *Policy) Allow(host string) bool {
if p == nil || len(p.hostGlob) == 0 {
return true
}
for _, re := range p.hostGlob {
if re.MatchString(host) {
return true
}
}
return false
}
// KeyPath возвращает эффективный путь к ключу для профиля: из профиля,
// иначе DefaultKeyPath, иначе "" (системный ключ KeyManager).
func (p *Policy) KeyPath(pr Profile) string {
if pr.KeyPath != "" {
return pr.KeyPath
}
if p != nil {
return p.DefaultKeyPath
}
return ""
}
// globToRegex превращает glob-паттерн хоста в regex с полным совпадением.
func globToRegex(pattern string) (*regexp.Regexp, error) {
re, err := regexp.Compile("^" + globToRe(pattern) + "$")
if err != nil {
return nil, err
}
return re, nil
}
func globToRe(pattern string) string {
var b strings.Builder
prev := rune(0)
for _, ch := range pattern {
switch ch {
case '*':
if prev != '*' {
b.WriteString(".*")
}
case '?':
b.WriteByte('.')
default:
b.WriteString(regexp.QuoteMeta(string(ch)))
}
prev = ch
}
return b.String()
}
// PolicyPathFromEnv - путь к per-agent ssh.json из FORGE_TENANT_CONFIG.
func PolicyPathFromEnv() (string, bool) {
dir := os.Getenv("FORGE_TENANT_CONFIG")
if dir == "" {
return "", false
}
return filepath.Join(dir, "ssh.json"), true
}
+117
View File
@@ -0,0 +1,117 @@
package ssh
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
)
func writePolicy(t *testing.T, content string) string {
t.Helper()
dir := t.TempDir()
p := filepath.Join(dir, "ssh.json")
if err := os.WriteFile(p, []byte(content), 0o600); err != nil {
t.Fatalf("write policy: %v", err)
}
return p
}
func TestLoadPolicy_MissingFileNil(t *testing.T) {
p, err := LoadPolicy(filepath.Join(t.TempDir(), "nope.json"))
if err != nil {
t.Fatalf("missing file: %v", err)
}
if p != nil {
t.Fatalf("expected nil policy for missing file, got %+v", p)
}
}
func TestLoadPolicy_ExpandsVars(t *testing.T) {
t.Setenv("FORGE_SSH_KEY", "/keys/agent")
p, err := LoadPolicy(writePolicy(t, `{
"profiles":[{"alias":"prod","host":"10.0.1.5","username":"deploy","key":"${FORGE_SSH_KEY}"}],
"allowed_hosts":["10.0.*"],
"default_key_path":"${FORGE_SSH_KEY}"
}`))
if err != nil {
t.Fatalf("load: %v", err)
}
pr, ok := p.Resolve("prod")
if !ok {
t.Fatal("profile prod not resolved")
}
if pr.Host != "10.0.1.5" || pr.KeyPath != "/keys/agent" {
t.Fatalf("profile = %+v", pr)
}
if p.DefaultKeyPath != "/keys/agent" {
t.Fatalf("default_key = %q", p.DefaultKeyPath)
}
}
func TestPolicy_Allowlist(t *testing.T) {
p, err := LoadPolicy(writePolicy(t, `{"allowed_hosts":["10.0.*","*.internal"]}`))
if err != nil {
t.Fatalf("load: %v", err)
}
if !p.Allow("10.0.1.5") || !p.Allow("db.internal") {
t.Fatal("expected allowed hosts to pass")
}
if p.Allow("192.168.0.1") || p.Allow("evil.com") {
t.Fatal("expected non-allowed hosts to be denied")
}
}
func TestPolicy_NoAllowlistAllowsAll(t *testing.T) {
p, err := LoadPolicy(writePolicy(t, `{"profiles":[{"alias":"a","host":"h","username":"u"}]}`))
if err != nil {
t.Fatalf("load: %v", err)
}
if !p.Allow("anything.example") {
t.Fatal("empty allowlist should allow all (backward compat)")
}
}
// testManager изолирует генерацию ключа в temp-каталог, чтобы тесты не
// создавали ./data/id_ed25519 в модуле.
func testManager(t *testing.T) *Manager {
t.Helper()
t.Setenv("SSH_MCP_KEY_PATH", filepath.Join(t.TempDir(), "id_ed25519"))
return NewManager("")
}
func TestConnect_ProfileResolution_DeniesUnknown(t *testing.T) {
m := testManager(t)
m.SetPolicy(&Policy{
Profiles: []Profile{{Alias: "prod", Host: "10.0.1.5", Username: "deploy"}},
})
// Неизвестный профиль не должен уходить в сеть.
_, err := m.Connect(context.Background(), ConnectOptions{Profile: "nope"})
if err == nil || !strings.Contains(err.Error(), "not found") {
t.Fatalf("expected profile-not-found error, got %v", err)
}
}
func TestConnect_Allowlist_DeniesHost(t *testing.T) {
m := testManager(t)
m.SetPolicy(&Policy{AllowedHosts: []string{"10.0.*"}})
// Host вне allowlist - deny до dial.
_, err := m.Connect(context.Background(), ConnectOptions{Host: "evil.example", Username: "x"})
if err == nil || !strings.Contains(err.Error(), "not in allowed_hosts") {
t.Fatalf("expected allowlist deny, got %v", err)
}
}
func TestConnect_ProfileWithoutPolicy_Denies(t *testing.T) {
m := testManager(t)
_, err := m.Connect(context.Background(), ConnectOptions{Profile: "prod"})
if err == nil {
t.Fatal("expected error when profile used without policy")
}
if !strings.Contains(err.Error(), "per-agent policy") {
t.Fatalf("unexpected error: %v", err)
}
}