Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
100 changes: 100 additions & 0 deletions packages/control-plane/src/cloudflare/session-platform.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,100 @@
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
import type { SqlDatabase } from "../db/sql-database";
import type { Logger } from "../logger";
import { createDurableObjectSessionPlatform } from "./session-platform";

/** Stand-in for the Workers runtime's request/response pair. */
class FakeRequestResponsePair {
constructor(
readonly request: string,
readonly response: string
) {}
}

function createFakeState() {
const storage = {
sql: { exec: vi.fn() },
transactionSync: vi.fn(<T>(closure: () => T): T => closure()),
getAlarm: vi.fn(async () => null),
setAlarm: vi.fn(async () => {}),
deleteAlarm: vi.fn(async () => {}),
};
const calls = {
id: { toString: () => "do-id" },
storage,
acceptWebSocket: vi.fn(),
getTags: vi.fn(() => ["sandbox", "sid:sb-1"]),
getWebSockets: vi.fn(() => []),
setWebSocketAutoResponse: vi.fn(),
waitUntil: vi.fn(),
};
const db = {} as SqlDatabase;
return { state: calls as unknown as DurableObjectState, storage, calls, db };
}

describe("createDurableObjectSessionPlatform", () => {
beforeEach(() => {
vi.stubGlobal("WebSocketRequestResponsePair", FakeRequestResponsePair);
});
afterEach(() => {
vi.unstubAllGlobals();
});

it("exposes the object's id, storage, alarm store, and the global store", () => {
const { state, storage, db } = createFakeState();

const platform = createDurableObjectSessionPlatform(state, db);

expect(platform.id).toBe("do-id");
expect(platform.storage).toBe(storage);
expect(platform.db).toBe(db);
expect(platform.alarmStore).toBe(storage);
expect(platform.storage.transactionSync(() => 42)).toBe(42);
expect(storage.transactionSync).toHaveBeenCalledTimes(1);
});

it("delegates socket acceptance, tags, and enumeration, passing the tag filter through", () => {
const { state, calls, db } = createFakeState();
const ws = {} as WebSocket;

const { sockets: host } = createDurableObjectSessionPlatform(state, db);
host.accept(ws, ["sandbox", "sid:sb-1"]);
host.sockets();
host.sockets("sandbox");

expect(calls.acceptWebSocket).toHaveBeenCalledWith(ws, ["sandbox", "sid:sb-1"]);
expect(host.tags(ws)).toEqual(["sandbox", "sid:sb-1"]);
expect(calls.getTags).toHaveBeenCalledWith(ws);
expect(calls.getWebSockets.mock.calls).toEqual([[undefined], ["sandbox"]]);
});

it("installs the auto-response as a request/response pair", () => {
const { state, calls, db } = createFakeState();

const platform = createDurableObjectSessionPlatform(state, db);
platform.sockets.setAutoResponse('{"type":"ping"}', '{"type":"pong"}');

expect(calls.setWebSocketAutoResponse).toHaveBeenCalledTimes(1);
const pair = calls.setWebSocketAutoResponse.mock.calls[0][0] as FakeRequestResponsePair;
expect(pair).toBeInstanceOf(FakeRequestResponsePair);
expect(pair.request).toBe('{"type":"ping"}');
expect(pair.response).toBe('{"type":"pong"}');
});

it("builds background tasks over the object's event lifetime that report to the given logger", async () => {
const { state, calls, db } = createFakeState();
const logger = { error: vi.fn() } as unknown as Logger;

const platform = createDurableObjectSessionPlatform(state, db);
platform.createBackgroundTasks(logger).submit(() => Promise.reject(new Error("boom")), {
name: "session.task",
});

expect(calls.waitUntil).toHaveBeenCalledTimes(1);
await calls.waitUntil.mock.calls[0][0];
expect(logger.error).toHaveBeenCalledWith(
"background_task.failed",
expect.objectContaining({ task_name: "session.task" })
);
});
});
29 changes: 29 additions & 0 deletions packages/control-plane/src/cloudflare/session-platform.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
import type { SqlDatabase } from "../db/sql-database";
import type { SessionPlatform } from "../session/platform";
import { createCloudflareBackgroundTasks } from "./background-tasks";

