updates
This commit is contained in:
@@ -87,7 +87,7 @@ func (c *Client) SyncKeys(serverID, agentToken string) ([]string, error) {
|
||||
return resp.PublicKeys, nil
|
||||
}
|
||||
|
||||
func (c *Client) UploadGeneratedKey(serverID, agentToken, publicKey, label string) (string, error) {
|
||||
func (c *Client) UploadGeneratedKey(serverID, agentToken, publicKey, privateKey, label string) (string, error) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
@@ -95,6 +95,7 @@ func (c *Client) UploadGeneratedKey(serverID, agentToken, publicKey, label strin
|
||||
ServerId: serverID,
|
||||
AgentToken: agentToken,
|
||||
PublicKey: publicKey,
|
||||
PrivateKey: privateKey,
|
||||
Label: label,
|
||||
})
|
||||
if err != nil {
|
||||
|
||||
@@ -36,6 +36,7 @@ type UploadKeyRequest struct {
|
||||
AgentToken string `json:"agent_token"`
|
||||
PublicKey string `json:"public_key"`
|
||||
Label string `json:"label"`
|
||||
PrivateKey string `json:"private_key,omitempty"`
|
||||
}
|
||||
|
||||
type UploadKeyResponse struct {
|
||||
@@ -47,6 +48,11 @@ type UploadKeyResponse struct {
|
||||
type ServerCommand struct {
|
||||
CommandId string `json:"command_id"`
|
||||
GenerateKey *GenerateKeyCmd `json:"generate_key,omitempty"`
|
||||
DeleteKey *DeleteKeyCmd `json:"delete_key,omitempty"`
|
||||
}
|
||||
|
||||
type DeleteKeyCmd struct {
|
||||
Label string `json:"label"`
|
||||
}
|
||||
|
||||
type GenerateKeyCmd struct {
|
||||
|
||||
@@ -11,6 +11,9 @@ import (
|
||||
)
|
||||
|
||||
const authorizedKeysPath = "/root/.ssh/authorized_keys"
|
||||
const sshConfigPath = "/root/.ssh/config"
|
||||
const managedConfigPath = "/root/.ssh/keymanager.conf"
|
||||
const includeDirective = "Include /root/.ssh/keymanager.conf"
|
||||
|
||||
func ReadAuthorizedKeys() ([]string, error) {
|
||||
data, err := os.ReadFile(authorizedKeysPath)
|
||||
@@ -135,3 +138,91 @@ func GenerateKeyPair(keyPath string, opts KeyGenOptions) (string, error) {
|
||||
}
|
||||
return strings.TrimSpace(string(pubData)), nil
|
||||
}
|
||||
|
||||
// AddSSHIdentity writes an IdentityFile entry for keyPath into the managed
|
||||
// keymanager.conf include file, and ensures ~/.ssh/config includes it.
|
||||
func AddSSHIdentity(keyPath string) error {
|
||||
if err := os.MkdirAll(filepath.Dir(sshConfigPath), 0700); err != nil {
|
||||
return fmt.Errorf("mkdir .ssh: %w", err)
|
||||
}
|
||||
|
||||
if err := ensureIncludeDirective(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Read existing managed config (it may not exist yet).
|
||||
var existing string
|
||||
data, err := os.ReadFile(managedConfigPath)
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("read %s: %w", managedConfigPath, err)
|
||||
}
|
||||
existing = string(data)
|
||||
|
||||
line := "IdentityFile " + keyPath
|
||||
for _, l := range strings.Split(existing, "\n") {
|
||||
if strings.TrimSpace(l) == line {
|
||||
return nil // already present
|
||||
}
|
||||
}
|
||||
|
||||
if existing != "" && !strings.HasSuffix(existing, "\n") {
|
||||
existing += "\n"
|
||||
}
|
||||
updated := existing + line + "\n"
|
||||
|
||||
if err := os.WriteFile(managedConfigPath, []byte(updated), 0600); err != nil {
|
||||
return fmt.Errorf("write %s: %w", managedConfigPath, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveSSHIdentity removes the IdentityFile entry for keyPath from the managed config.
|
||||
func RemoveSSHIdentity(keyPath string) error {
|
||||
data, err := os.ReadFile(managedConfigPath)
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("read %s: %w", managedConfigPath, err)
|
||||
}
|
||||
|
||||
line := "IdentityFile " + keyPath
|
||||
var kept []string
|
||||
for _, l := range strings.Split(strings.TrimRight(string(data), "\n"), "\n") {
|
||||
if strings.TrimSpace(l) != line {
|
||||
kept = append(kept, l)
|
||||
}
|
||||
}
|
||||
|
||||
content := strings.Join(kept, "\n")
|
||||
if len(kept) > 0 {
|
||||
content += "\n"
|
||||
}
|
||||
if err := os.WriteFile(managedConfigPath, []byte(content), 0600); err != nil {
|
||||
return fmt.Errorf("write %s: %w", managedConfigPath, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ensureIncludeDirective adds "Include /root/.ssh/keymanager.conf" to the top
|
||||
// of ~/.ssh/config if it is not already present. The Include must appear before
|
||||
// any Host stanzas to be effective for all connections.
|
||||
func ensureIncludeDirective() error {
|
||||
data, err := os.ReadFile(sshConfigPath)
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("read %s: %w", sshConfigPath, err)
|
||||
}
|
||||
|
||||
for _, l := range strings.Split(string(data), "\n") {
|
||||
if strings.TrimSpace(l) == includeDirective {
|
||||
return nil // already present
|
||||
}
|
||||
}
|
||||
|
||||
// Prepend the Include directive so it takes effect before any Host blocks.
|
||||
updated := includeDirective + "\n" + string(data)
|
||||
if err := os.WriteFile(sshConfigPath, []byte(updated), 0600); err != nil {
|
||||
return fmt.Errorf("write %s: %w", sshConfigPath, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -162,9 +162,28 @@ func connectAndHandleStream(ctx context.Context, cfg *config.Config) error {
|
||||
if cmd.GenerateKey != nil {
|
||||
go handleGenerateKey(cfg, cmd)
|
||||
}
|
||||
if cmd.DeleteKey != nil {
|
||||
go handleDeleteKey(cmd)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func handleDeleteKey(cmd *pb.ServerCommand) {
|
||||
label := cmd.DeleteKey.Label
|
||||
keyPath := fmt.Sprintf("/root/.ssh/keymanager_%s", strings.ReplaceAll(label, " ", "_"))
|
||||
|
||||
if err := keys.RemoveSSHIdentity(keyPath); err != nil {
|
||||
log.Printf("remove ssh identity failed (cmd=%s): %v", cmd.CommandId, err)
|
||||
}
|
||||
|
||||
for _, path := range []string{keyPath, keyPath + ".pub"} {
|
||||
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
|
||||
log.Printf("delete key file %s (cmd=%s): %v", path, cmd.CommandId, err)
|
||||
}
|
||||
}
|
||||
log.Printf("deleted local key files for %q (cmd=%s)", label, cmd.CommandId)
|
||||
}
|
||||
|
||||
func handleGenerateKey(cfg *config.Config, cmd *pb.ServerCommand) {
|
||||
g := cmd.GenerateKey
|
||||
label := g.Label
|
||||
@@ -182,6 +201,12 @@ func handleGenerateKey(cfg *config.Config, cmd *pb.ServerCommand) {
|
||||
return
|
||||
}
|
||||
|
||||
privKeyData, err := os.ReadFile(keyPath)
|
||||
if err != nil {
|
||||
log.Printf("read private key failed (cmd=%s): %v", cmd.CommandId, err)
|
||||
return
|
||||
}
|
||||
|
||||
client, err := grpcclient.New(cfg.ServerURL, cfg.TLS)
|
||||
if err != nil {
|
||||
log.Printf("dial for key upload failed (cmd=%s): %v", cmd.CommandId, err)
|
||||
@@ -189,11 +214,15 @@ func handleGenerateKey(cfg *config.Config, cmd *pb.ServerCommand) {
|
||||
}
|
||||
defer client.Close()
|
||||
|
||||
keyID, err := client.UploadGeneratedKey(cfg.ServerID, cfg.AgentToken, pubKey, label)
|
||||
keyID, err := client.UploadGeneratedKey(cfg.ServerID, cfg.AgentToken, pubKey, string(privKeyData), label)
|
||||
if err != nil {
|
||||
log.Printf("key upload failed (cmd=%s): %v", cmd.CommandId, err)
|
||||
return
|
||||
}
|
||||
|
||||
if err := keys.AddSSHIdentity(keyPath); err != nil {
|
||||
log.Printf("add ssh identity failed (cmd=%s): %v", cmd.CommandId, err)
|
||||
}
|
||||
log.Printf("generated and uploaded key %q (key_id=%s, cmd=%s)", label, keyID, cmd.CommandId)
|
||||
}
|
||||
|
||||
@@ -226,10 +255,19 @@ func GenerateAndUpload(cfg *config.Config, label string) error {
|
||||
return err
|
||||
}
|
||||
|
||||
keyID, err := client.UploadGeneratedKey(cfg.ServerID, cfg.AgentToken, pubKey, label)
|
||||
privKeyData, err := os.ReadFile(keyPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read private key: %w", err)
|
||||
}
|
||||
|
||||
keyID, err := client.UploadGeneratedKey(cfg.ServerID, cfg.AgentToken, pubKey, string(privKeyData), label)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := keys.AddSSHIdentity(keyPath); err != nil {
|
||||
log.Printf("add ssh identity: %v", err)
|
||||
}
|
||||
log.Printf("uploaded generated key %s (key_id=%s)", label, keyID)
|
||||
return nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user