Files
HostKeeper/pkg/knownhosts/knownhosts.go
T
swanadiva 4a8b4a2bc4 test: Phase 3 comprehensive test coverage — 105 tests, zero race conditions
- Add test/crypto/ (19 tests): AES-256-GCM encrypt/decrypt, PBKDF2, IsEncrypted, HashPassword
- Add test/knownhosts/ (14 tests): TOFU verify, MITM detection, CRUD, persistence
- Add test/storage/ (20 tests): KeyPair/Snippet CRUD, encryption, MergeStrategy
- Add test/config/ (9 tests): config lifecycle, path getters
- Add test/tui/ (10 tests): WrapFooter, ClampWidth, TruncateStr
- Fix knownhosts deadlock: Add/Remove use saveInternal()
- Export responsive.go functions for testing
- Add docs/TEST_PLAN.md with full scenario documentation

Coverage: crypto 0%→100%, knownhosts 0%→100%, storage 40%→90%, config 22%→80%, tui 40%→70%
2026-06-29 12:05:02 +07:00

193 lines
4.5 KiB
Go

package knownhosts
import (
"encoding/base64"
"encoding/json"
"fmt"
"net"
"os"
"path/filepath"
"sync"
"time"
cryptossh "golang.org/x/crypto/ssh"
)
// KnownHosts manages the known_hosts file
type KnownHosts struct {
path string
hosts map[string]*HostKey // key = "hostname:port"
mu sync.RWMutex
}
// HostKey represents a stored host key
type HostKey struct {
Hostname string `json:"hostname"`
Port int `json:"port"`
KeyType string `json:"key_type"`
KeyData string `json:"key_data"` // Base64-encoded raw key
AddedAt time.Time `json:"added_at"`
}
// New creates a new KnownHosts manager
func New(dataDir string) (*KnownHosts, error) {
path := filepath.Join(dataDir, "known_hosts")
kh := &KnownHosts{
path: path,
hosts: make(map[string]*HostKey),
}
if err := kh.load(); err != nil {
// File doesn't exist yet, that's OK
if !os.IsNotExist(err) {
return nil, err
}
}
return kh, nil
}
// Load reads the known_hosts file
func (kh *KnownHosts) load() error {
kh.mu.Lock()
defer kh.mu.Unlock()
data, err := os.ReadFile(kh.path)
if err != nil {
return err
}
var hosts []*HostKey
if err := json.Unmarshal(data, &hosts); err != nil {
return err
}
for _, h := range hosts {
key := fmt.Sprintf("%s:%d", h.Hostname, h.Port)
kh.hosts[key] = h
}
return nil
}
// Save writes the known_hosts file
func (kh *KnownHosts) Save() error {
kh.mu.Lock()
defer kh.mu.Unlock()
return kh.saveInternal()
}
// saveInternal writes the known_hosts file without locking (caller must hold lock)
func (kh *KnownHosts) saveInternal() error {
var hosts []*HostKey
for _, h := range kh.hosts {
hosts = append(hosts, h)
}
data, err := json.MarshalIndent(hosts, "", " ")
if err != nil {
return err
}
return os.WriteFile(kh.path, data, 0600)
}
// Verify checks if a host key is known and matches
func (kh *KnownHosts) Verify(hostname string, port int, remoteKey cryptossh.PublicKey) (bool, *HostKey) {
kh.mu.RLock()
defer kh.mu.RUnlock()
key := fmt.Sprintf("%s:%d", hostname, port)
stored, ok := kh.hosts[key]
if !ok {
return false, nil // Unknown host
}
// Compare key type and data
remoteType := remoteKey.Type()
remoteData := base64.StdEncoding.EncodeToString(remoteKey.Marshal())
if stored.KeyType != remoteType || stored.KeyData != remoteData {
return false, stored // Key mismatch — potential MITM
}
return true, stored // Key matches
}
// Add stores a new host key
func (kh *KnownHosts) Add(hostname string, port int, remoteKey cryptossh.PublicKey) error {
kh.mu.Lock()
defer kh.mu.Unlock()
key := fmt.Sprintf("%s:%d", hostname, port)
kh.hosts[key] = &HostKey{
Hostname: hostname,
Port: port,
KeyType: remoteKey.Type(),
KeyData: base64.StdEncoding.EncodeToString(remoteKey.Marshal()),
AddedAt: time.Now(),
}
return kh.saveInternal()
}
// Remove removes a host key
func (kh *KnownHosts) Remove(hostname string, port int) error {
kh.mu.Lock()
defer kh.mu.Unlock()
key := fmt.Sprintf("%s:%d", hostname, port)
delete(kh.hosts, key)
return kh.saveInternal()
}
// Get returns the stored host key for a given host
func (kh *KnownHosts) Get(hostname string, port int) *HostKey {
kh.mu.RLock()
defer kh.mu.RUnlock()
key := fmt.Sprintf("%s:%d", hostname, port)
return kh.hosts[key]
}
// HostKeyCallback returns a crypto/ssh HostKeyCallback for use in SSH config
func (kh *KnownHosts) HostKeyCallback(autoAdd bool) cryptossh.HostKeyCallback {
return func(hostname string, remote net.Addr, remoteKey cryptossh.PublicKey) error {
// Extract port from address
_, portStr, err := net.SplitHostPort(remote.String())
if err != nil {
return fmt.Errorf("failed to parse address: %w", err)
}
port := 22
fmt.Sscanf(portStr, "%d", &port)
// Check if host is known
matches, stored := kh.Verify(hostname, port, remoteKey)
if matches {
return nil // Key matches, connection OK
}
if stored == nil {
// Unknown host — auto-add if enabled
if autoAdd {
if err := kh.Add(hostname, port, remoteKey); err != nil {
return fmt.Errorf("failed to add host key: %w", err)
}
return nil
}
return fmt.Errorf("host key not found for %s:%d — run 'hostkeeper trust %s' to add", hostname, port, hostname)
}
// Key mismatch — potential MITM attack
return fmt.Errorf("WARNING: host key mismatch for %s:%d!\n"+
"Stored key type: %s\n"+
"Remote key type: %s\n"+
"This could indicate a MITM attack.\n"+
"Run 'hostkeeper trust --remove %s' and try again.",
hostname, port, stored.KeyType, remoteKey.Type(), hostname)
}
}