/**
* A Durable Object's storage, hibernatable sockets, alarm, and event lifetime
* as the session platform, over the deployment's global store.
*/
export function createDurableObjectSessionPlatform(
ctx: DurableObjectState,
db: SqlDatabase
): SessionPlatform {
return {
id: ctx.id.toString(),
storage: ctx.storage,
db,
alarmStore: ctx.storage,
sockets: {
accept: (ws, tags) => ctx.acceptWebSocket(ws, tags),
tags: (ws) => ctx.getTags(ws),
sockets: (tag) => ctx.getWebSockets(tag),
// Hibernation-level auto-response: matched by the runtime without
// waking the object.
setAutoResponse: (request, response) =>
ctx.setWebSocketAutoResponse(new WebSocketRequestResponsePair(request, response)),
},
createBackgroundTasks: (log) => createCloudflareBackgroundTasks(ctx, log),
};
}
111 changes: 48 additions & 63 deletions packages/control-plane/src/session/components.ts
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ import { requireRepoSecretsEncryptionKey, requireTokenEncryptionKey } from "../e
import type { Env, ClientInfo } from "../types";
import type { SessionRow } from "./types";
import type { SqlDatabase } from "../db/sql-database";
import type { SessionPlatform } from "./platform";
import { SessionCoreRepository } from "./session-core-repository";
import { SandboxRepository } from "./sandbox-repository";
import { SessionAttachmentRepository } from "./session-attachment-repository";
Expand Down Expand Up @@ -80,7 +81,6 @@ import { CallbackNotificationService } from "./callback-notification-service";
import { UserEnvResolver } from "./user-env-resolver";
import { resolveSessionRepoId } from "./repo-id-resolution";
import { Scheduler } from "../scheduler/scheduler";
import { createCloudflareBackgroundTasks } from "../cloudflare/background-tasks";
import { PresenceService } from "./presence-service";
import { SessionMessageQueue } from "./message-queue";
import { SandboxArtifactEventHandler } from "./sandbox-events/artifact.handler";
Expand Down Expand Up @@ -138,13 +138,6 @@ import { AuthorizationError, AuthorizationService } from "../authorization/servi
*/
const WS_AUTH_TIMEOUT_MS = 30000; // 30 seconds

/** The platform surface the session graph is built over. */
export interface SessionPlatform {
ctx: DurableObjectState;
sql: SqlStorage;
db: SqlDatabase | null;
}

/**
* What the platform adapter (SessionDO) is allowed to touch. Everything else
* stays inside the factory; `internals` exists for integration tests that
Expand Down Expand Up @@ -211,9 +204,16 @@ function resolveExecutionTimeoutMs(

/** Build the session runtime, including authorization verification and lease expiry handling. */
export function createSessionRuntime(platform: SessionPlatform, env: Env): SessionRuntime {
const { ctx, sql, db } = platform;
const durableObjectId = ctx.id.toString();
const transaction = <T>(closure: () => T): T => ctx.storage.transactionSync(closure);
const {
id: durableObjectId,
storage,
db,
alarmStore,
sockets: socketHost,
createBackgroundTasks,
} = platform;
const { sql } = storage;
const transaction = <T>(closure: () => T): T => storage.transactionSync(closure);

// Tier 1 — repositories and alarm persistence (leaves over SqlStorage).
const attachmentRepository = new SessionAttachmentRepository(sql);
Expand Down Expand Up @@ -249,29 +249,27 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
createLogger("session-do", {}, parseLogLevel(env.LOG_LEVEL)),
getPublicSessionId
);
const backgroundTasks = createCloudflareBackgroundTasks(ctx, log);
const backgroundTasks = createBackgroundTasks(log);
// The sandbox repository validates the status it reads and warns on anything
// unmodelled, so it needs the session logger — and it owns encrypt-at-rest
// for access secrets, so it takes the key.
const sandboxRepository = new SandboxRepository(sql, log, repoSecretsEncryptionKey);

