Skip to content
33 changes: 32 additions & 1 deletion apps/cli/src/api/apiMachine.transports.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ import {
import { logger } from '@/ui/logger';
import type { Machine } from './types';

const { configurationMock, mockAxiosGet, mockAxiosIsAxiosError, mockAxiosPost, mockIo } = vi.hoisted(() => ({
const { configurationMock, mockAxiosGet, mockAxiosIsAxiosError, mockAxiosPost, mockIo, rpcHandlerConfigs } = vi.hoisted(() => ({
configurationMock: {
apiServerUrl: 'http://localhost:3005',
activeServerDir: '',
Expand All @@ -29,6 +29,7 @@ const { configurationMock, mockAxiosGet, mockAxiosIsAxiosError, mockAxiosPost, m
emitWithAck: vi.fn(),
io: { on: vi.fn() },
})),
rpcHandlerConfigs: [] as Array<Record<string, unknown>>,
}));

vi.mock('socket.io-client', () => ({
Expand Down Expand Up @@ -63,6 +64,9 @@ vi.mock('@/rpc/handlers/machineFileBrowser/registerMachineFileBrowserHandlers',
vi.mock('./machine/rpcHandlers', () => ({ registerMachineRpcHandlers: vi.fn() }));
vi.mock('./rpc/RpcHandlerManager', () => ({
RpcHandlerManager: class {
constructor(config: Record<string, unknown>) {
rpcHandlerConfigs.push(config);
}
registerHandler() {}
onSocketConnect() {}
onSocketDisconnect() {}
Expand All @@ -88,6 +92,7 @@ describe('ApiMachineClient transports', () => {
mockAxiosPost.mockResolvedValue({ status: 200, data: { success: true, applied: true } });
mockAxiosGet.mockResolvedValue({ status: 200, data: { machine: null } });
bindApiSessionSocketMock(mockIo, createApiSessionSocketStub());
rpcHandlerConfigs.length = 0;
});

afterEach(() => {
Expand Down Expand Up @@ -123,6 +128,32 @@ describe('ApiMachineClient transports', () => {
expect(opts.autoConnect).toBe(false);
});

it('configures machine RPC to project only strict completed-stop transport proof', async () => {
const mod = await import('./apiMachine');
const { RPC_METHODS } = await import('@happier-dev/protocol/rpc');

new mod.ApiMachineClient('fake-token', {
id: 'test-machine',
encryptionKey: new Uint8Array(32),
encryptionVariant: 'legacy',
metadata: null,
metadataVersion: 0,
daemonState: null,
daemonStateVersion: 0,
});

const projector = rpcHandlerConfigs.at(-1)?.projectTransportAcknowledgement;
expect(projector).toBeTypeOf('function');
expect((projector as (input: { method: string; result: unknown }) => unknown)({
method: `test-machine:${RPC_METHODS.STOP_SESSION}`,
result: { status: 'stopped' },
})).toEqual({ kind: 'session.stop', status: 'stopped' });
expect((projector as (input: { method: string; result: unknown }) => unknown)({
method: `test-machine:${RPC_METHODS.STOP_SESSION}`,
result: { status: 'requested' },
})).toBeNull();
});

it('serializes machine refresh errors without dumping axios request details', async () => {
const mod = await import('./apiMachine');

Expand Down
2 changes: 2 additions & 0 deletions apps/cli/src/api/apiMachine.ts
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ import { recoverDaemonTerminalSessionMutationJournals } from './session/mutation

import type { DaemonToServerEvents, ServerToDaemonEvents } from './machine/socketTypes';
import { authorizeMachineRpcRequest } from './machine/machineRpcAuthorization';
import { projectMachineRpcTransportAcknowledgement } from './machine/projectMachineRpcTransportAcknowledgement';
import { registerMachineRpcHandlers, type MachineRpcHandlerDeps, type MachineRpcHandlers } from './machine/rpcHandlers';
import { resolveMachineRpcWorkingDirectory } from './machine/resolveMachineRpcWorkingDirectory';
import type { Socket } from 'socket.io-client';
Expand Down Expand Up @@ -255,6 +256,7 @@ export class ApiMachineClient {
}
},
authorizeRequest: authorizeMachineRpcRequest,
projectTransportAcknowledgement: projectMachineRpcTransportAcknowledgement,
});

const machineRpcWorkingDirectory = resolveMachineRpcWorkingDirectory();
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
import { describe, expect, it } from 'vitest';
import { RPC_METHODS } from '@happier-dev/protocol/rpc';

import { projectMachineRpcTransportAcknowledgement } from './projectMachineRpcTransportAcknowledgement';

describe('projectMachineRpcTransportAcknowledgement', () => {
it('projects proof only for a strict completed stop result', () => {
expect(projectMachineRpcTransportAcknowledgement({
method: `machine-1:${RPC_METHODS.STOP_SESSION}`,
result: { status: 'stopped' },
})).toEqual({ kind: 'session.stop', status: 'stopped' });

for (const result of [
{ status: 'requested' },
{ status: 'not_found' },
{ status: 'incomplete', reason: 'runner_exit_timeout' },
]) {
expect(projectMachineRpcTransportAcknowledgement({
method: `machine-1:${RPC_METHODS.STOP_SESSION}`,
result,
})).toBeNull();
}
});

it('does not project proof for another method or a lookalike result', () => {
expect(projectMachineRpcTransportAcknowledgement({
method: 'machine-1:other-method',
result: { status: 'stopped' },
})).toBeNull();
expect(projectMachineRpcTransportAcknowledgement({
method: `machine-1:${RPC_METHODS.STOP_SESSION}`,
result: { status: 'stopped', extra: true },
})).toBeNull();
});
});
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
import { StopSessionResultSchema } from '@happier-dev/protocol';
import {
RPC_METHODS,
resolveSocketRpcSessionWriteAuthorizationMethod,
} from '@happier-dev/protocol/rpc';
import type { SocketRpcTransportAcknowledgementV1 } from '@happier-dev/protocol/socketRpc';

export function projectMachineRpcTransportAcknowledgement(input: Readonly<{
method: string;
result: unknown;
}>): SocketRpcTransportAcknowledgementV1 | null {
if (
resolveSocketRpcSessionWriteAuthorizationMethod(input.method)
!== RPC_METHODS.STOP_SESSION
) {
return null;
}
const parsed = StopSessionResultSchema.safeParse(input.result);
return parsed.success && parsed.data.status === 'stopped'
? { kind: 'session.stop', status: 'stopped' }
: null;
}
39 changes: 39 additions & 0 deletions apps/cli/src/api/rpc/RpcHandlerManager.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,45 @@ describe('RpcHandlerManager.handleRequest (plaintext)', () => {
});

describe('RpcHandlerManager.handleRequest (encrypted)', () => {
it('keeps the encrypted result opaque while exposing requested transport acknowledgement metadata', async () => {
const encryptionKey = new Uint8Array(32).fill(3);
const rpc = new RpcHandlerManager({
scopePrefix: 'machine-1',
encryptionKey,
encryptionVariant: 'dataKey',
logger: () => {},
projectTransportAcknowledgement: ({ method, result }) => (
method === 'machine-1:stop-session'
&& typeof result === 'object'
&& result !== null
&& (result as { status?: unknown }).status === 'stopped'
? { kind: 'session.stop' as const, status: 'stopped' as const }
: null
),
});

rpc.registerHandler('stop-session', async () => ({ status: 'stopped' }));

const res = await rpc.handleRequest({
method: 'machine-1:stop-session',
params: encodeBase64(encrypt(encryptionKey, 'dataKey', { sessionId: 'sess_1' })),
transportResponseEnvelopeVersion: 1,
});

expect(res).toEqual({
v: 1,
result: expect.any(String),
acknowledgement: { kind: 'session.stop', status: 'stopped' },
});
expect(
decrypt(
encryptionKey,
'dataKey',
decodeBase64((res as { result: string }).result),
),
).toEqual({ status: 'stopped' });
});

it('rejects encrypted requests when the authorization hook rejects the decrypted params', async () => {
const encryptionKey = new Uint8Array(32).fill(5);
const authorizeRequest = vi.fn(() => ({
Expand Down
80 changes: 72 additions & 8 deletions apps/cli/src/api/rpc/RpcHandlerManager.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,11 @@ import {
type RpcAuthorizationResult,
} from './types';
import { Socket } from 'socket.io-client';
import { SOCKET_RPC_EVENTS } from '@happier-dev/protocol/socketRpc';
import {
SOCKET_RPC_EVENTS,
SOCKET_RPC_TRANSPORT_RESPONSE_ENVELOPE_VERSION_V1,
type SocketRpcTransportAcknowledgementV1,
} from '@happier-dev/protocol/socketRpc';
import {
RPC_ERROR_CODES,
RPC_ERROR_MESSAGES,
Expand All @@ -30,6 +34,7 @@ export class RpcHandlerManager {
private readonly logger: (message: string, data?: any) => void;
private readonly onRegistrationError: RpcHandlerConfig['onRegistrationError'];
private readonly authorizeRequest: RpcHandlerConfig['authorizeRequest'];
private readonly projectTransportAcknowledgement: RpcHandlerConfig['projectTransportAcknowledgement'];
private socket: Socket | null = null;
private inFlightRequestCount = 0;
private idleResolvers = new Set<() => void>();
Expand All @@ -42,6 +47,7 @@ export class RpcHandlerManager {
this.logger = config.logger || ((msg, data) => defaultLogger.debug(msg, data));
this.onRegistrationError = config.onRegistrationError;
this.authorizeRequest = config.authorizeRequest;
this.projectTransportAcknowledgement = config.projectTransportAcknowledgement;
}

private encodeResponse(response: unknown): unknown {
Expand Down Expand Up @@ -83,7 +89,7 @@ export class RpcHandlerManager {
if (!handler) {
this.logger('[RPC] [ERROR] Method not found', { method: request.method });
const errorResponse = { error: RPC_ERROR_MESSAGES.METHOD_NOT_FOUND, errorCode: RPC_ERROR_CODES.METHOD_NOT_FOUND };
return this.encodeResponse(errorResponse);
return this.encodeTransportResponse(request, errorResponse);
}

// Decrypt the incoming params (unless session is plaintext).
Expand All @@ -96,7 +102,7 @@ export class RpcHandlerManager {
const errorResponse = {
error: 'Invalid RPC params',
};
return this.encodeResponse(errorResponse);
return this.encodeTransportResponse(request, errorResponse);
}

const authorizationResult: RpcAuthorizationResult = this.authorizeRequest
Expand All @@ -107,7 +113,7 @@ export class RpcHandlerManager {
})
: { ok: true };
if (authorizationResult.ok !== true) {
return this.encodeResponse({
return this.encodeTransportResponse(request, {
error: authorizationResult.error,
...(authorizationResult.errorCode ? { errorCode: authorizationResult.errorCode } : {}),
});
Expand All @@ -119,17 +125,28 @@ export class RpcHandlerManager {
this.logger('[RPC] Handler returned', { method: request.method, hasResult: result !== undefined });

// Encrypt and return the response
const response = this.encodeResponse(result);
if (this.encryptionMode !== 'plain' && typeof response === 'string') {
this.logger('[RPC] Sending encrypted response', { method: request.method, responseLength: response.length });
const acknowledgement = this.projectAcknowledgement(request, decryptedParams, result);
const response = this.encodeTransportResponse(request, result, acknowledgement);
if (this.encryptionMode !== 'plain') {
const encodedResult = request.transportResponseEnvelopeVersion
=== SOCKET_RPC_TRANSPORT_RESPONSE_ENVELOPE_VERSION_V1
&& response
&& typeof response === 'object'
&& !Array.isArray(response)
? (response as { result?: unknown }).result
: response;
this.logger('[RPC] Sending encrypted response', {
method: request.method,
responseLength: typeof encodedResult === 'string' ? encodedResult.length : 0,
});
}
return response;
} catch (error) {
this.logger('[RPC] [ERROR] Error handling request', { error });
const errorResponse = {
error: error instanceof Error ? error.message : 'Unknown error'
};
return this.encodeResponse(errorResponse);
return this.encodeTransportResponse(request, errorResponse);
} finally {
this.finishInFlightRequest();
}
Expand Down Expand Up @@ -224,6 +241,53 @@ export class RpcHandlerManager {
return `${this.scopePrefix}:${method}`;
}

private encodeTransportResponse(
request: RpcRequest,
result: unknown,
acknowledgement: SocketRpcTransportAcknowledgementV1 | null = null,
): unknown {
const encodedResult = this.encodeResponse(result);
if (
request.transportResponseEnvelopeVersion
!== SOCKET_RPC_TRANSPORT_RESPONSE_ENVELOPE_VERSION_V1
) {
return encodedResult;
}
return {
v: SOCKET_RPC_TRANSPORT_RESPONSE_ENVELOPE_VERSION_V1,
result: encodedResult,
...(acknowledgement ? { acknowledgement } : {}),
};
}

private projectAcknowledgement(
request: RpcRequest,
params: unknown,
result: unknown,
): SocketRpcTransportAcknowledgementV1 | null {
if (
request.transportResponseEnvelopeVersion
!== SOCKET_RPC_TRANSPORT_RESPONSE_ENVELOPE_VERSION_V1
|| !this.projectTransportAcknowledgement
) {
return null;
}
try {
return this.projectTransportAcknowledgement({
method: request.method,
params,
result,
...(request.authorization ? { authorization: request.authorization } : {}),
});
} catch (error) {
this.logger('[RPC] Transport acknowledgement projection failed', {
method: request.method,
error,
});
return null;
}
}

private beginInFlightRequest(): void {
this.inFlightRequestCount += 1;
}
Expand Down
16 changes: 11 additions & 5 deletions apps/cli/src/api/rpc/types.ts
Original file line number Diff line number Diff line change
@@ -1,4 +1,8 @@
import type { SocketRpcAuthorizationContext } from '@happier-dev/protocol';
import type {
SocketRpcRequestPayload,
SocketRpcTransportAcknowledgementV1,
} from '@happier-dev/protocol/socketRpc';

/**
* Common RPC types and interfaces for both session and machine clients
Expand Down Expand Up @@ -34,11 +38,7 @@ export type RpcHandlerMap = Map<string, RpcHandler>;
/**
* RPC request data from server
*/
export interface RpcRequest {
method: string;
params: unknown;
authorization?: SocketRpcAuthorizationContext;
}
export type RpcRequest = SocketRpcRequestPayload;

/**
* RPC response callback
Expand All @@ -60,6 +60,12 @@ export interface RpcHandlerConfig {
params: unknown;
authorization?: SocketRpcAuthorizationContext;
}>) => RpcAuthorizationResult | Promise<RpcAuthorizationResult>;
projectTransportAcknowledgement?: (request: Readonly<{
method: string;
params: unknown;
result: unknown;
authorization?: SocketRpcAuthorizationContext;
}>) => SocketRpcTransportAcknowledgementV1 | null;
}

export type RpcAuthorizationResult =
Expand Down
9 changes: 2 additions & 7 deletions apps/cli/src/api/types.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import { z } from 'zod'
import { UsageSchema } from '@/api/usage'
import { SOCKET_RPC_EVENTS } from '@happier-dev/protocol/socketRpc'
import { SOCKET_RPC_EVENTS, type SocketRpcRequestPayload as ProtocolSocketRpcRequestPayload } from '@happier-dev/protocol/socketRpc'
import { SentFromSchema } from '@happier-dev/protocol'
import type { ExecutionRunPublicState } from '@happier-dev/protocol'
import type {
Expand All @@ -19,7 +19,6 @@ import type {
SessionRollbackRangesV1,
SessionUsageLimitRecoveryV1,
SessionTerminalMetadata,
SocketRpcAuthorizationContext,
SessionMessageRole,
ProviderSessionInfoV1,
SessionRuntimeActivityState,
Expand Down Expand Up @@ -140,11 +139,7 @@ export type UpdateMachineBody = Extract<Update['body'], { t: 'update-machine' }>
export const SessionBroadcastSchema = SessionBroadcastContainerSchema
export type SessionBroadcast = SessionBroadcastContainer

export interface SocketRpcRequestPayload {
method: string
params: unknown
authorization?: SocketRpcAuthorizationContext
}
export type SocketRpcRequestPayload = ProtocolSocketRpcRequestPayload

export interface SocketRpcCallPayload extends SocketRpcRequestPayload {
timeoutMs?: number
Expand Down
Loading
Loading