|
| 1 | +import * as path from 'path' |
| 2 | +import * as fs from 'fs/promises' |
| 3 | +import * as crypto from 'crypto' |
| 4 | +import type { IPCHandler } from './ipcServer' |
| 5 | +import type { Database } from 'bun:sqlite' |
| 6 | +import { logger } from '../utils/logger' |
| 7 | +import { getWorkspacePath } from '@opencode-manager/shared/config/env' |
| 8 | +import { broadcastSSHHostKeyRequest } from '../services/sse-aggregator' |
| 9 | +import { executeCommand } from '../utils/process' |
| 10 | +import { parseSSHHost, normalizeHostPort } from '../utils/ssh-key-manager' |
| 11 | + |
| 12 | +interface SSHHostKeyRequest { |
| 13 | + id: string |
| 14 | + host: string |
| 15 | + ip: string |
| 16 | + keyType: string |
| 17 | + fingerprint: string |
| 18 | + timestamp: number |
| 19 | + isKeyChanged: boolean |
| 20 | +} |
| 21 | + |
| 22 | +export class SSHHostKeyHandler implements IPCHandler { |
| 23 | + private pendingRequests = new Map<string, { |
| 24 | + request: SSHHostKeyRequest |
| 25 | + resolve: (value: boolean) => void |
| 26 | + timeout: ReturnType<typeof setTimeout> |
| 27 | + }>() |
| 28 | + private readonly timeoutMs: number |
| 29 | + private knownHostsPath: string |
| 30 | + private database: Database |
| 31 | + |
| 32 | + constructor(database: Database, timeoutMs: number = 120_000) { |
| 33 | + this.database = database |
| 34 | + this.timeoutMs = timeoutMs |
| 35 | + const configDir = path.join(getWorkspacePath(), 'config') |
| 36 | + this.knownHostsPath = path.join(configDir, 'known_hosts') |
| 37 | + this.ensureKnownHostsFile() |
| 38 | + logger.info(`SSHHostKeyHandler initialized with timeout=${timeoutMs}ms, known_hosts=${this.knownHostsPath}`) |
| 39 | + } |
| 40 | + |
| 41 | + private async ensureKnownHostsFile(): Promise<void> { |
| 42 | + try { |
| 43 | + const configDir = path.join(getWorkspacePath(), 'config') |
| 44 | + await fs.mkdir(configDir, { recursive: true }) |
| 45 | + try { |
| 46 | + await fs.access(this.knownHostsPath) |
| 47 | + } catch { |
| 48 | + await fs.writeFile(this.knownHostsPath, '', { mode: 0o600 }) |
| 49 | + logger.info(`Created known_hosts file at ${this.knownHostsPath}`) |
| 50 | + } |
| 51 | + } catch (error) { |
| 52 | + logger.error('Failed to ensure known_hosts file:', error) |
| 53 | + } |
| 54 | + } |
| 55 | + |
| 56 | + async verifyHostKeyBeforeOperation(repoUrl: string): Promise<boolean> { |
| 57 | + const { host, port } = parseSSHHost(repoUrl) |
| 58 | + const hostPort = normalizeHostPort(host, port) |
| 59 | + |
| 60 | + const trustedHost = this.getTrustedHost(hostPort) |
| 61 | + if (trustedHost) { |
| 62 | + logger.info(`Host ${hostPort} already trusted, skipping verification`) |
| 63 | + return true |
| 64 | + } |
| 65 | + |
| 66 | + try { |
| 67 | + const publicKey = await this.fetchHostPublicKey(host, port) |
| 68 | + logger.info(`Fetched public key for ${hostPort}`) |
| 69 | + |
| 70 | + const parts = publicKey.split(' ') |
| 71 | + const keyType = parts[1] || 'UNKNOWN' |
| 72 | + const requestId = crypto.randomBytes(16).toString('hex') |
| 73 | + const hostKeyRequest: SSHHostKeyRequest = { |
| 74 | + id: requestId, |
| 75 | + host: hostPort, |
| 76 | + ip: '', |
| 77 | + keyType, |
| 78 | + fingerprint: publicKey, |
| 79 | + timestamp: Date.now(), |
| 80 | + isKeyChanged: false |
| 81 | + } |
| 82 | + |
| 83 | + logger.info(`Broadcasting SSH host key request: ${requestId} for host=${hostPort}`) |
| 84 | + broadcastSSHHostKeyRequest({ ...hostKeyRequest, requestId, action: 'verify' }) |
| 85 | + |
| 86 | + return new Promise<boolean>((resolve) => { |
| 87 | + const timeout = setTimeout(() => { |
| 88 | + logger.info(`SSH host key request timed out: ${requestId}, rejecting connection`) |
| 89 | + this.pendingRequests.delete(requestId) |
| 90 | + resolve(false) |
| 91 | + }, this.timeoutMs) |
| 92 | + |
| 93 | + this.pendingRequests.set(requestId, { request: hostKeyRequest, resolve, timeout }) |
| 94 | + }) |
| 95 | + } catch (error) { |
| 96 | + logger.warn(`Failed to fetch host key for ${hostPort}, rejecting connection:`, (error as Error).message) |
| 97 | + return false |
| 98 | + } |
| 99 | + } |
| 100 | + |
| 101 | + private async fetchHostPublicKey(host: string, port?: string): Promise<string> { |
| 102 | + const portArgs = port ? ['-p', port] : [] |
| 103 | + const output = await executeCommand(['ssh-keyscan', '-t', 'ed25519,rsa,ecdsa', ...portArgs, host], { silent: true }) |
| 104 | + |
| 105 | + const bracketedHost = port && port !== '22' ? `[${host}]:${port}` : host |
| 106 | + const lines = output.trim().split('\n') |
| 107 | + for (const line of lines) { |
| 108 | + if (line.startsWith(host) || line.startsWith(bracketedHost)) { |
| 109 | + return line |
| 110 | + } |
| 111 | + } |
| 112 | + |
| 113 | + throw new Error('No valid host keys found') |
| 114 | + } |
| 115 | + |
| 116 | + async handle(request: unknown): Promise<unknown> { |
| 117 | + const response = request as { requestId: string; response: 'accept' | 'reject' } |
| 118 | + return await this.respond(response) |
| 119 | + } |
| 120 | + |
| 121 | + async respond(response: { requestId: string; response: 'accept' | 'reject' }): Promise<{ success: boolean; error?: string }> { |
| 122 | + const pending = this.pendingRequests.get(response.requestId) |
| 123 | + if (!pending) { |
| 124 | + return { success: false, error: 'Request not found or expired' } |
| 125 | + } |
| 126 | + |
| 127 | + clearTimeout(pending.timeout) |
| 128 | + this.pendingRequests.delete(response.requestId) |
| 129 | + |
| 130 | + if (response.response === 'accept') { |
| 131 | + await this.addToKnownHosts(pending.request.host, pending.request.fingerprint) |
| 132 | + this.saveTrustedHost(pending.request.host, pending.request.fingerprint) |
| 133 | + logger.info(`Accepted SSH host key for ${pending.request.host}`) |
| 134 | + } else { |
| 135 | + logger.info(`Rejected SSH host key for ${pending.request.host}`) |
| 136 | + } |
| 137 | + |
| 138 | + pending.resolve(response.response === 'accept') |
| 139 | + return { success: true } |
| 140 | + } |
| 141 | + |
| 142 | + private async addToKnownHosts(host: string, publicKey: string): Promise<void> { |
| 143 | + try { |
| 144 | + await fs.appendFile(this.knownHostsPath, publicKey + '\n') |
| 145 | + logger.info(`Added host to known_hosts: ${host}`) |
| 146 | + } catch (error) { |
| 147 | + logger.error(`Failed to add host to known_hosts: ${error}`) |
| 148 | + } |
| 149 | + } |
| 150 | + |
| 151 | + private async loadFromDatabaseToKnownHosts(): Promise<void> { |
| 152 | + try { |
| 153 | + const hosts = this.database.prepare('SELECT * FROM trusted_ssh_hosts').all() as Array<{ |
| 154 | + id: number |
| 155 | + host: string |
| 156 | + key_type: string |
| 157 | + public_key: string |
| 158 | + created_at: number |
| 159 | + updated_at: number |
| 160 | + }> |
| 161 | + |
| 162 | + const entries = hosts.map(h => h.public_key).join('\n') |
| 163 | + await fs.writeFile(this.knownHostsPath, entries + '\n', { mode: 0o600 }) |
| 164 | + logger.info(`Loaded ${hosts.length} trusted hosts from database to known_hosts`) |
| 165 | + } catch (error) { |
| 166 | + logger.error('Failed to load trusted hosts from database:', error) |
| 167 | + } |
| 168 | + } |
| 169 | + |
| 170 | + private getTrustedHost(host: string): { key_type: string; public_key: string } | null { |
| 171 | + try { |
| 172 | + const result = this.database.prepare('SELECT key_type, public_key FROM trusted_ssh_hosts WHERE host = ?').get(host) as { |
| 173 | + key_type: string |
| 174 | + public_key: string |
| 175 | + } | undefined |
| 176 | + return result || null |
| 177 | + } catch (error) { |
| 178 | + logger.error(`Failed to get trusted host ${host}:`, error) |
| 179 | + return null |
| 180 | + } |
| 181 | + } |
| 182 | + |
| 183 | + private saveTrustedHost(host: string, publicKey: string): void { |
| 184 | + try { |
| 185 | + const parts = publicKey.split(' ') |
| 186 | + const keyType = parts[1] || 'UNKNOWN' |
| 187 | + const now = Date.now() |
| 188 | + const existing = this.getTrustedHost(host) |
| 189 | + if (existing) { |
| 190 | + this.database.prepare('UPDATE trusted_ssh_hosts SET key_type = ?, public_key = ?, updated_at = ? WHERE host = ?') |
| 191 | + .run(keyType, publicKey, now, host) |
| 192 | + logger.info(`Updated trusted host in database: ${host}`) |
| 193 | + } else { |
| 194 | + this.database.prepare('INSERT INTO trusted_ssh_hosts (host, key_type, public_key, created_at, updated_at) VALUES (?, ?, ?, ?, ?)') |
| 195 | + .run(host, keyType, publicKey, now, now) |
| 196 | + logger.info(`Saved new trusted host to database: ${host}`) |
| 197 | + } |
| 198 | + } catch (error) { |
| 199 | + logger.error(`Failed to save trusted host ${host}:`, error) |
| 200 | + } |
| 201 | + } |
| 202 | + |
| 203 | + async initialize(): Promise<void> { |
| 204 | + await this.ensureKnownHostsFile() |
| 205 | + await this.loadFromDatabaseToKnownHosts() |
| 206 | + logger.info('SSHHostKeyHandler initialized with known_hosts from database') |
| 207 | + } |
| 208 | + |
| 209 | + getKnownHostsPath(): string { |
| 210 | + return this.knownHostsPath |
| 211 | + } |
| 212 | + |
| 213 | + getEnv(): Record<string, string> { |
| 214 | + return { |
| 215 | + KNOWN_HOSTS_PATH: this.knownHostsPath |
| 216 | + } |
| 217 | + } |
| 218 | + |
| 219 | + getPendingCount(): number { |
| 220 | + return this.pendingRequests.size |
| 221 | + } |
| 222 | +} |
| 223 | + |
| 224 | +export function createSSHHostKeyHandler(database: Database, timeoutMs?: number): SSHHostKeyHandler { |
| 225 | + return new SSHHostKeyHandler(database, timeoutMs) |
| 226 | +} |
0 commit comments