// Tier 2 — sockets and alarm scheduling.
const alarmScheduler = createEarliestAlarmScheduler(ctx.storage, alarmDeadlines);
const alarmScheduler = createEarliestAlarmScheduler(alarmStore, alarmDeadlines);
const wsManager: SessionWebSocketManager = new SessionWebSocketManagerImpl(
ctx,
socketHost,
sandboxRepository,
wsClientMappingRepository,
alarmScheduler,
log,
{ authTimeoutMs: WS_AUTH_TIMEOUT_MS }
);
// Hibernation-level ping/pong: the runtime answers keepalives without
// waking the Durable Object. Platform-global wiring, so it lives here.
ctx.setWebSocketAutoResponse(
new WebSocketRequestResponsePair(
JSON.stringify({ type: "ping" }),
JSON.stringify({ type: "pong", timestamp: Date.now() })
)
// Platform-level ping/pong: keepalives are answered without waking the
// runtime. Session-wide wiring, so it lives here.
socketHost.setAutoResponse(
JSON.stringify({ type: "ping" }),
JSON.stringify({ type: "pong", timestamp: Date.now() })
);

// Tier 3 — outbound delivery over the socket registry.
Expand All @@ -288,8 +286,8 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi

// Shared single instances/closures — every consumer below takes these
// rather than re-deriving its own copy.
const sessionIndexStore = db ? new SessionIndexStore(db) : null;
const sessionPullRequestStore = db ? new SessionPullRequestStore(db) : null;
const sessionIndexStore = new SessionIndexStore(db);
const sessionPullRequestStore = new SessionPullRequestStore(db);
const resolveRepoId = (sessionRow: SessionRow) =>
resolveSessionRepoId(sessionRow, sessionCoreRepository, sourceControlProvider);

Expand Down Expand Up @@ -332,7 +330,7 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
terminalMessageCompletedAt: completedAt,
});

const userScmTokenStore = db ? new UserScmTokenStore(db, tokenEncryptionKey) : null;
const userScmTokenStore = new UserScmTokenStore(db, tokenEncryptionKey);
const participantService = new ParticipantService({
repository: participantRepository,
getProcessingMessageAuthor: () => messageRepository.getProcessingMessageAuthor(),
Expand All @@ -342,14 +340,12 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
userScmTokenStore,
});

