diff --git a/src/main/managed-data-accounts/service-enrollment-persistence.test.ts b/src/main/managed-data-accounts/service-enrollment-persistence.test.ts new file mode 100644 index 00000000000..9375071d2f1 --- /dev/null +++ b/src/main/managed-data-accounts/service-enrollment-persistence.test.ts @@ -0,0 +1,209 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import * as filesystem from 'node:fs' +import { + existsSync, + mkdirSync, + mkdtempSync, + readFileSync, + readdirSync, + rmSync, + writeFileSync +} from 'node:fs' +import { tmpdir } from 'node:os' +import { join } from 'node:path' +import SyncDatabase from '../sqlite/sync-database' +import * as secureFile from '../../shared/secure-file' +import { ManagedDataAccountService } from './service' + +vi.mock('node:fs', async (importOriginal) => { + const actual = await importOriginal() + return { ...actual } +}) + +let root: string +let source: string +let service: ManagedDataAccountService + +beforeEach(() => { + root = mkdtempSync(join(tmpdir(), 'orca-enrollment-persistence-')) + source = join(root, 'source') + mkdirSync(join(source, 'devin'), { recursive: true }) + writeFileSync(join(source, 'devin', 'credentials.toml'), 'windsurf_api_key = "test-only-key"\n') + mkdirSync(join(source, 'opencode'), { recursive: true }) + const database = new SyncDatabase(join(source, 'opencode', 'opencode.db')) + database.exec( + 'CREATE TABLE session_v2 (id TEXT); CREATE TABLE credential (integration_id TEXT, value TEXT)' + ) + database + .prepare('INSERT INTO credential VALUES (?, ?)') + .run('google', '{"type":"key","key":"test-only-key"}') + database.close() + service = new ManagedDataAccountService(join(root, 'managed')) +}) + +afterEach(() => { + vi.restoreAllMocks() + rmSync(root, { recursive: true, force: true }) +}) + +describe.each(['opencode', 'devin'] as const)('committed %s enrollment', (provider) => { + it.each(['after write', 'unrestricted'])( + 'keeps captured credentials when metadata persistence reports %s failure', + async (failure) => { + const metadata = join(root, 'managed', provider, 'accounts.json') + const write = secureFile.writeSecureFile + vi.spyOn(secureFile, 'writeSecureFile').mockImplementation((...args) => { + const result = write(...args) + if (args[0] !== metadata) { + return result + } + if (failure === 'unrestricted') { + return false + } + throw new Error('post-publication write failure') + }) + await expect(service.add(provider, source, 'Work')).rejects.toThrow( + failure === 'unrestricted' ? 'metadata permissions' : 'post-publication' + ) + const state = service.list(provider) + expect(state.accounts).toHaveLength(1) + expect(state.activeAccountId).toBe(state.accounts[0].id) + const directory = join(root, 'managed', provider, state.accounts[0].id) + expect(existsSync(directory)).toBe(true) + expect(service.launchEnvironment(provider).XDG_DATA_HOME).toBe(join(directory, 'data')) + } + ) + + it('cleans an unregistered profile when metadata fails before publication', async () => { + const metadata = join(root, 'managed', provider, 'accounts.json') + const write = secureFile.writeSecureFile + vi.spyOn(secureFile, 'writeSecureFile').mockImplementation((...args) => { + if (args[0] === metadata) { + throw new Error('pre-publication write failure') + } + return write(...args) + }) + await expect(service.add(provider, source, 'Work')).rejects.toThrow('pre-publication') + expect(service.list(provider)).toEqual({ accounts: [], activeAccountId: null }) + expect(readdirSync(join(root, 'managed', provider))).toEqual([]) + }) + + it('keeps credentials when published metadata is unreadable', async () => { + const metadata = join(root, 'managed', provider, 'accounts.json') + const write = secureFile.writeSecureFile + let published = '' + vi.spyOn(secureFile, 'writeSecureFile').mockImplementation((...args) => { + const result = write(...args) + if (args[0] !== metadata) { + return result + } + published = readFileSync(metadata, 'utf8') + writeFileSync(metadata, '{') + throw new Error('post-publication metadata unreadable') + }) + await expect(service.add(provider, source, 'Work')).rejects.toThrow('metadata unreadable') + expect( + readdirSync(join(root, 'managed', provider)).filter((name) => name !== 'accounts.json') + ).toHaveLength(1) + writeFileSync(metadata, published) + expect(service.launchEnvironment(provider).XDG_DATA_HOME).toBeTruthy() + }) + + it('keeps credentials when metadata existence cannot be checked', async () => { + const metadata = join(root, 'managed', provider, 'accounts.json') + const write = secureFile.writeSecureFile + const exists = filesystem.existsSync + const stat = filesystem.lstatSync + let inaccessible = false + vi.spyOn(filesystem, 'existsSync').mockImplementation((path) => + inaccessible && path === metadata ? false : exists(path) + ) + vi.spyOn(filesystem, 'lstatSync').mockImplementation((...args) => { + if (inaccessible && args[0] === metadata) { + throw Object.assign(new Error('Metadata access denied'), { code: 'EACCES' }) + } + return stat(...args) + }) + vi.spyOn(secureFile, 'writeSecureFile').mockImplementation((...args) => { + const result = write(...args) + if (args[0] !== metadata) { + return result + } + inaccessible = true + return false + }) + await expect(service.add(provider, source, 'Work')).rejects.toThrow('metadata permissions') + inaccessible = false + const state = service.list(provider) + expect(exists(join(root, 'managed', provider, state.accounts[0].id))).toBe(true) + expect(service.launchEnvironment(provider).XDG_DATA_HOME).toBeTruthy() + }) + + it('preserves a profile registered with another UUID case', async () => { + const metadata = join(root, 'managed', provider, 'accounts.json') + const write = secureFile.writeSecureFile + const writer = vi.spyOn(secureFile, 'writeSecureFile').mockImplementation((...args) => { + const result = write(...args) + if (args[0] !== metadata) { + return result + } + const state = service.list(provider) + writeFileSync( + metadata, + JSON.stringify({ + ...state, + accounts: state.accounts.map((account) => ({ ...account, id: account.id.toUpperCase() })), + activeAccountId: state.activeAccountId?.toUpperCase() + }) + ) + throw new Error('post-publication UUID case change') + }) + await expect(service.add(provider, source, 'Work')).rejects.toThrow('UUID case change') + writer.mockRestore() + const registered = service.list(provider).accounts[0] + const dataHome = join(root, 'managed', provider, registered.id.toLowerCase(), 'data') + expect(existsSync(dataHome)).toBe(true) + expect(service.launchEnvironment(provider).XDG_DATA_HOME).toBe(dataHome) + expect( + service.transcriptEnvironments(provider).map((environment) => environment.XDG_DATA_HOME) + ).toEqual([dataHome]) + expect((await service.select(provider, registered.id.toLowerCase())).activeAccountId).toBe( + registered.id + ) + expect(service.launchEnvironment(provider).XDG_DATA_HOME).toBe(dataHome) + expect((await service.select(provider, registered.id.toUpperCase())).activeAccountId).toBe( + registered.id + ) + expect(service.launchEnvironment(provider).XDG_DATA_HOME).toBe(dataHome) + }) + + it('selects the registered UUID spelling and searches its transcript first', async () => { + const personal = (await service.add(provider, source, 'Personal')).accounts[0] + const work = (await service.add(provider, source, 'Work')).accounts[1] + const selected = await service.select(provider, work.id.toUpperCase()) + expect(selected.activeAccountId).toBe(work.id) + expect(service.list(provider).activeAccountId).toBe(work.id) + expect( + service.transcriptEnvironments(provider).map((environment) => environment.XDG_DATA_HOME) + ).toEqual( + [work, personal].map((account) => + join(root, 'managed', provider, account.id.toLowerCase(), 'data') + ) + ) + }) + + it('isolates a throwing listener after enrollment has committed', async () => { + const warn = vi.spyOn(console, 'warn').mockImplementation(() => {}) + service.onChanged(() => { + throw new Error('private-listener-detail') + }) + const healthy = vi.fn() + service.onChanged(healthy) + const state = await service.add(provider, source, 'Work') + expect(service.list(provider)).toEqual(state) + expect(service.launchEnvironment(provider).XDG_DATA_HOME).toBeTruthy() + expect(healthy).toHaveBeenCalledOnce() + expect(warn).toHaveBeenCalledOnce() + expect(warn.mock.calls.flat().join(' ')).not.toContain('private-listener-detail') + }) +}) diff --git a/src/main/managed-data-accounts/service.test.ts b/src/main/managed-data-accounts/service.test.ts index 5ab7f12d04c..bc96800a5aa 100644 --- a/src/main/managed-data-accounts/service.test.ts +++ b/src/main/managed-data-accounts/service.test.ts @@ -491,15 +491,18 @@ describe('managed data accounts', () => { expect(service.transcriptEnvironments('devin')).toEqual([secondEnvironment]) }) - it('rejects a credential symlink without touching its target', async () => { - const original = join(source, 'devin', 'credentials.toml') - const target = join(root, 'private.toml') - writeFileSync(target, readFileSync(original)) - rmSync(original) - symlinkSync(target, original) - await expect(service.add('devin', source, 'Work')).rejects.toThrow('regular file') - expect(readFileSync(target, 'utf8')).toContain('test-only-key') - }) + it.skipIf(process.platform === 'win32')( + 'rejects a credential symlink without touching its target', + async () => { + const original = join(source, 'devin', 'credentials.toml') + const target = join(root, 'private.toml') + writeFileSync(target, readFileSync(original)) + rmSync(original) + symlinkSync(target, original) + await expect(service.add('devin', source, 'Work')).rejects.toThrow('regular file') + expect(readFileSync(target, 'utf8')).toContain('test-only-key') + } + ) it('keeps credential parse errors out of RPC messages', async () => { writeFileSync( diff --git a/src/main/managed-data-accounts/service.ts b/src/main/managed-data-accounts/service.ts index 5ebf2ae91f2..749fb023876 100644 --- a/src/main/managed-data-accounts/service.ts +++ b/src/main/managed-data-accounts/service.ts @@ -69,8 +69,7 @@ export class ManagedDataAccountService { if (!existsSync(path)) { return { accounts: [], activeAccountId: null } } - this.assertOwned(path) - return stateSchema.parse(JSON.parse(readFileSync(path, 'utf8'))) + return this.readState(path) } add( @@ -98,7 +97,24 @@ export class ManagedDataAccountService { activeAccountId: id }) } catch (error) { - rmSync(directory, { recursive: true, force: true }) + let registered = true + try { + registered = this.readState(join(this.root, provider, 'accounts.json')).accounts.some( + (account) => account.id.toLowerCase() === id.toLowerCase() + ) + } catch (metadataError) { + if ( + metadataError instanceof Error && + 'code' in metadataError && + metadataError.code === 'ENOENT' + ) { + registered = false + } + // Unreadable metadata cannot prove that this profile is unregistered. + } + if (!registered) { + rmSync(directory, { recursive: true, force: true }) + } throw error } }) @@ -110,10 +126,9 @@ export class ManagedDataAccountService { ): Promise { return this.mutate(async () => { const state = this.list(provider) - if (accountId !== null) { - this.requireAccount(provider, accountId) - } - return this.persist(provider, { ...state, activeAccountId: accountId }) + const activeAccountId = + accountId === null ? null : this.requireAccount(provider, accountId).id + return this.persist(provider, { ...state, activeAccountId }) }) } @@ -206,7 +221,7 @@ export class ManagedDataAccountService { provider: ManagedDataAccountProvider, accountId: string ): Record { - const directory = this.requireAccount(provider, accountId) + const { directory } = this.requireAccount(provider, accountId) return { XDG_DATA_HOME: join(directory, 'data'), XDG_STATE_HOME: join(directory, 'state'), @@ -219,13 +234,19 @@ export class ManagedDataAccountService { return () => this.listeners.delete(listener) } - private requireAccount(provider: ManagedDataAccountProvider, id: string): string { - if (!this.list(provider).accounts.some((account) => account.id === id)) { + private requireAccount( + provider: ManagedDataAccountProvider, + id: string + ): { id: string; directory: string } { + const account = this.list(provider).accounts.find( + (registered) => registered.id.toLowerCase() === id.toLowerCase() + ) + if (!account) { throw new Error('Managed account not found.') } - const directory = join(this.root, provider, id) + const directory = join(this.root, provider, account.id.toLowerCase()) this.assertOwned(directory) - return directory + return { id: account.id, directory } } private persist( @@ -237,6 +258,11 @@ export class ManagedDataAccountService { return checked } + private readState(path: string): ManagedDataAccountsState { + this.assertOwned(path) + return stateSchema.parse(JSON.parse(readFileSync(path, 'utf8'))) + } + private writeState( provider: ManagedDataAccountProvider, state: ManagedDataAccountsState @@ -254,7 +280,11 @@ export class ManagedDataAccountService { private notifyChanged(): void { for (const listener of this.listeners) { - listener() + try { + listener() + } catch { + console.warn('[managed-data-accounts] Account change listener failed.') + } } }