Initial commit: sshkeeper v0.1.0
Console SSH connection manager for Linux. Features: - TUI (Bubble Tea) with server list, add/edit form, test/save - CLI commands: add, list, show, edit, delete, connect, test, search, import, export, run, group, template, vault, ssh-config - Encrypted vault (Argon2id + XChaCha20-Poly1305) for passwords - PTY-wrapper for password auth - SQLite (modernc, no CGO) for server profiles - XDG-compatible paths - OpenSSH config generation - Import from ~/.ssh/config
This commit is contained in:
@@ -0,0 +1,106 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/BurntSushi/toml"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
SSH SSHConfig `toml:"ssh"`
|
||||
Vault VaultConfig `toml:"vault"`
|
||||
UI UIConfig `toml:"ui"`
|
||||
|
||||
// resolved paths
|
||||
ConfigDir string `toml:"-"`
|
||||
DataDir string `toml:"-"`
|
||||
}
|
||||
|
||||
type SSHConfig struct {
|
||||
Binary string `toml:"binary"`
|
||||
ConnectTimeoutSec int `toml:"connect_timeout_seconds"`
|
||||
TestCommand string `toml:"test_command"`
|
||||
}
|
||||
|
||||
type VaultConfig struct {
|
||||
AutoLockMinutes int `toml:"auto_lock_minutes"`
|
||||
}
|
||||
|
||||
type UIConfig struct {
|
||||
ShowSecurityHints bool `toml:"show_security_hints"`
|
||||
}
|
||||
|
||||
func defaultConfig() *Config {
|
||||
return &Config{
|
||||
SSH: SSHConfig{
|
||||
Binary: "/usr/bin/ssh",
|
||||
ConnectTimeoutSec: 10,
|
||||
TestCommand: "echo SSHKEEPER_OK",
|
||||
},
|
||||
Vault: VaultConfig{
|
||||
AutoLockMinutes: 15,
|
||||
},
|
||||
UI: UIConfig{
|
||||
ShowSecurityHints: false,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func Load() (*Config, error) {
|
||||
cfg := defaultConfig()
|
||||
|
||||
// XDG paths
|
||||
configDir := os.Getenv("XDG_CONFIG_HOME")
|
||||
if configDir == "" {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
configDir = filepath.Join(home, ".config", "sshkeeper")
|
||||
}
|
||||
cfg.ConfigDir = configDir
|
||||
|
||||
dataDir := os.Getenv("XDG_DATA_HOME")
|
||||
if dataDir == "" {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dataDir = filepath.Join(home, ".local", "share", "sshkeeper")
|
||||
}
|
||||
cfg.DataDir = dataDir
|
||||
|
||||
// Ensure dirs exist
|
||||
if err := os.MkdirAll(configDir, 0700); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.MkdirAll(dataDir, 0700); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
configFile := filepath.Join(configDir, "config.toml")
|
||||
|
||||
// Write default config if not exists
|
||||
if _, err := os.Stat(configFile); os.IsNotExist(err) {
|
||||
f, err := os.OpenFile(configFile, os.O_CREATE|os.O_WRONLY, 0600)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
if err := toml.NewEncoder(f).Encode(cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// Parse existing config
|
||||
if _, err := toml.DecodeFile(configFile, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Re-apply paths since toml decode might overwrite
|
||||
cfg.ConfigDir = configDir
|
||||
cfg.DataDir = dataDir
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package config
|
||||
|
||||
import "path/filepath"
|
||||
|
||||
func DBPath(dataDir string) string {
|
||||
return filepath.Join(dataDir, "sshkeeper.db")
|
||||
}
|
||||
|
||||
func VaultPath(dataDir string) string {
|
||||
return filepath.Join(dataDir, "vault.bin")
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"embed"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
var migrationsFS embed.FS
|
||||
|
||||
type DB struct {
|
||||
conn *sql.DB
|
||||
}
|
||||
|
||||
func Open(dataDir string) (*DB, error) {
|
||||
dbPath := filepath.Join(dataDir, "sshkeeper.db")
|
||||
|
||||
conn, err := sql.Open("sqlite", dbPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open database: %w", err)
|
||||
}
|
||||
|
||||
if err := conn.Ping(); err != nil {
|
||||
return nil, fmt.Errorf("ping database: %w", err)
|
||||
}
|
||||
|
||||
db := &DB{conn: conn}
|
||||
|
||||
if err := db.migrate(); err != nil {
|
||||
return nil, fmt.Errorf("migrate: %w", err)
|
||||
}
|
||||
|
||||
os.Chmod(dbPath, 0600)
|
||||
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func (db *DB) Close() error {
|
||||
return db.conn.Close()
|
||||
}
|
||||
|
||||
func (db *DB) migrate() error {
|
||||
entries, err := migrationsFS.ReadDir("migrations")
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migrations dir: %w", err)
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
content, err := migrationsFS.ReadFile("migrations/" + entry.Name())
|
||||
if err != nil {
|
||||
return fmt.Errorf("read migration %s: %w", entry.Name(), err)
|
||||
}
|
||||
if _, err := db.conn.Exec(string(content)); err != nil {
|
||||
return fmt.Errorf("exec migration %s: %w", entry.Name(), err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
CREATE TABLE IF NOT EXISTS servers (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
alias TEXT NOT NULL UNIQUE,
|
||||
display_name TEXT NOT NULL DEFAULT '',
|
||||
host TEXT NOT NULL,
|
||||
port INTEGER NOT NULL DEFAULT 22,
|
||||
user TEXT NOT NULL DEFAULT '',
|
||||
auth_method TEXT NOT NULL DEFAULT 'key',
|
||||
identity_file TEXT NOT NULL DEFAULT '',
|
||||
proxy_jump TEXT NOT NULL DEFAULT '',
|
||||
group_name TEXT NOT NULL DEFAULT '',
|
||||
notes TEXT NOT NULL DEFAULT '',
|
||||
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
last_connected_at DATETIME,
|
||||
last_test_at DATETIME,
|
||||
last_test_status TEXT NOT NULL DEFAULT 'unknown',
|
||||
last_test_error TEXT NOT NULL DEFAULT ''
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tags (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS server_tags (
|
||||
server_id INTEGER NOT NULL REFERENCES servers(id) ON DELETE CASCADE,
|
||||
tag_id INTEGER NOT NULL REFERENCES tags(id) ON DELETE CASCADE,
|
||||
PRIMARY KEY (server_id, tag_id)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS forwards (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
server_id INTEGER NOT NULL REFERENCES servers(id) ON DELETE CASCADE,
|
||||
type TEXT NOT NULL DEFAULT 'local',
|
||||
local_addr TEXT NOT NULL DEFAULT '',
|
||||
local_port INTEGER NOT NULL DEFAULT 0,
|
||||
remote_addr TEXT NOT NULL DEFAULT '',
|
||||
remote_port INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS command_templates (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
server_id INTEGER NOT NULL REFERENCES servers(id) ON DELETE CASCADE,
|
||||
name TEXT NOT NULL,
|
||||
command TEXT NOT NULL
|
||||
);
|
||||
@@ -0,0 +1,248 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
|
||||
"github.com/mirivlad/sshkeeper/internal/model"
|
||||
)
|
||||
|
||||
func (db *DB) CreateServer(s *model.Server) error {
|
||||
result, err := db.conn.Exec(`
|
||||
INSERT INTO servers (alias, display_name, host, port, user, auth_method, identity_file, proxy_jump, group_name, notes)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
s.Alias, s.DisplayName, s.Host, s.Port, s.User, s.AuthMethod, s.IdentityFile, s.ProxyJump, s.GroupName, s.Notes)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.ID, _ = result.LastInsertId()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (db *DB) UpdateServer(s *model.Server) error {
|
||||
_, err := db.conn.Exec(`
|
||||
UPDATE servers SET
|
||||
display_name=?, host=?, port=?, user=?, auth_method=?,
|
||||
identity_file=?, proxy_jump=?, group_name=?, notes=?, updated_at=CURRENT_TIMESTAMP
|
||||
WHERE alias=?`,
|
||||
s.DisplayName, s.Host, s.Port, s.User, s.AuthMethod,
|
||||
s.IdentityFile, s.ProxyJump, s.GroupName, s.Notes, s.Alias)
|
||||
return err
|
||||
}
|
||||
|
||||
func (db *DB) DeleteServer(alias string) error {
|
||||
_, err := db.conn.Exec("DELETE FROM servers WHERE alias=?", alias)
|
||||
return err
|
||||
}
|
||||
|
||||
func (db *DB) GetServer(alias string) (*model.Server, error) {
|
||||
var s model.Server
|
||||
var lastConnected, lastTest sql.NullTime
|
||||
err := db.conn.QueryRow(`
|
||||
SELECT id, alias, display_name, host, port, user, auth_method,
|
||||
identity_file, proxy_jump, group_name, notes,
|
||||
created_at, updated_at, last_connected_at,
|
||||
last_test_at, last_test_status, last_test_error
|
||||
FROM servers WHERE alias=?`, alias).Scan(
|
||||
&s.ID, &s.Alias, &s.DisplayName, &s.Host, &s.Port, &s.User, &s.AuthMethod,
|
||||
&s.IdentityFile, &s.ProxyJump, &s.GroupName, &s.Notes,
|
||||
&s.CreatedAt, &s.UpdatedAt, &lastConnected,
|
||||
&lastTest, &s.LastTestStatus, &s.LastTestError)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if lastConnected.Valid {
|
||||
s.LastConnectedAt = &lastConnected.Time
|
||||
}
|
||||
if lastTest.Valid {
|
||||
s.LastTestAt = &lastTest.Time
|
||||
}
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
func (db *DB) ListServers() ([]*model.Server, error) {
|
||||
rows, err := db.conn.Query(`
|
||||
SELECT id, alias, display_name, host, port, user, auth_method,
|
||||
identity_file, proxy_jump, group_name, notes,
|
||||
created_at, updated_at, last_connected_at,
|
||||
last_test_at, last_test_status, last_test_error
|
||||
FROM servers ORDER BY alias`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var servers []*model.Server
|
||||
for rows.Next() {
|
||||
var s model.Server
|
||||
var lastConnected, lastTest sql.NullTime
|
||||
err := rows.Scan(
|
||||
&s.ID, &s.Alias, &s.DisplayName, &s.Host, &s.Port, &s.User, &s.AuthMethod,
|
||||
&s.IdentityFile, &s.ProxyJump, &s.GroupName, &s.Notes,
|
||||
&s.CreatedAt, &s.UpdatedAt, &lastConnected,
|
||||
&lastTest, &s.LastTestStatus, &s.LastTestError)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if lastConnected.Valid {
|
||||
s.LastConnectedAt = &lastConnected.Time
|
||||
}
|
||||
if lastTest.Valid {
|
||||
s.LastTestAt = &lastTest.Time
|
||||
}
|
||||
servers = append(servers, &s)
|
||||
}
|
||||
return servers, rows.Err()
|
||||
}
|
||||
|
||||
func (db *DB) SearchServers(query string) ([]*model.Server, error) {
|
||||
pattern := "%" + query + "%"
|
||||
rows, err := db.conn.Query(`
|
||||
SELECT id, alias, display_name, host, port, user, auth_method,
|
||||
identity_file, proxy_jump, group_name, notes,
|
||||
created_at, updated_at, last_connected_at,
|
||||
last_test_at, last_test_status, last_test_error
|
||||
FROM servers
|
||||
WHERE alias LIKE ? OR display_name LIKE ? OR host LIKE ? OR user LIKE ? OR group_name LIKE ?
|
||||
ORDER BY alias`, pattern, pattern, pattern, pattern, pattern)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var servers []*model.Server
|
||||
for rows.Next() {
|
||||
var s model.Server
|
||||
var lastConnected, lastTest sql.NullTime
|
||||
err := rows.Scan(
|
||||
&s.ID, &s.Alias, &s.DisplayName, &s.Host, &s.Port, &s.User, &s.AuthMethod,
|
||||
&s.IdentityFile, &s.ProxyJump, &s.GroupName, &s.Notes,
|
||||
&s.CreatedAt, &s.UpdatedAt, &lastConnected,
|
||||
&lastTest, &s.LastTestStatus, &s.LastTestError)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if lastConnected.Valid {
|
||||
s.LastConnectedAt = &lastConnected.Time
|
||||
}
|
||||
if lastTest.Valid {
|
||||
s.LastTestAt = &lastTest.Time
|
||||
}
|
||||
servers = append(servers, &s)
|
||||
}
|
||||
return servers, rows.Err()
|
||||
}
|
||||
|
||||
func (db *DB) UpdateTestResult(alias string, status model.TestStatus, testErr string) error {
|
||||
_, err := db.conn.Exec(`
|
||||
UPDATE servers SET last_test_at=CURRENT_TIMESTAMP, last_test_status=?, last_test_error=?
|
||||
WHERE alias=?`, status, testErr, alias)
|
||||
return err
|
||||
}
|
||||
|
||||
func (db *DB) UpdateLastConnected(alias string) error {
|
||||
_, err := db.conn.Exec("UPDATE servers SET last_connected_at=CURRENT_TIMESTAMP WHERE alias=?", alias)
|
||||
return err
|
||||
}
|
||||
|
||||
// Tag methods
|
||||
func (db *DB) AddTagToServer(serverID int64, tagName string) error {
|
||||
var tagID int64
|
||||
err := db.conn.QueryRow("SELECT id FROM tags WHERE name=?", tagName).Scan(&tagID)
|
||||
if err == sql.ErrNoRows {
|
||||
result, err := db.conn.Exec("INSERT INTO tags (name) VALUES (?)", tagName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tagID, _ = result.LastInsertId()
|
||||
} else if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = db.conn.Exec("INSERT OR IGNORE INTO server_tags (server_id, tag_id) VALUES (?, ?)", serverID, tagID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (db *DB) GetServerTags(serverID int64) ([]string, error) {
|
||||
rows, err := db.conn.Query(`
|
||||
SELECT t.name FROM tags t
|
||||
JOIN server_tags st ON st.tag_id = t.id
|
||||
WHERE st.server_id = ?
|
||||
ORDER BY t.name`, serverID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var tags []string
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tags = append(tags, name)
|
||||
}
|
||||
return tags, rows.Err()
|
||||
}
|
||||
|
||||
// Forward methods
|
||||
func (db *DB) AddForward(serverID int64, fwdType model.ForwardType, localAddr string, localPort int, remoteAddr string, remotePort int) error {
|
||||
_, err := db.conn.Exec(`
|
||||
INSERT INTO forwards (server_id, type, local_addr, local_port, remote_addr, remote_port)
|
||||
VALUES (?, ?, ?, ?, ?, ?)`,
|
||||
serverID, fwdType, localAddr, localPort, remoteAddr, remotePort)
|
||||
return err
|
||||
}
|
||||
|
||||
func (db *DB) GetForwards(serverID int64) ([]*model.Forward, error) {
|
||||
rows, err := db.conn.Query(`
|
||||
SELECT id, server_id, type, local_addr, local_port, remote_addr, remote_port
|
||||
FROM forwards WHERE server_id=?`, serverID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var forwards []*model.Forward
|
||||
for rows.Next() {
|
||||
var f model.Forward
|
||||
if err := rows.Scan(&f.ID, &f.ServerID, &f.Type, &f.LocalAddr, &f.LocalPort, &f.RemoteAddr, &f.RemotePort); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
forwards = append(forwards, &f)
|
||||
}
|
||||
return forwards, rows.Err()
|
||||
}
|
||||
|
||||
// Ensure time import is used
|
||||
var _ time.Time
|
||||
|
||||
// Command template methods
|
||||
func (db *DB) AddCommandTemplate(serverID int64, name, command string) error {
|
||||
_, err := db.conn.Exec(
|
||||
"INSERT INTO command_templates (server_id, name, command) VALUES (?, ?, ?)",
|
||||
serverID, name, command)
|
||||
return err
|
||||
}
|
||||
|
||||
func (db *DB) GetCommandTemplates(serverAlias string) ([]*model.CommandTemplate, error) {
|
||||
rows, err := db.conn.Query(`
|
||||
SELECT ct.id, ct.server_id, ct.name, ct.command
|
||||
FROM command_templates ct
|
||||
JOIN servers s ON s.id = ct.server_id
|
||||
WHERE s.alias = ?
|
||||
ORDER BY ct.name`, serverAlias)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var templates []*model.CommandTemplate
|
||||
for rows.Next() {
|
||||
var t model.CommandTemplate
|
||||
if err := rows.Scan(&t.ID, &t.ServerID, &t.Name, &t.Command); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
templates = append(templates, &t)
|
||||
}
|
||||
return templates, rows.Err()
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type AuthMethod string
|
||||
|
||||
const (
|
||||
AuthPassword AuthMethod = "password"
|
||||
AuthKey AuthMethod = "key"
|
||||
AuthKeyPassphrase AuthMethod = "key_passphrase"
|
||||
AuthAgent AuthMethod = "agent"
|
||||
)
|
||||
|
||||
type TestStatus string
|
||||
|
||||
const (
|
||||
TestUnknown TestStatus = "unknown"
|
||||
TestOK TestStatus = "ok"
|
||||
TestFailed TestStatus = "failed"
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
ID int64 `json:"id"`
|
||||
Alias string `json:"alias"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Host string `json:"host"`
|
||||
Port int `json:"port"`
|
||||
User string `json:"user"`
|
||||
AuthMethod AuthMethod `json:"auth_method"`
|
||||
IdentityFile string `json:"identity_file"`
|
||||
ProxyJump string `json:"proxy_jump"`
|
||||
GroupName string `json:"group_name"`
|
||||
Notes string `json:"notes"`
|
||||
Tags []string `json:"tags"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
LastConnectedAt *time.Time `json:"last_connected_at"`
|
||||
LastTestAt *time.Time `json:"last_test_at"`
|
||||
LastTestStatus TestStatus `json:"last_test_status"`
|
||||
LastTestError string `json:"last_test_error"`
|
||||
}
|
||||
|
||||
type SecretType string
|
||||
|
||||
const (
|
||||
SecretSSHPassword SecretType = "ssh_password"
|
||||
SecretKeyPassphrase SecretType = "key_passphrase"
|
||||
SecretSudoPassword SecretType = "sudo_password"
|
||||
SecretCustom SecretType = "custom_secret"
|
||||
)
|
||||
|
||||
type Secret struct {
|
||||
ID string `json:"id"`
|
||||
Type SecretType `json:"type"`
|
||||
Nonce []byte `json:"nonce"`
|
||||
Data []byte `json:"data"`
|
||||
}
|
||||
|
||||
type ForwardType string
|
||||
|
||||
const (
|
||||
ForwardLocal ForwardType = "local"
|
||||
ForwardRemote ForwardType = "remote"
|
||||
ForwardDynamic ForwardType = "dynamic"
|
||||
)
|
||||
|
||||
type Forward struct {
|
||||
ID int64 `json:"id"`
|
||||
ServerID int64 `json:"server_id"`
|
||||
Type ForwardType `json:"type"`
|
||||
LocalAddr string `json:"local_addr"`
|
||||
LocalPort int `json:"local_port"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
RemotePort int `json:"remote_port"`
|
||||
}
|
||||
|
||||
type Tag struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
type CommandTemplate struct {
|
||||
ID int64 `json:"id"`
|
||||
ServerID int64 `json:"server_id"`
|
||||
Name string `json:"name"`
|
||||
Command string `json:"command"`
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package ssh
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mirivlad/sshkeeper/internal/config"
|
||||
"github.com/mirivlad/sshkeeper/internal/model"
|
||||
)
|
||||
|
||||
type VaultFunc func(serverAlias string, secretType string) (string, error)
|
||||
|
||||
func Connect(cfg *config.Config, server *model.Server, getVault VaultFunc) error {
|
||||
args := buildArgs(server)
|
||||
|
||||
switch server.AuthMethod {
|
||||
case model.AuthPassword:
|
||||
password, err := getVault(server.Alias, "ssh_password")
|
||||
if err != nil {
|
||||
return fmt.Errorf("get password from vault: %w", err)
|
||||
}
|
||||
return ConnectWithPassword(cfg.SSH.Binary, args, password)
|
||||
|
||||
case model.AuthKeyPassphrase:
|
||||
// For key+passphrase, we need to handle the passphrase
|
||||
// For now, let ssh-agent handle it or prompt normally
|
||||
// TODO: use ssh-agent or similar
|
||||
fallthrough
|
||||
|
||||
default:
|
||||
// key, agent, key+passphrase - direct execution
|
||||
cmd := exec.Command(cfg.SSH.Binary, args...)
|
||||
cmd.Stdin = os.Stdin
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
return fmt.Errorf("start ssh: %w", err)
|
||||
}
|
||||
|
||||
return cmd.Wait()
|
||||
}
|
||||
}
|
||||
|
||||
func Test(cfg *config.Config, server *model.Server, getVault VaultFunc) (bool, string) {
|
||||
args := buildArgs(server)
|
||||
args = append(args, "-o", fmt.Sprintf("ConnectTimeout=%d", cfg.SSH.ConnectTimeoutSec))
|
||||
|
||||
switch server.AuthMethod {
|
||||
case model.AuthPassword:
|
||||
// For password auth, we can't use BatchMode
|
||||
// Use a short timeout and try to connect
|
||||
args = append(args, "-o", "NumberOfPasswordPrompts=1")
|
||||
password, err := getVault(server.Alias, "ssh_password")
|
||||
if err != nil {
|
||||
return false, fmt.Sprintf("vault error: %v", err)
|
||||
}
|
||||
return testWithPassword(cfg, args, password)
|
||||
|
||||
default:
|
||||
// key, agent, key+passphrase
|
||||
args = append(args, "-o", "BatchMode=yes")
|
||||
args = append(args, cfg.SSH.TestCommand)
|
||||
|
||||
cmd := exec.Command(cfg.SSH.Binary, args...)
|
||||
cmd.Stdin = nil
|
||||
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return false, strings.TrimSpace(string(output))
|
||||
}
|
||||
|
||||
result := strings.TrimSpace(string(output))
|
||||
if result == "SSHKEEPER_OK" {
|
||||
return true, ""
|
||||
}
|
||||
return false, result
|
||||
}
|
||||
}
|
||||
|
||||
func testWithPassword(cfg *config.Config, args []string, password string) (bool, string) {
|
||||
// For password test, we use PTY approach with a short timeout
|
||||
// This is a simplified version - in production, use ConnectWithPassword
|
||||
// with a test command
|
||||
args = append(args, cfg.SSH.TestCommand)
|
||||
|
||||
cmd := exec.Command(cfg.SSH.Binary, args...)
|
||||
cmd.Stdin = nil
|
||||
cmd.Stdout = nil
|
||||
cmd.Stderr = nil
|
||||
|
||||
// Use a timeout
|
||||
done := make(chan error, 1)
|
||||
if err := cmd.Start(); err != nil {
|
||||
return false, err.Error()
|
||||
}
|
||||
|
||||
go func() {
|
||||
done <- cmd.Wait()
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
return false, err.Error()
|
||||
}
|
||||
return true, ""
|
||||
case <-time.After(time.Duration(cfg.SSH.ConnectTimeoutSec) * time.Second):
|
||||
cmd.Process.Kill()
|
||||
return false, "connection timeout"
|
||||
}
|
||||
}
|
||||
|
||||
func buildArgs(server *model.Server) []string {
|
||||
var args []string
|
||||
|
||||
args = append(args, "-p", fmt.Sprintf("%d", server.Port))
|
||||
|
||||
if server.IdentityFile != "" {
|
||||
args = append(args, "-i", server.IdentityFile)
|
||||
}
|
||||
|
||||
if server.ProxyJump != "" {
|
||||
args = append(args, "-J", server.ProxyJump)
|
||||
}
|
||||
|
||||
// Disable strict host key checking for first connection
|
||||
// In production, this should be configurable
|
||||
args = append(args, "-o", "StrictHostKeyChecking=accept-new")
|
||||
|
||||
target := fmt.Sprintf("%s@%s", server.User, server.Host)
|
||||
args = append(args, target)
|
||||
|
||||
return args
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
package ssh
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mirivlad/sshkeeper/internal/model"
|
||||
)
|
||||
|
||||
// GenerateConfig creates OpenSSH config content from server profiles
|
||||
func GenerateConfig(servers []*model.Server) (string, error) {
|
||||
var sb strings.Builder
|
||||
|
||||
sb.WriteString("# Generated by sshkeeper. Do not edit manually.\n")
|
||||
sb.WriteString(fmt.Sprintf("# Generated at: %s\n\n", time.Now().Format(time.RFC3339)))
|
||||
|
||||
for _, s := range servers {
|
||||
sb.WriteString(fmt.Sprintf("Host %s\n", s.Alias))
|
||||
sb.WriteString(fmt.Sprintf(" HostName %s\n", s.Host))
|
||||
if s.Port != 22 {
|
||||
sb.WriteString(fmt.Sprintf(" Port %d\n", s.Port))
|
||||
}
|
||||
if s.User != "" {
|
||||
sb.WriteString(fmt.Sprintf(" User %s\n", s.User))
|
||||
}
|
||||
if s.IdentityFile != "" && s.AuthMethod != model.AuthPassword {
|
||||
sb.WriteString(fmt.Sprintf(" IdentityFile %s\n", s.IdentityFile))
|
||||
}
|
||||
if s.ProxyJump != "" {
|
||||
sb.WriteString(fmt.Sprintf(" ProxyJump %s\n", s.ProxyJump))
|
||||
}
|
||||
sb.WriteString("\n")
|
||||
}
|
||||
|
||||
return sb.String(), nil
|
||||
}
|
||||
|
||||
// WriteConfig writes the generated config to ~/.ssh/config.d/sshkeeper.conf
|
||||
func WriteConfig(servers []*model.Server) error {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
configD := home + "/.ssh/config.d"
|
||||
if err := os.MkdirAll(configD, 0700); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
content, err := GenerateConfig(servers)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
configFile := configD + "/sshkeeper.conf"
|
||||
tmpFile := configFile + ".tmp"
|
||||
|
||||
if err := os.WriteFile(tmpFile, []byte(content), 0600); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return os.Rename(tmpFile, configFile)
|
||||
}
|
||||
|
||||
// InstallInclude adds "Include ~/.ssh/config.d/*.conf" to ~/.ssh/config
|
||||
func InstallInclude() error {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
sshDir := home + "/.ssh"
|
||||
configD := sshDir + "/config.d"
|
||||
mainConfig := sshDir + "/config"
|
||||
|
||||
if err := os.MkdirAll(configD, 0700); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
includeLine := "Include ~/.ssh/config.d/*.conf"
|
||||
|
||||
// Check if already included
|
||||
if data, err := os.ReadFile(mainConfig); err == nil {
|
||||
if strings.Contains(string(data), "Include ~/.ssh/config.d") {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// Prepend include line
|
||||
f, err := os.OpenFile(mainConfig, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
_, err = f.WriteString("\n" + includeLine + "\n")
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package ssh
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/mirivlad/sshkeeper/internal/model"
|
||||
)
|
||||
|
||||
// ImportFromSSHConfig parses ~/.ssh/config and returns server profiles
|
||||
func ImportFromSSHConfig() ([]*model.Server, error) {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
configPath := filepath.Join(home, ".ssh", "config")
|
||||
f, err := os.Open(configPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, fmt.Errorf("~/.ssh/config not found")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
var servers []*model.Server
|
||||
var current *model.Server
|
||||
|
||||
scanner := bufio.NewScanner(f)
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" || strings.HasPrefix(line, "#") {
|
||||
continue
|
||||
}
|
||||
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) < 2 {
|
||||
continue
|
||||
}
|
||||
|
||||
key := strings.ToLower(fields[0])
|
||||
value := strings.Join(fields[1:], " ")
|
||||
|
||||
switch key {
|
||||
case "host":
|
||||
if current != nil && current.Host != "" {
|
||||
servers = append(servers, current)
|
||||
}
|
||||
// Skip wildcard hosts and patterns
|
||||
if strings.Contains(value, "*") || strings.Contains(value, "?") {
|
||||
current = nil
|
||||
continue
|
||||
}
|
||||
current = &model.Server{
|
||||
Alias: value,
|
||||
Host: value,
|
||||
Port: 22,
|
||||
User: "",
|
||||
AuthMethod: model.AuthKey,
|
||||
}
|
||||
|
||||
case "hostname":
|
||||
if current != nil {
|
||||
current.Host = value
|
||||
}
|
||||
|
||||
case "port":
|
||||
if current != nil {
|
||||
if port, err := strconv.Atoi(value); err == nil {
|
||||
current.Port = port
|
||||
}
|
||||
}
|
||||
|
||||
case "user":
|
||||
if current != nil {
|
||||
current.User = value
|
||||
}
|
||||
|
||||
case "identityfile":
|
||||
if current != nil {
|
||||
current.IdentityFile = value
|
||||
if current.AuthMethod == model.AuthKey {
|
||||
current.AuthMethod = model.AuthKey
|
||||
}
|
||||
}
|
||||
|
||||
case "proxyjump":
|
||||
if current != nil {
|
||||
current.ProxyJump = value
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Don't forget the last host
|
||||
if current != nil && current.Host != "" {
|
||||
servers = append(servers, current)
|
||||
}
|
||||
|
||||
return servers, scanner.Err()
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
package ssh
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"regexp"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/creack/pty"
|
||||
"golang.org/x/term"
|
||||
)
|
||||
|
||||
var passwordPromptRe = regexp.MustCompile(`(?i)(password|passphrase).*:\s*$`)
|
||||
|
||||
// ConnectWithPassword runs SSH through a PTY, detects the password prompt,
|
||||
// sends the password, and then bridges the user terminal to the SSH session.
|
||||
func ConnectWithPassword(sshBinary string, args []string, password string) error {
|
||||
// Start SSH with PTY
|
||||
cmd := exec.Command(sshBinary, args...)
|
||||
cmd.Env = os.Environ()
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{
|
||||
Setsid: true,
|
||||
Setctty: true,
|
||||
}
|
||||
|
||||
ptmx, err := pty.Start(cmd)
|
||||
if err != nil {
|
||||
return fmt.Errorf("start ssh with pty: %w", err)
|
||||
}
|
||||
defer ptmx.Close()
|
||||
|
||||
// Save terminal state and set to raw
|
||||
oldState, err := term.MakeRaw(int(os.Stdin.Fd()))
|
||||
if err != nil {
|
||||
return fmt.Errorf("set raw terminal: %w", err)
|
||||
}
|
||||
defer term.Restore(int(os.Stdin.Fd()), oldState)
|
||||
|
||||
// Channel to signal when password has been sent
|
||||
passwordSent := make(chan bool, 1)
|
||||
done := make(chan error, 1)
|
||||
|
||||
// Read from PTY, detect password prompt
|
||||
go func() {
|
||||
buf := make([]byte, 4096)
|
||||
var accumulated strings.Builder
|
||||
|
||||
for {
|
||||
n, err := ptmx.Read(buf)
|
||||
if n > 0 {
|
||||
data := buf[:n]
|
||||
accumulated.Write(data)
|
||||
|
||||
// Write to stdout
|
||||
os.Stdout.Write(data)
|
||||
|
||||
// Check for password prompt
|
||||
if !<-passwordSent {
|
||||
text := accumulated.String()
|
||||
if passwordPromptRe.MatchString(text) {
|
||||
passwordSent <- true
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
ptmx.Write([]byte(password + "\r"))
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
// Reset accumulated buffer periodically to avoid unbounded growth
|
||||
if accumulated.Len() > 8192 {
|
||||
s := accumulated.String()
|
||||
accumulated.Reset()
|
||||
accumulated.WriteString(s[len(s)-2048:])
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if err != io.EOF {
|
||||
done <- err
|
||||
} else {
|
||||
done <- nil
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// Copy stdin to PTY
|
||||
go func() {
|
||||
io.Copy(ptmx, os.Stdin)
|
||||
}()
|
||||
|
||||
// Wait for command completion
|
||||
err = cmd.Wait()
|
||||
passwordSent <- false // signal to stop
|
||||
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,680 @@
|
||||
package tui
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/charmbracelet/bubbles/list"
|
||||
"github.com/charmbracelet/bubbles/spinner"
|
||||
"github.com/charmbracelet/bubbles/textinput"
|
||||
"github.com/charmbracelet/bubbletea"
|
||||
"github.com/charmbracelet/lipgloss"
|
||||
"github.com/mirivlad/sshkeeper/internal/model"
|
||||
)
|
||||
|
||||
// --- Styles ---
|
||||
|
||||
var (
|
||||
titleStyle = lipgloss.NewStyle().
|
||||
Bold(true).
|
||||
Foreground(lipgloss.Color("12")).
|
||||
MarginLeft(2)
|
||||
|
||||
selectedStyle = lipgloss.NewStyle().
|
||||
Foreground(lipgloss.Color("15")).
|
||||
Background(lipgloss.Color("4")).
|
||||
Bold(true)
|
||||
|
||||
normalStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("15"))
|
||||
|
||||
testOKStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("10")).Bold(true)
|
||||
testFailStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("9")).Bold(true)
|
||||
|
||||
helpStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("14")).MarginLeft(2)
|
||||
|
||||
errorStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("9")).Bold(true)
|
||||
successStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("10")).Bold(true)
|
||||
|
||||
focusedStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("12")).Bold(true)
|
||||
blurredStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("7"))
|
||||
)
|
||||
|
||||
// --- Messages ---
|
||||
|
||||
type serversLoadedMsg struct {
|
||||
servers []*model.Server
|
||||
err error
|
||||
}
|
||||
|
||||
type testDoneMsg struct {
|
||||
ok bool
|
||||
err string
|
||||
}
|
||||
|
||||
type saveDoneMsg struct {
|
||||
err error
|
||||
}
|
||||
|
||||
// connectRequestMsg — TUI requests a connect action to be handled outside
|
||||
type connectRequestMsg struct {
|
||||
server *model.Server
|
||||
}
|
||||
|
||||
// --- Server list item ---
|
||||
|
||||
type serverItem struct {
|
||||
server *model.Server
|
||||
}
|
||||
|
||||
func (i serverItem) Title() string { return i.server.Alias }
|
||||
func (i serverItem) Description() string { return fmt.Sprintf("%s@%s:%d %s", i.server.User, i.server.Host, i.server.Port, i.server.AuthMethod) }
|
||||
func (i serverItem) FilterValue() string { return i.server.Alias + " " + i.server.DisplayName + " " + i.server.Host + " " + i.server.User }
|
||||
|
||||
// --- External callbacks ---
|
||||
|
||||
var (
|
||||
ListServers func() ([]*model.Server, error)
|
||||
SearchServers func(query string) ([]*model.Server, error)
|
||||
DeleteServer func(alias string) error
|
||||
TestConnection func(server *model.Server) (bool, string)
|
||||
SaveServer func(server *model.Server, password string) error
|
||||
)
|
||||
|
||||
// --- Screen type ---
|
||||
|
||||
type screen int
|
||||
|
||||
const (
|
||||
screenList screen = iota
|
||||
screenForm
|
||||
screenSearch
|
||||
)
|
||||
|
||||
// --- Result type — returned from TUI to caller ---
|
||||
|
||||
type TUIResult struct {
|
||||
Server *model.Server
|
||||
Action string // "connect"
|
||||
}
|
||||
|
||||
// --- Main TUI model ---
|
||||
|
||||
type tuiModel struct {
|
||||
screen screen
|
||||
list list.Model
|
||||
servers []*model.Server
|
||||
searchInput textinput.Model
|
||||
form *formModel
|
||||
err error
|
||||
success string
|
||||
width int
|
||||
height int
|
||||
result *TUIResult
|
||||
}
|
||||
|
||||
func New(servers []*model.Server) *tuiModel {
|
||||
items := make([]list.Item, len(servers))
|
||||
for i, s := range servers {
|
||||
items[i] = serverItem{server: s}
|
||||
}
|
||||
|
||||
l := list.New(items, list.NewDefaultDelegate(), 0, 0)
|
||||
l.Title = "sshkeeper"
|
||||
l.SetShowStatusBar(false)
|
||||
l.SetFilteringEnabled(false)
|
||||
l.Styles.Title = titleStyle
|
||||
|
||||
search := textinput.New()
|
||||
search.Placeholder = "Search..."
|
||||
search.CharLimit = 64
|
||||
|
||||
return &tuiModel{
|
||||
screen: screenList,
|
||||
list: l,
|
||||
servers: servers,
|
||||
searchInput: search,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *tuiModel) Result() *TUIResult {
|
||||
return m.result
|
||||
}
|
||||
|
||||
func (m *tuiModel) Init() tea.Cmd {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *tuiModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
switch msg := msg.(type) {
|
||||
|
||||
case tea.WindowSizeMsg:
|
||||
m.width = msg.Width
|
||||
m.height = msg.Height
|
||||
m.list.SetSize(msg.Width, msg.Height-4)
|
||||
if m.form != nil {
|
||||
m.form.width = msg.Width
|
||||
m.form.height = msg.Height
|
||||
}
|
||||
return m, nil
|
||||
|
||||
case serversLoadedMsg:
|
||||
if msg.err != nil {
|
||||
m.err = msg.err
|
||||
} else {
|
||||
m.servers = msg.servers
|
||||
items := make([]list.Item, len(msg.servers))
|
||||
for i, s := range msg.servers {
|
||||
items[i] = serverItem{server: s}
|
||||
}
|
||||
m.list.SetItems(items)
|
||||
}
|
||||
return m, nil
|
||||
|
||||
case connectRequestMsg:
|
||||
// Store result and quit TUI — caller will handle the connect
|
||||
m.result = &TUIResult{
|
||||
Server: msg.server,
|
||||
Action: "connect",
|
||||
}
|
||||
return m, tea.Quit
|
||||
|
||||
case testDoneMsg:
|
||||
if m.form != nil {
|
||||
m.form.testing = false
|
||||
if msg.ok {
|
||||
m.form.testResult = "Connection OK."
|
||||
m.form.testOK = true
|
||||
} else {
|
||||
m.form.testResult = fmt.Sprintf("Connection failed:\n%s", msg.err)
|
||||
m.form.testOK = false
|
||||
}
|
||||
m.form.testResultTime = time.Now()
|
||||
m.form.err = nil
|
||||
}
|
||||
return m, nil
|
||||
|
||||
case saveDoneMsg:
|
||||
if m.form != nil {
|
||||
m.form.saving = false
|
||||
if msg.err != nil {
|
||||
m.form.err = msg.err
|
||||
m.form.saved = false
|
||||
} else {
|
||||
m.form.saved = true
|
||||
m.form.savedTime = time.Now()
|
||||
m.form.err = nil
|
||||
}
|
||||
}
|
||||
return m, nil
|
||||
|
||||
case tea.KeyMsg:
|
||||
switch m.screen {
|
||||
case screenList:
|
||||
return m.updateList(msg)
|
||||
case screenForm:
|
||||
return m.updateForm(msg)
|
||||
case screenSearch:
|
||||
return m.updateSearch(msg)
|
||||
}
|
||||
}
|
||||
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m *tuiModel) updateList(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
switch msg.String() {
|
||||
case "q", "ctrl+c":
|
||||
return m, tea.Quit
|
||||
|
||||
case "/":
|
||||
m.screen = screenSearch
|
||||
m.searchInput.Focus()
|
||||
return m, nil
|
||||
|
||||
case "a":
|
||||
m.form = newFormModel(m.width, m.height)
|
||||
m.screen = screenForm
|
||||
return m, nil
|
||||
|
||||
case "e":
|
||||
if item, ok := m.list.SelectedItem().(serverItem); ok {
|
||||
m.form = newEditFormModel(item.server, m.width, m.height)
|
||||
m.screen = screenForm
|
||||
}
|
||||
return m, nil
|
||||
|
||||
case "d":
|
||||
if item, ok := m.list.SelectedItem().(serverItem); ok {
|
||||
return m, func() tea.Msg {
|
||||
err := DeleteServer(item.server.Alias)
|
||||
if err != nil {
|
||||
return saveDoneMsg{err: err}
|
||||
}
|
||||
servers, err := ListServers()
|
||||
return serversLoadedMsg{servers: servers, err: err}
|
||||
}
|
||||
}
|
||||
|
||||
case "t":
|
||||
if item, ok := m.list.SelectedItem().(serverItem); ok {
|
||||
return m, func() tea.Msg {
|
||||
ok, testErr := TestConnection(item.server)
|
||||
return testDoneMsg{ok: ok, err: testErr}
|
||||
}
|
||||
}
|
||||
|
||||
case "enter":
|
||||
if item, ok := m.list.SelectedItem().(serverItem); ok {
|
||||
// Request connect — TUI will quit and caller handles it
|
||||
return m, func() tea.Msg {
|
||||
return connectRequestMsg{server: item.server}
|
||||
}
|
||||
}
|
||||
|
||||
default:
|
||||
var cmd tea.Cmd
|
||||
m.list, cmd = m.list.Update(msg)
|
||||
return m, cmd
|
||||
}
|
||||
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m *tuiModel) updateSearch(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
switch msg.String() {
|
||||
case "esc":
|
||||
m.screen = screenList
|
||||
m.searchInput.Blur()
|
||||
m.searchInput.SetValue("")
|
||||
return m, nil
|
||||
|
||||
case "enter":
|
||||
m.screen = screenList
|
||||
m.searchInput.Blur()
|
||||
query := m.searchInput.Value()
|
||||
if query != "" {
|
||||
return m, func() tea.Msg {
|
||||
servers, err := SearchServers(query)
|
||||
return serversLoadedMsg{servers: servers, err: err}
|
||||
}
|
||||
}
|
||||
return m, func() tea.Msg {
|
||||
servers, err := ListServers()
|
||||
return serversLoadedMsg{servers: servers, err: err}
|
||||
}
|
||||
|
||||
default:
|
||||
var cmd tea.Cmd
|
||||
m.searchInput, cmd = m.searchInput.Update(msg)
|
||||
return m, cmd
|
||||
}
|
||||
}
|
||||
|
||||
func (m *tuiModel) updateForm(msg tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
switch msg.String() {
|
||||
case "esc":
|
||||
m.screen = screenList
|
||||
m.form = nil
|
||||
m.err = nil
|
||||
m.success = ""
|
||||
return m, nil
|
||||
}
|
||||
|
||||
updated, cmd := m.form.Update(msg)
|
||||
if fm, ok := updated.(*formModel); ok {
|
||||
m.form = fm
|
||||
}
|
||||
return m, cmd
|
||||
}
|
||||
|
||||
func (m *tuiModel) View() string {
|
||||
var b strings.Builder
|
||||
|
||||
switch m.screen {
|
||||
case screenList:
|
||||
b.WriteString(m.list.View())
|
||||
b.WriteString("\n")
|
||||
b.WriteString(helpStyle.Render("Enter connect | a add | e edit | d delete | t test | / search | q quit"))
|
||||
|
||||
case screenSearch:
|
||||
b.WriteString("Search: " + m.searchInput.View() + "\n")
|
||||
b.WriteString(helpStyle.Render("Enter search | Esc cancel"))
|
||||
|
||||
case screenForm:
|
||||
b.WriteString(m.form.View())
|
||||
}
|
||||
|
||||
if m.err != nil {
|
||||
b.WriteString("\n" + errorStyle.Render(fmt.Sprintf("Error: %v", m.err)))
|
||||
m.err = nil
|
||||
}
|
||||
if m.success != "" {
|
||||
b.WriteString("\n" + successStyle.Render(m.success))
|
||||
m.success = ""
|
||||
}
|
||||
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// --- Form model ---
|
||||
|
||||
type formModel struct {
|
||||
edit bool
|
||||
server *model.Server
|
||||
inputs []textinput.Model
|
||||
password textinput.Model
|
||||
focusIdx int
|
||||
testResult string
|
||||
testOK bool
|
||||
testResultTime time.Time
|
||||
testing bool
|
||||
saving bool
|
||||
saved bool
|
||||
savedTime time.Time
|
||||
err error
|
||||
spinner spinner.Model
|
||||
width int
|
||||
height int
|
||||
}
|
||||
|
||||
func newFormModel(w, h int) *formModel {
|
||||
inputs := make([]textinput.Model, 10)
|
||||
labels := []string{
|
||||
"Alias",
|
||||
"Display Name",
|
||||
"Host",
|
||||
"Port",
|
||||
"User",
|
||||
"Auth Method (password/key/key_passphrase/agent)",
|
||||
"Identity File",
|
||||
"ProxyJump",
|
||||
"Group",
|
||||
"Notes",
|
||||
}
|
||||
for i, label := range labels {
|
||||
inputs[i] = textinput.New()
|
||||
inputs[i].Placeholder = label
|
||||
inputs[i].CharLimit = 128
|
||||
}
|
||||
|
||||
pw := textinput.New()
|
||||
pw.Placeholder = "Password / Passphrase (stored in vault)"
|
||||
pw.CharLimit = 256
|
||||
pw.EchoMode = textinput.EchoPassword
|
||||
|
||||
s := spinner.New()
|
||||
s.Spinner = spinner.Dot
|
||||
s.Style = lipgloss.NewStyle().Foreground(lipgloss.Color("12"))
|
||||
|
||||
inputs[0].Focus()
|
||||
|
||||
return &formModel{
|
||||
inputs: inputs,
|
||||
password: pw,
|
||||
focusIdx: 0,
|
||||
spinner: s,
|
||||
width: w,
|
||||
height: h,
|
||||
}
|
||||
}
|
||||
|
||||
func newEditFormModel(s *model.Server, w, h int) *formModel {
|
||||
fm := newFormModel(w, h)
|
||||
fm.edit = true
|
||||
fm.server = s
|
||||
fm.inputs[0].SetValue(s.Alias)
|
||||
fm.inputs[1].SetValue(s.DisplayName)
|
||||
fm.inputs[2].SetValue(s.Host)
|
||||
fm.inputs[3].SetValue(fmt.Sprintf("%d", s.Port))
|
||||
fm.inputs[4].SetValue(s.User)
|
||||
fm.inputs[5].SetValue(string(s.AuthMethod))
|
||||
fm.inputs[6].SetValue(s.IdentityFile)
|
||||
fm.inputs[7].SetValue(s.ProxyJump)
|
||||
fm.inputs[8].SetValue(s.GroupName)
|
||||
fm.inputs[9].SetValue(s.Notes)
|
||||
return fm
|
||||
}
|
||||
|
||||
func (fm *formModel) Init() tea.Cmd {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (fm *formModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
// Handle test/save completion
|
||||
switch msg := msg.(type) {
|
||||
case testDoneMsg:
|
||||
fm.testing = false
|
||||
if msg.ok {
|
||||
fm.testResult = "Connection OK."
|
||||
fm.testOK = true
|
||||
} else {
|
||||
fm.testResult = fmt.Sprintf("Connection failed:\n%s", msg.err)
|
||||
fm.testOK = false
|
||||
}
|
||||
fm.testResultTime = time.Now()
|
||||
fm.err = nil
|
||||
return fm, nil
|
||||
case saveDoneMsg:
|
||||
fm.saving = false
|
||||
if msg.err != nil {
|
||||
fm.err = msg.err
|
||||
fm.saved = false
|
||||
} else {
|
||||
fm.saved = true
|
||||
fm.savedTime = time.Now()
|
||||
fm.err = nil
|
||||
}
|
||||
return fm, nil
|
||||
}
|
||||
|
||||
// Handle spinner tick while testing/saving
|
||||
if fm.testing || fm.saving {
|
||||
var cmd tea.Cmd
|
||||
fm.spinner, cmd = fm.spinner.Update(msg)
|
||||
if _, ok := msg.(tea.KeyMsg); ok {
|
||||
return fm, cmd
|
||||
}
|
||||
return fm, cmd
|
||||
}
|
||||
|
||||
switch msg := msg.(type) {
|
||||
case tea.KeyMsg:
|
||||
switch msg.String() {
|
||||
case "tab", "down":
|
||||
fm.focusIdx++
|
||||
total := len(fm.inputs) + 3
|
||||
if fm.focusIdx >= total {
|
||||
fm.focusIdx = 0
|
||||
}
|
||||
fm.updateFocus()
|
||||
return fm, nil
|
||||
|
||||
case "shift+tab", "up":
|
||||
fm.focusIdx--
|
||||
if fm.focusIdx < 0 {
|
||||
total := len(fm.inputs) + 3
|
||||
fm.focusIdx = total - 1
|
||||
}
|
||||
fm.updateFocus()
|
||||
return fm, nil
|
||||
|
||||
case "enter":
|
||||
switch {
|
||||
case fm.focusIdx == len(fm.inputs)+1:
|
||||
return fm, fm.runTest()
|
||||
case fm.focusIdx == len(fm.inputs)+2:
|
||||
return fm, fm.runSave()
|
||||
default:
|
||||
fm.focusIdx++
|
||||
total := len(fm.inputs) + 3
|
||||
if fm.focusIdx >= total {
|
||||
fm.focusIdx = 0
|
||||
}
|
||||
fm.updateFocus()
|
||||
return fm, nil
|
||||
}
|
||||
|
||||
case "esc":
|
||||
return fm, nil
|
||||
}
|
||||
}
|
||||
|
||||
if fm.focusIdx < len(fm.inputs) {
|
||||
var cmd tea.Cmd
|
||||
fm.inputs[fm.focusIdx], cmd = fm.inputs[fm.focusIdx].Update(msg)
|
||||
return fm, cmd
|
||||
}
|
||||
|
||||
if fm.focusIdx == len(fm.inputs) {
|
||||
var cmd tea.Cmd
|
||||
fm.password, cmd = fm.password.Update(msg)
|
||||
return fm, cmd
|
||||
}
|
||||
|
||||
return fm, nil
|
||||
}
|
||||
|
||||
func (fm *formModel) updateFocus() {
|
||||
for i := range fm.inputs {
|
||||
fm.inputs[i].Blur()
|
||||
fm.inputs[i].Prompt = blurredStyle.Render(fm.inputs[i].Placeholder + ": ")
|
||||
}
|
||||
fm.password.Blur()
|
||||
fm.password.Prompt = blurredStyle.Render(fm.password.Placeholder + ": ")
|
||||
|
||||
if fm.focusIdx < len(fm.inputs) {
|
||||
fm.inputs[fm.focusIdx].Focus()
|
||||
fm.inputs[fm.focusIdx].Prompt = focusedStyle.Render(fm.inputs[fm.focusIdx].Placeholder + "> ")
|
||||
} else if fm.focusIdx == len(fm.inputs) {
|
||||
fm.password.Focus()
|
||||
fm.password.Prompt = focusedStyle.Render(fm.password.Placeholder + "> ")
|
||||
}
|
||||
}
|
||||
|
||||
func (fm *formModel) runTest() tea.Cmd {
|
||||
fm.testing = true
|
||||
fm.testResult = ""
|
||||
fm.err = nil
|
||||
fm.saved = false
|
||||
|
||||
s := fm.buildServer()
|
||||
return tea.Batch(
|
||||
fm.spinner.Tick,
|
||||
func() tea.Msg {
|
||||
if s.AuthMethod == model.AuthPassword && fm.password.Value() == "" {
|
||||
return testDoneMsg{ok: false, err: "Password is required for password auth."}
|
||||
}
|
||||
ok, testErr := TestConnection(s)
|
||||
return testDoneMsg{ok: ok, err: testErr}
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
func (fm *formModel) runSave() tea.Cmd {
|
||||
fm.saving = true
|
||||
fm.err = nil
|
||||
fm.saved = false
|
||||
fm.testResult = ""
|
||||
|
||||
s := fm.buildServer()
|
||||
pw := fm.password.Value()
|
||||
|
||||
return tea.Batch(
|
||||
fm.spinner.Tick,
|
||||
func() tea.Msg {
|
||||
if s.Alias == "" {
|
||||
return saveDoneMsg{err: fmt.Errorf("alias is required")}
|
||||
}
|
||||
if s.Host == "" {
|
||||
return saveDoneMsg{err: fmt.Errorf("host is required")}
|
||||
}
|
||||
err := SaveServer(s, pw)
|
||||
return saveDoneMsg{err: err}
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
func (fm *formModel) buildServer() *model.Server {
|
||||
port := 22
|
||||
fmt.Sscanf(fm.inputs[3].Value(), "%d", &port)
|
||||
authMethod := model.AuthMethod(fm.inputs[5].Value())
|
||||
if authMethod == "" {
|
||||
authMethod = model.AuthKey
|
||||
}
|
||||
return &model.Server{
|
||||
Alias: fm.inputs[0].Value(),
|
||||
DisplayName: fm.inputs[1].Value(),
|
||||
Host: fm.inputs[2].Value(),
|
||||
Port: port,
|
||||
User: fm.inputs[4].Value(),
|
||||
AuthMethod: authMethod,
|
||||
IdentityFile: fm.inputs[6].Value(),
|
||||
ProxyJump: fm.inputs[7].Value(),
|
||||
GroupName: fm.inputs[8].Value(),
|
||||
Notes: fm.inputs[9].Value(),
|
||||
}
|
||||
}
|
||||
|
||||
func (fm *formModel) View() string {
|
||||
var b strings.Builder
|
||||
|
||||
title := "Add Server"
|
||||
if fm.edit {
|
||||
title = "Edit Server: " + fm.server.Alias
|
||||
}
|
||||
b.WriteString(titleStyle.Render(title))
|
||||
b.WriteString("\n\n")
|
||||
|
||||
for i := range fm.inputs {
|
||||
b.WriteString(fm.inputs[i].View())
|
||||
b.WriteString("\n")
|
||||
}
|
||||
|
||||
b.WriteString(fm.password.View())
|
||||
b.WriteString("\n")
|
||||
|
||||
showResults := time.Since(fm.testResultTime) < 10*time.Second || time.Since(fm.savedTime) < 10*time.Second
|
||||
|
||||
if fm.testing {
|
||||
b.WriteString("\n" + fm.spinner.View() + " Testing connection...\n")
|
||||
} else if fm.saving {
|
||||
b.WriteString("\n" + fm.spinner.View() + " Saving...\n")
|
||||
} else if showResults {
|
||||
if fm.testResult != "" {
|
||||
b.WriteString("\n")
|
||||
if fm.testOK {
|
||||
b.WriteString(testOKStyle.Render("✓ " + fm.testResult))
|
||||
} else {
|
||||
b.WriteString(testFailStyle.Render("✗ " + fm.testResult))
|
||||
}
|
||||
b.WriteString("\n")
|
||||
}
|
||||
if fm.saved {
|
||||
b.WriteString("\n" + successStyle.Render("✓ Saved.") + "\n")
|
||||
}
|
||||
if fm.err != nil {
|
||||
b.WriteString("\n" + errorStyle.Render(fmt.Sprintf("✗ Error: %v", fm.err)) + "\n")
|
||||
}
|
||||
}
|
||||
|
||||
testBtn := "[ Test ]"
|
||||
saveBtn := "[ Save ]"
|
||||
|
||||
if fm.focusIdx == len(fm.inputs)+1 {
|
||||
testBtn = selectedStyle.Render(testBtn)
|
||||
} else {
|
||||
testBtn = normalStyle.Render(testBtn)
|
||||
}
|
||||
|
||||
if fm.focusIdx == len(fm.inputs)+2 {
|
||||
saveBtn = selectedStyle.Render(saveBtn)
|
||||
} else {
|
||||
saveBtn = normalStyle.Render(saveBtn)
|
||||
}
|
||||
|
||||
b.WriteString("\n" + testBtn + " " + saveBtn + "\n\n")
|
||||
b.WriteString(helpStyle.Render("Tab/↓ next | ↑ prev | Enter select | Esc back"))
|
||||
|
||||
return b.String()
|
||||
}
|
||||
@@ -0,0 +1,463 @@
|
||||
package vault
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/argon2"
|
||||
"golang.org/x/crypto/chacha20poly1305"
|
||||
)
|
||||
|
||||
const (
|
||||
currentVersion = 1
|
||||
saltLen = 32
|
||||
nonceLen = 24
|
||||
keyLen = 32
|
||||
)
|
||||
|
||||
type KDFMeta struct {
|
||||
Name string `json:"name"`
|
||||
MemoryKiB int `json:"memory_kib"`
|
||||
Iterations int `json:"iterations"`
|
||||
Parallelism int `json:"parallelism"`
|
||||
Salt string `json:"salt"`
|
||||
}
|
||||
|
||||
type Record struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
Nonce string `json:"nonce"`
|
||||
Ciphertext string `json:"ciphertext"`
|
||||
}
|
||||
|
||||
type VaultFile struct {
|
||||
Version int `json:"version"`
|
||||
KDF KDFMeta `json:"kdf"`
|
||||
Records []Record `json:"records"`
|
||||
}
|
||||
|
||||
type Vault struct {
|
||||
mu sync.Mutex
|
||||
path string
|
||||
masterKey []byte
|
||||
records map[string][]byte // id -> plaintext
|
||||
modified bool
|
||||
}
|
||||
|
||||
func New(path string) *Vault {
|
||||
return &Vault{
|
||||
path: path,
|
||||
records: make(map[string][]byte),
|
||||
}
|
||||
}
|
||||
|
||||
// Exists checks if vault file exists and has content
|
||||
func Exists(path string) bool {
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return info.Size() > 0
|
||||
}
|
||||
|
||||
// Create initializes a new vault with a master password
|
||||
func Create(path string, masterPassword string) error {
|
||||
salt := make([]byte, saltLen)
|
||||
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
|
||||
return fmt.Errorf("generate salt: %w", err)
|
||||
}
|
||||
|
||||
kdf := KDFMeta{
|
||||
Name: "argon2id",
|
||||
MemoryKiB: 4096,
|
||||
Iterations: 2,
|
||||
Parallelism: 1,
|
||||
Salt: base64.StdEncoding.EncodeToString(salt),
|
||||
}
|
||||
|
||||
fmt.Print("Deriving key...")
|
||||
|
||||
key := argon2.IDKey([]byte(masterPassword), salt, uint32(kdf.Iterations), uint32(kdf.MemoryKiB)*1024, uint8(kdf.Parallelism), keyLen)
|
||||
|
||||
// Verify key is valid by doing a test encrypt/decrypt
|
||||
vf := VaultFile{
|
||||
Version: currentVersion,
|
||||
KDF: kdf,
|
||||
Records: []Record{},
|
||||
}
|
||||
|
||||
data, err := json.Marshal(vf)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal vault: %w", err)
|
||||
}
|
||||
|
||||
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0600)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create vault file: %w", err)
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
if _, err := f.Write(data); err != nil {
|
||||
return fmt.Errorf("write vault: %w", err)
|
||||
}
|
||||
|
||||
// Clear key from memory
|
||||
for i := range key {
|
||||
key[i] = 0
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Unlock decrypts the vault with master password
|
||||
func (v *Vault) Unlock(masterPassword string) error {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
data, err := os.ReadFile(v.path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read vault file: %w", err)
|
||||
}
|
||||
|
||||
var vf VaultFile
|
||||
if err := json.Unmarshal(data, &vf); err != nil {
|
||||
return fmt.Errorf("parse vault: %w", err)
|
||||
}
|
||||
|
||||
if vf.Version != currentVersion {
|
||||
return fmt.Errorf("unsupported vault version: %d", vf.Version)
|
||||
}
|
||||
|
||||
salt, err := base64.StdEncoding.DecodeString(vf.KDF.Salt)
|
||||
if err != nil {
|
||||
return fmt.Errorf("decode salt: %w", err)
|
||||
}
|
||||
|
||||
key := argon2.IDKey([]byte(masterPassword), salt, uint32(vf.KDF.Iterations), uint32(vf.KDF.MemoryKiB)*1024, uint8(vf.KDF.Parallelism), keyLen)
|
||||
|
||||
// Try to decrypt first record to verify password
|
||||
if len(vf.Records) > 0 {
|
||||
if _, err := decryptRecord(key, vf.Records[0]); err != nil {
|
||||
return fmt.Errorf("invalid master password")
|
||||
}
|
||||
}
|
||||
|
||||
v.masterKey = key
|
||||
v.records = make(map[string][]byte)
|
||||
|
||||
for _, rec := range vf.Records {
|
||||
plaintext, err := decryptRecord(key, rec)
|
||||
if err != nil {
|
||||
return fmt.Errorf("decrypt record %s: %w", rec.ID, err)
|
||||
}
|
||||
v.records[rec.ID] = plaintext
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Lock clears the master key and records from memory
|
||||
func (v *Vault) Lock() {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
if v.masterKey != nil {
|
||||
for i := range v.masterKey {
|
||||
v.masterKey[i] = 0
|
||||
}
|
||||
}
|
||||
v.masterKey = nil
|
||||
v.records = make(map[string][]byte)
|
||||
}
|
||||
|
||||
// IsUnlocked returns whether the vault is currently unlocked
|
||||
func (v *Vault) IsUnlocked() bool {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
return v.masterKey != nil
|
||||
}
|
||||
|
||||
// Put stores a secret in memory (not persisted until Save)
|
||||
func (v *Vault) Put(id string, secretType string, plaintext []byte) error {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
if v.masterKey == nil {
|
||||
return fmt.Errorf("vault is locked")
|
||||
}
|
||||
|
||||
v.records[id] = plaintext
|
||||
v.modified = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get retrieves a secret
|
||||
func (v *Vault) Get(id string) ([]byte, error) {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
if v.masterKey == nil {
|
||||
return nil, fmt.Errorf("vault is locked")
|
||||
}
|
||||
|
||||
data, ok := v.records[id]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("secret not found: %s", id)
|
||||
}
|
||||
|
||||
result := make([]byte, len(data))
|
||||
copy(result, data)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// Delete removes a secret
|
||||
func (v *Vault) Delete(id string) {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
delete(v.records, id)
|
||||
v.modified = true
|
||||
}
|
||||
|
||||
// Save persists encrypted vault to disk
|
||||
func (v *Vault) Save() error {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
if v.masterKey == nil {
|
||||
return fmt.Errorf("vault is locked")
|
||||
}
|
||||
|
||||
salt, err := base64.StdEncoding.DecodeString(v.getSalt())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
kdf := KDFMeta{
|
||||
Name: "argon2id",
|
||||
MemoryKiB: 4096,
|
||||
Iterations: 2,
|
||||
Parallelism: 1,
|
||||
Salt: base64.StdEncoding.EncodeToString(salt),
|
||||
}
|
||||
|
||||
fmt.Print("Deriving key...")
|
||||
|
||||
var records []Record
|
||||
for id, plaintext := range v.records {
|
||||
rec, err := encryptRecord(v.masterKey, id, plaintext)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypt record %s: %w", id, err)
|
||||
}
|
||||
records = append(records, rec)
|
||||
}
|
||||
|
||||
vf := VaultFile{
|
||||
Version: currentVersion,
|
||||
KDF: kdf,
|
||||
Records: records,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(vf)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal vault: %w", err)
|
||||
}
|
||||
|
||||
tmpPath := v.path + ".tmp"
|
||||
f, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0600)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create temp vault: %w", err)
|
||||
}
|
||||
|
||||
if _, err := f.Write(data); err != nil {
|
||||
f.Close()
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("write vault: %w", err)
|
||||
}
|
||||
f.Close()
|
||||
|
||||
if err := os.Rename(tmpPath, v.path); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return fmt.Errorf("rename vault: %w", err)
|
||||
}
|
||||
|
||||
v.modified = false
|
||||
return nil
|
||||
}
|
||||
|
||||
// ChangePassword re-encrypts the vault with a new master password
|
||||
func (v *Vault) ChangePassword(newPassword string) error {
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
if v.masterKey == nil {
|
||||
return fmt.Errorf("vault is locked")
|
||||
}
|
||||
|
||||
salt := make([]byte, saltLen)
|
||||
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
|
||||
return fmt.Errorf("generate salt: %w", err)
|
||||
}
|
||||
|
||||
newKey := argon2.IDKey([]byte(newPassword), salt, 3, 8192*1024, 1, keyLen)
|
||||
|
||||
kdf := KDFMeta{
|
||||
Name: "argon2id",
|
||||
MemoryKiB: 4096,
|
||||
Iterations: 2,
|
||||
Parallelism: 1,
|
||||
Salt: base64.StdEncoding.EncodeToString(salt),
|
||||
}
|
||||
|
||||
fmt.Print("Deriving key...")
|
||||
|
||||
var records []Record
|
||||
for id, plaintext := range v.records {
|
||||
rec, err := encryptRecord(newKey, id, plaintext)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypt record: %w", err)
|
||||
}
|
||||
records = append(records, rec)
|
||||
}
|
||||
|
||||
vf := VaultFile{
|
||||
Version: currentVersion,
|
||||
KDF: kdf,
|
||||
Records: records,
|
||||
}
|
||||
|
||||
data, err := json.Marshal(vf)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
tmpPath := v.path + ".tmp"
|
||||
f, err := os.OpenFile(tmpPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := f.Write(data); err != nil {
|
||||
f.Close()
|
||||
os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
f.Close()
|
||||
|
||||
if err := os.Rename(tmpPath, v.path); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
|
||||
// Swap key
|
||||
for i := range v.masterKey {
|
||||
v.masterKey[i] = 0
|
||||
}
|
||||
v.masterKey = newKey
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Helper to get salt from existing vault
|
||||
func (v *Vault) getSalt() string {
|
||||
data, err := os.ReadFile(v.path)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
var vf VaultFile
|
||||
if err := json.Unmarshal(data, &vf); err != nil {
|
||||
return ""
|
||||
}
|
||||
return vf.KDF.Salt
|
||||
}
|
||||
|
||||
func encryptRecord(key []byte, id string, plaintext []byte) (Record, error) {
|
||||
aead, err := chacha20poly1305.NewX(key)
|
||||
if err != nil {
|
||||
return Record{}, err
|
||||
}
|
||||
|
||||
nonce := make([]byte, nonceLen)
|
||||
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
|
||||
return Record{}, err
|
||||
}
|
||||
|
||||
ciphertext := aead.Seal(nil, nonce, plaintext, []byte(id))
|
||||
|
||||
return Record{
|
||||
ID: id,
|
||||
Nonce: base64.StdEncoding.EncodeToString(nonce),
|
||||
Ciphertext: base64.StdEncoding.EncodeToString(ciphertext),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func decryptRecord(key []byte, rec Record) ([]byte, error) {
|
||||
aead, err := chacha20poly1305.NewX(key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nonce, err := base64.StdEncoding.DecodeString(rec.Nonce)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode nonce: %w", err)
|
||||
}
|
||||
|
||||
ciphertext, err := base64.StdEncoding.DecodeString(rec.Ciphertext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode ciphertext: %w", err)
|
||||
}
|
||||
|
||||
plaintext, err := aead.Open(nil, nonce, ciphertext, []byte(rec.ID))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decrypt failed: %w", err)
|
||||
}
|
||||
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
// VerifyPassword checks if a master password is correct without unlocking
|
||||
func VerifyPassword(path string, masterPassword string) (bool, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
var vf VaultFile
|
||||
if err := json.Unmarshal(data, &vf); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
salt, err := base64.StdEncoding.DecodeString(vf.KDF.Salt)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
key := argon2.IDKey([]byte(masterPassword), salt, uint32(vf.KDF.Iterations), uint32(vf.KDF.MemoryKiB)*1024, uint8(vf.KDF.Parallelism), keyLen)
|
||||
defer func() {
|
||||
for i := range key {
|
||||
key[i] = 0
|
||||
}
|
||||
}()
|
||||
|
||||
if len(vf.Records) == 0 {
|
||||
// Empty vault, try a test encryption
|
||||
return true, nil
|
||||
}
|
||||
|
||||
_, err = decryptRecord(key, vf.Records[0])
|
||||
return err == nil, nil
|
||||
}
|
||||
|
||||
// Constant-time comparison to prevent timing attacks
|
||||
func SecureCompare(a, b string) bool {
|
||||
return subtle.ConstantTimeCompare([]byte(a), []byte(b)) == 1
|
||||
}
|
||||
|
||||
// Ensure time import is used
|
||||
var _ time.Duration
|
||||
Reference in New Issue
Block a user