const scheduler = db ? new Scheduler(db, env, backgroundTasks) : undefined;
const scheduler = new Scheduler(db, env, backgroundTasks);
const callbackService = new CallbackNotificationService({
repository: sessionCoreRepository,
messageRepository,
env,
completeAutomationRun: scheduler
? (completion) => scheduler.runComplete(completion)
: undefined,
completeAutomationRun: (completion) => scheduler.runComplete(completion),
log,
getSessionId: () => resolvePublicSessionId(sessionCoreRepository.getSession(), durableObjectId),
});
Expand Down Expand Up @@ -558,7 +554,7 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
// service around the request-scoped log, so these stay functions.
const refreshOpenAIToken = async (sessionRow: SessionRow, requestLog: Logger) => {
const service = new OpenAITokenRefreshService(
db!,
db,
repoSecretsEncryptionKey,
resolveRepoId,
requestLog
Expand All @@ -567,7 +563,7 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
};
const refreshXaiToken = async (sessionRow: SessionRow, requestLog: Logger) => {
const service = new XaiTokenRefreshService(
db!,
db,
repoSecretsEncryptionKey,
resolveRepoId,
requestLog
Expand All @@ -585,7 +581,6 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
sandboxRepository,
sandboxEventProcessor,
messenger,
Boolean(db),
refreshOpenAIToken,
refreshXaiToken,
getScmCredentials,
Expand Down Expand Up @@ -646,7 +641,7 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
pushBranchToRemote: (pushSpec) => pushService.pushBranchToRemote(pushSpec),
messenger,
appName: resolveAppName(env),
sessionPullRequests: sessionPullRequestStore ?? undefined,
sessionPullRequests: sessionPullRequestStore,
resolveScmSettings: (repo) => resolveScmSettings(db, repo),
});

Expand Down Expand Up @@ -694,7 +689,6 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi
schedulePullRequestRefresh,
scmProviderName,
resolveAuthorization: async (userId) => {
if (!db) return { kind: "unavailable" };
try {
const authorization = await new AuthorizationService(db).getEffectiveAuthorization(userId);
return authorization.suspendedAt === null
Expand Down Expand Up @@ -868,7 +862,7 @@ export function createSessionRuntime(platform: SessionPlatform, env: Env): Sessi

interface LifecycleManagerDeps {
env: Env;
db: SqlDatabase | null;
db: SqlDatabase;
/** The latched public-session-id resolver shared with the session logger. */
getSessionId: () => string;
/** The repository, satisfying the manager's storage port structurally. */
Expand Down Expand Up @@ -913,34 +907,27 @@ function createLifecycleManager(deps: LifecycleManagerDeps): SandboxLifecycleMan
env.WORKER_URL ||
`https://open-inspect-control-plane.${env.CF_ACCOUNT_ID || "workers"}.workers.dev`;

// Create D1-backed lookups if database is available
let mcpServerLookup: McpServerLookup | undefined;
if (db) {
const mcpStore = new McpServerStore(db, repoSecretsEncryptionKey);
mcpServerLookup = {
getDecryptedForSession: (repositories) => mcpStore.getDecryptedForSession(repositories),
};
}
const mcpStore = new McpServerStore(db, repoSecretsEncryptionKey);
const mcpServerLookup: McpServerLookup = {
getDecryptedForSession: (repositories) => mcpStore.getDecryptedForSession(repositories),
};

// Session-scoped gate: resolved from the primary member (the scalar mirror
// this lookup is called with) — see resolveSessionScopedSettings for the
// per-feature scope rules. Token absence short-circuits to false so a
// misconfigured deployment never installs a tool that would 503 on every call.
let slackAgentNotifyLookup: SlackAgentNotifyLookup | undefined;
if (db) {
const tokenPresent = !!env.SLACK_BOT_TOKEN;
const settingsStore = new IntegrationSettingsStore(db);
slackAgentNotifyLookup = {
isEnabledForRepo: async (repoOwner, repoName) => {
if (!tokenPresent) return false;
const settings =
repoOwner && repoName
? (await settingsStore.getResolvedConfig("slack", `${repoOwner}/${repoName}`)).settings
: ((await settingsStore.getGlobal("slack"))?.defaults ?? {});
return resolveSlackSettings(settings).agentNotificationsEnabled;
},
};
}
const tokenPresent = !!env.SLACK_BOT_TOKEN;
const settingsStore = new IntegrationSettingsStore(db);
const slackAgentNotifyLookup: SlackAgentNotifyLookup = {
isEnabledForRepo: async (repoOwner, repoName) => {
if (!tokenPresent) return false;
const settings =
repoOwner && repoName
? (await settingsStore.getResolvedConfig("slack", `${repoOwner}/${repoName}`)).settings
: ((await settingsStore.getGlobal("slack"))?.defaults ?? {});
return resolveSlackSettings(settings).agentNotificationsEnabled;
},
};

const sandboxDashboardUrlBuilder =
sandboxBackend === "modal"
Expand All @@ -966,13 +953,11 @@ function createLifecycleManager(deps: LifecycleManagerDeps): SandboxLifecycleMan
sandboxDashboardUrlBuilder,
};

// Create the image lookup if D1 is available and the provider supports
// prebuilt images.
let imageBuildLookup: ImageBuildLookup | undefined;
// The image lookup exists only for providers that support prebuilt images.
const imageBuildProvider = resolveImageBuildProvider(sandboxBackend);
if (db && imageBuildProvider) {
imageBuildLookup = createImageBuildLookup(db, imageBuildProvider);
}
const imageBuildLookup: ImageBuildLookup | undefined = imageBuildProvider
? createImageBuildLookup(db, imageBuildProvider)
: undefined;

return new SandboxLifecycleManager(
provider,
Expand Down
Loading
Loading