Initial commit: forge-tools-ssh — MCP-сервер для администрирования по SSH
This commit is contained in:
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 = ""
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user