fix: refresh oauth tokens

This commit is contained in:
johnyeo
2026-06-11 20:06:21 +01:00
parent 76096a57c9
commit 67de792bf8
14 changed files with 494 additions and 24 deletions

View File

@@ -14,6 +14,7 @@ import {
type UserMessageContentBlock,
} from "../../../harness/claudeManaged/session/userMessage.js";
import { ensureAutumnVault } from "../../../harness/claudeManaged/vaults/ensureAutumnVault.js";
import { cancelPendingSessionApprovals } from "../../../internal/approvals/actions/cancelPendingSessionApprovals.js";
import { containsSecret } from "../../../internal/sandbox/tool/guardrails.js";
import { db } from "../../../lib/db.js";
import { createBraintrustLogger } from "../../../providers/braintrust/index.js";
@@ -51,6 +52,27 @@ const redactSecrets = ({
const buildNudgeText = ({ toolName }: { toolName: string }) =>
`Call the ${toolName} tool now with the exact args from your preview. It will pause for user approval automatically — do not ask for confirmation or repeat the summary.`;
const isWaitingOnSessionResponseError = (error: unknown) =>
error instanceof Error &&
error.message.includes("waiting on responses to events");
const interruptSession = async ({
client,
sessionId,
}: {
client: Anthropic;
sessionId: string;
}) =>
await driveSessionTurn({
autumnMcpServerName: claudeManagedConfig.autumnMcpServerName,
client,
kickoff: () =>
client.beta.sessions.events.send(sessionId, {
events: [{ type: "user.interrupt" }],
}),
sessionId,
});
const mergeTurnOutcomes = (
first: SessionTurnOutcome,
second: SessionTurnOutcome,
@@ -99,7 +121,7 @@ const buildMessageText = ({
export const claudeManagedEngine: AgentEngine = {
name: "claude-managed",
run: async ({ ctx, params }) => {
const { env, logger, onAction, org, thread, token } = ctx;
const { env, logger, onAction, org, providerUserId, thread, token } = ctx;
const threadKey = [
thread.provider,
@@ -169,6 +191,29 @@ export const claudeManagedEngine: AgentEngine = {
}
const activeSessionId = sessionId;
if (!newSession) {
const { cancelledCount } = await cancelPendingSessionApprovals({
client,
db,
logger,
providerUserId,
query: {
channelId: thread.channelId,
env,
orgId: org.id,
provider: thread.provider,
runId: activeSessionId,
workspaceId: thread.workspaceId,
},
sessionId: activeSessionId,
});
if (cancelledCount > 0) {
await onAction?.(
"Cancelled the previous approval because new instructions were received.",
);
}
}
logger.info("Starting Claude Managed agent", {
event: "leaf.agent_started",
context: { env, org_id: org.id, provider: thread.provider },
@@ -232,8 +277,28 @@ export const claudeManagedEngine: AgentEngine = {
});
};
const driveTurnWithInterruptRetry = async ({
content,
span,
}: {
content: UserMessageContentBlock[];
span?: Span;
}) => {
try {
return await driveTurn({ content, span });
} catch (error) {
if (!isWaitingOnSessionResponseError(error)) throw error;
logger.warn("Interrupting blocked Claude Managed session", {
event: "leaf.claude_managed_session_interrupted",
context: { env, org_id: org.id },
});
await interruptSession({ client, sessionId: activeSessionId });
return await driveTurn({ content, span });
}
};
const runTurn = async ({ span }: { span?: Span }) => {
const first = await driveTurn({
const first = await driveTurnWithInterruptRetry({
content: buildUserMessageContent({
attachments: params.attachments,
text,

View File

@@ -26,6 +26,7 @@ export const runMessage = async ({
installation,
logger = rootLogger,
onAction,
providerUserId,
recentMessages,
text,
channelId,
@@ -87,6 +88,7 @@ export const runMessage = async ({
id: installation.org_id,
slug: installation.org_slug ?? undefined,
},
providerUserId,
thread: {
channelId,
provider: installation.provider,

View File

@@ -1,5 +1,5 @@
import type { AutumnLogger } from "@autumn/logging";
import type { AppEnv } from "@autumn/shared";
import type { AppEnv, ChatProvider } from "@autumn/shared";
import type { AgentHarnessName } from "../../lib/chatAgentConfig.js";
import type { AgentOutput, ChatContextMessage } from "../../types.js";
@@ -17,7 +17,7 @@ export type MessageAttachment = {
export type ThreadRef = {
channelId: string;
provider: string;
provider: ChatProvider;
threadId: string;
workspaceId: string;
};
@@ -31,6 +31,7 @@ export type MessageContext = {
logger: AutumnLogger;
onAction?: (message: string) => Promise<void> | void;
org: { id: string; slug?: string };
providerUserId: string;
thread: ThreadRef;
timestamp: number;
/** Org+env OAuth access token used for Autumn MCP auth. */

View File

@@ -154,6 +154,7 @@ const runAndReply = async ({
installation,
logger,
onAction: logAction,
providerUserId,
recentMessages,
text,
channelId,

View File

@@ -98,6 +98,21 @@ export const cmaRepo = {
return row?.vault_id;
},
getVault: async ({
db,
env,
orgId,
}: {
db: ChatDb;
env: AppEnv;
orgId: string;
}) => {
const row = await db.query.cmaVaults.findFirst({
where: and(eq(cmaVaults.org_id, orgId), eq(cmaVaults.env, env)),
});
return row;
},
upsertVault: async ({
credentialId,
db,

View File

@@ -30,12 +30,38 @@ const tokenEndpoint = () => {
return endpoint;
};
export const isCmaVaultStale = ({
credentialUpdatedAt,
vaultUpdatedAt,
}: {
credentialUpdatedAt: number;
vaultUpdatedAt?: number | null;
}) => !vaultUpdatedAt || credentialUpdatedAt > vaultUpdatedAt;
const buildCredentialAuth = ({
accessToken,
credential,
mcpServerUrl,
}: {
accessToken: string;
credential: NonNullable<
Awaited<ReturnType<typeof getChatOAuthCredentialByInstallationEnv>>
>;
mcpServerUrl: string;
}) => ({
type: "mcp_oauth" as const,
mcp_server_url: mcpServerUrl,
access_token: accessToken,
refresh: {
client_id: credential.oauth_client_id,
refresh_token: decrypt(credential.refresh_token),
token_endpoint: tokenEndpoint(),
token_endpoint_auth: { type: "none" as const },
},
});
// Mirrors the org's Autumn MCP OAuth credential into a CMA vault so Anthropic
// injects it after egress — the credential never enters the sandbox. Created
// once per (org, env); Anthropic auto-refreshes via the stored refresh token.
// NOTE: if Better Auth rotates the refresh token out-of-band, the vault copy goes
// stale and the next session reports an MCP auth error — re-seed by clearing the
// stored vault id. (Acceptable for v1; revisit if it bites.)
// injects it after egress. Resync when our local OAuth credential rotates.
export const ensureAutumnVault = async ({
client,
env,
@@ -49,9 +75,6 @@ export const ensureAutumnVault = async ({
provider: string;
workspaceId: string;
}): Promise<string> => {
const existing = await cmaRepo.getVaultId({ db, env, orgId });
if (existing) return existing;
const installation = await db.query.chatInstallations.findFirst({
where: and(
eq(chatInstallations.org_id, orgId),
@@ -76,22 +99,55 @@ export const ensureAutumnVault = async ({
}
const mcpServerUrl = new URL("/mcp", chatEnv.MCP_SERVER_URL).toString();
const existing = await cmaRepo.getVault({ db, env, orgId });
if (
existing &&
!isCmaVaultStale({
credentialUpdatedAt: credential.updated_at,
vaultUpdatedAt: existing.updated_at,
})
) {
return existing.vault_id;
}
const auth = buildCredentialAuth({ accessToken, credential, mcpServerUrl });
if (existing) {
const updated = await client.beta.vaults.credentials.update(
existing.credential_id,
{
vault_id: existing.vault_id,
auth: {
type: "mcp_oauth",
access_token: auth.access_token,
refresh: {
refresh_token: auth.refresh.refresh_token,
scope: credential.scopes.join(" "),
},
},
metadata: {
credential_updated_at: String(credential.updated_at),
},
},
);
await cmaRepo.upsertVault({
credentialId: updated.id,
db,
env,
orgId,
vaultId: existing.vault_id,
});
return existing.vault_id;
}
const vault = await client.beta.vaults.create({
display_name: `autumn/${orgId}/${env}`,
metadata: { app: "leaf", env, orgId },
});
const created = await client.beta.vaults.credentials.create(vault.id, {
display_name: `autumn-mcp/${orgId}/${env}`,
auth: {
type: "mcp_oauth",
mcp_server_url: mcpServerUrl,
access_token: accessToken,
refresh: {
client_id: credential.oauth_client_id,
refresh_token: decrypt(credential.refresh_token),
token_endpoint: tokenEndpoint(),
token_endpoint_auth: { type: "none" },
},
auth,
metadata: {
credential_updated_at: String(credential.updated_at),
},
});
await cmaRepo.upsertVault({

View File

@@ -0,0 +1,110 @@
import type Anthropic from "@anthropic-ai/sdk";
import type { AutumnLogger } from "@autumn/logging";
import type { ChatApproval } from "@autumn/shared";
import { claudeManagedConfig } from "../../../harness/claudeManaged/config.js";
import {
driveSessionTurn,
type SessionTurnOutcome,
} from "../../../harness/claudeManaged/session/driveSessionTurn.js";
import type { ChatDb } from "../../../lib/db.js";
import { chatApprovalRepo } from "../repos/chatApprovalRepo.js";
const STALE_APPROVAL_DENY_MESSAGE =
"User sent new instructions before approving this action.";
type ListPendingApprovalsInput = Parameters<
typeof chatApprovalRepo.listPendingForRun
>[0];
type CancelApprovalInput = Parameters<typeof chatApprovalRepo.cancel>[0];
type CancelPendingSessionApprovalsDeps = {
cancelApproval: (input: CancelApprovalInput) => Promise<ChatApproval | undefined>;
driveTurn: typeof driveSessionTurn;
listPendingApprovals: (
input: ListPendingApprovalsInput,
) => Promise<ChatApproval[]>;
};
const defaultDeps: CancelPendingSessionApprovalsDeps = {
cancelApproval: chatApprovalRepo.cancel,
driveTurn: driveSessionTurn,
listPendingApprovals: chatApprovalRepo.listPendingForRun,
};
export const cancelPendingSessionApprovalsWithDeps = async ({
client,
db,
logger,
providerUserId,
query,
sessionId,
deps = defaultDeps,
}: {
client: Anthropic;
db: ChatDb;
logger: AutumnLogger;
providerUserId: string;
query: Omit<ListPendingApprovalsInput, "db">;
sessionId: string;
deps?: CancelPendingSessionApprovalsDeps;
}) => {
const approvals = await deps.listPendingApprovals({ ...query, db });
if (approvals.length === 0) return { cancelledCount: 0 };
const executableApprovals = approvals.filter(
(approval): approval is ChatApproval & { tool_call_id: string } =>
Boolean(approval.tool_call_id),
);
let outcome: SessionTurnOutcome | undefined;
if (executableApprovals.length > 0) {
try {
outcome = await deps.driveTurn({
autumnMcpServerName: claudeManagedConfig.autumnMcpServerName,
client,
kickoff: () =>
client.beta.sessions.events.send(sessionId, {
events: executableApprovals.map((approval) => ({
deny_message: STALE_APPROVAL_DENY_MESSAGE,
result: "deny" as const,
tool_use_id: approval.tool_call_id,
type: "user.tool_confirmation" as const,
})),
}),
sessionId,
});
} catch (error) {
logger.warn("Failed to deny stale Claude Managed approval", {
event: "leaf.approval_auto_cancel_deny_failed",
error,
});
}
}
for (const approval of approvals) {
await deps.cancelApproval({
approvalId: approval.id,
db,
providerUserId,
});
}
logger.info("Cancelled stale pending approvals before new user message", {
event: "leaf.approval_auto_cancelled",
data: {
cancelled_count: approvals.length,
denied_count: executableApprovals.length,
had_session_error: Boolean(outcome?.errorMessage),
},
});
return { cancelledCount: approvals.length };
};
export const cancelPendingSessionApprovals = async (
input: Omit<
Parameters<typeof cancelPendingSessionApprovalsWithDeps>[0],
"deps"
>,
) => cancelPendingSessionApprovalsWithDeps(input);

View File

@@ -3,6 +3,7 @@ import { claimChatApproval } from "./claimChatApproval.js";
import { finalizeChatApproval } from "./finalizeChatApproval.js";
import { getChatApproval } from "./getChatApproval.js";
import { insertChatApproval } from "./insertChatApproval.js";
import { listPendingChatApprovalsForRun } from "./listPendingChatApprovalsForRun.js";
export const chatApprovalRepo = {
cancel: cancelChatApproval,
@@ -10,4 +11,5 @@ export const chatApprovalRepo = {
finalize: finalizeChatApproval,
get: getChatApproval,
insert: insertChatApproval,
listPendingForRun: listPendingChatApprovalsForRun,
} as const;

View File

@@ -0,0 +1,34 @@
import { type AppEnv, type ChatProvider, chatApprovals } from "@autumn/shared";
import { and, desc, eq, gt } from "drizzle-orm";
import type { ChatDb } from "../../../lib/db.js";
export const listPendingChatApprovalsForRun = async ({
channelId,
db,
env,
orgId,
provider,
runId,
workspaceId,
}: {
channelId: string;
db: ChatDb;
env: AppEnv;
orgId: string;
provider: ChatProvider;
runId: string;
workspaceId: string;
}) =>
await db.query.chatApprovals.findMany({
orderBy: desc(chatApprovals.created_at),
where: and(
eq(chatApprovals.org_id, orgId),
eq(chatApprovals.provider, provider),
eq(chatApprovals.workspace_id, workspaceId),
eq(chatApprovals.channel_id, channelId),
eq(chatApprovals.env, env),
eq(chatApprovals.run_id, runId),
eq(chatApprovals.status, "pending"),
gt(chatApprovals.expires_at, Date.now()),
),
});

View File

@@ -75,6 +75,7 @@ export type BotMessage = {
installation: LeafChatInstallation;
logger?: AutumnLogger;
onAction?: (message: string) => Promise<void> | void;
providerUserId: string;
recentMessages?: ChatContextMessage[];
text: string;
channelId: string;

View File

@@ -57,6 +57,7 @@ const main = async () => {
data: { message },
}),
org: { id: orgId },
providerUserId: "smoke",
thread: {
channelId: "smoke",
provider,

View File

@@ -19,6 +19,9 @@ const { autumnChatInstructions } = await import(
const { createFirecrawlTools } = await import(
"../../../src/agent/tools/firecrawl.js"
);
const { isCmaVaultStale } = await import(
"../../../src/harness/claudeManaged/vaults/ensureAutumnVault.js"
);
const execute = async (
tool: { execute?: (...args: never[]) => Promise<unknown> } | undefined,
@@ -175,3 +178,26 @@ describe("Firecrawl tools", () => {
expect((result as { markdown: string }).markdown.length).toBe(12_000);
});
});
describe("Claude Managed vault sync", () => {
test("treats the vault as stale when local OAuth credentials are newer", () => {
expect(
isCmaVaultStale({
credentialUpdatedAt: 2000,
vaultUpdatedAt: 1000,
}),
).toBe(true);
expect(
isCmaVaultStale({
credentialUpdatedAt: 1000,
vaultUpdatedAt: 2000,
}),
).toBe(false);
expect(
isCmaVaultStale({
credentialUpdatedAt: 1000,
vaultUpdatedAt: null,
}),
).toBe(true);
});
});

View File

@@ -1,4 +1,6 @@
import { describe, expect, test } from "bun:test";
import type Anthropic from "@anthropic-ai/sdk";
import type { AutumnLogger } from "@autumn/logging";
import { AppEnv, type ChatApproval } from "@autumn/shared";
import type { ActionEvent } from "chat";
import { approvalErrorResult } from "../../../src/internal/approvals/utils/approvalErrors.js";
@@ -14,6 +16,15 @@ const setLeafTestEnv = () => {
process.env.SLACK_SIGNING_SECRET ??= "test-slack-signing-secret";
};
const testLogger = {
child: () => testLogger,
debug: () => {},
error: () => {},
info: () => {},
warn: () => {},
warning: () => {},
} as unknown as AutumnLogger;
describe("approval flow", () => {
test("maps suspended destructive tool output to a pending approval request", () => {
const request = approvalRequestFromOutput({
@@ -174,4 +185,120 @@ describe("approval flow", () => {
expect(JSON.stringify(edits[1])).toContain("Attach plan failed");
expect(JSON.stringify(edits[1])).toContain("Missing email.");
});
test("denies and cancels stale pending approvals before a new user message", async () => {
setLeafTestEnv();
const { cancelPendingSessionApprovalsWithDeps } = await import(
"../../../src/internal/approvals/actions/cancelPendingSessionApprovals.js"
);
const approval = {
id: "approval_1",
tool_call_id: "tool_use_1",
} as ChatApproval;
const sentEvents: unknown[] = [];
const cancelled: unknown[] = [];
const client = {
beta: {
sessions: {
events: {
send: async (_sessionId: string, body: { events: unknown[] }) => {
sentEvents.push(...body.events);
},
},
},
},
} as unknown as Anthropic;
const result = await cancelPendingSessionApprovalsWithDeps({
client,
db: {} as never,
logger: testLogger,
providerUserId: "U1",
query: {
channelId: "C1",
env: AppEnv.Sandbox,
orgId: "org_1",
provider: "slack",
runId: "sesn_1",
workspaceId: "T1",
},
sessionId: "sesn_1",
deps: {
cancelApproval: async (input) => {
cancelled.push(input);
return approval;
},
driveTurn: async ({ kickoff }) => {
await kickoff();
return {
textParts: [],
usage: {
cacheCreationInputTokens: 0,
cacheReadInputTokens: 0,
inputTokens: 0,
outputTokens: 0,
},
};
},
listPendingApprovals: async () => [approval],
},
});
expect(result.cancelledCount).toBe(1);
expect(sentEvents).toEqual([
{
deny_message: "User sent new instructions before approving this action.",
result: "deny",
tool_use_id: "tool_use_1",
type: "user.tool_confirmation",
},
]);
expect(cancelled).toEqual([
{
approvalId: "approval_1",
db: {},
providerUserId: "U1",
},
]);
});
test("cancels pending approvals without a tool call id", async () => {
setLeafTestEnv();
const { cancelPendingSessionApprovalsWithDeps } = await import(
"../../../src/internal/approvals/actions/cancelPendingSessionApprovals.js"
);
const approval = {
id: "approval_1",
tool_call_id: null,
} as ChatApproval;
const cancelled: unknown[] = [];
await cancelPendingSessionApprovalsWithDeps({
client: {} as Anthropic,
db: {} as never,
logger: testLogger,
providerUserId: "U1",
query: {
channelId: "C1",
env: AppEnv.Sandbox,
orgId: "org_1",
provider: "slack",
runId: "sesn_1",
workspaceId: "T1",
},
sessionId: "sesn_1",
deps: {
cancelApproval: async (input) => {
cancelled.push(input);
return approval;
},
driveTurn: async () => {
throw new Error("should not drive session");
},
listPendingApprovals: async () => [approval],
},
});
expect(cancelled).toHaveLength(1);
});
});

View File

@@ -1,10 +1,12 @@
import {
AppEnv,
FeatureUsageType,
type ProductItem,
type ProductV2,
} from "@autumn/shared";
import { items } from "@tests/utils/fixtures/items";
import { products } from "@tests/utils/fixtures/products";
import { clearOrg } from "@tests/utils/setup/clearOrg.js";
import defaultCtx, {
createTestContext,
type TestContext,
@@ -329,6 +331,18 @@ export const buildKnowledgePlatformProducts = () => {
};
};
const withProductGroup = ({
products,
group,
}: {
products: ProductV2[];
group: string;
}) =>
products.map((product) => ({
...product,
group,
}));
export const initKnowledgePlatformScenario = async ({
customerId = "agent-knowledge-platform",
attachPlan = "trial",
@@ -398,8 +412,11 @@ export const seedKnowledgePlatformCustomers = async ({
await initScenario({
setup: [
s.products({
list: Object.values(catalog.plans),
prefix: productPrefix,
list: withProductGroup({
products: Object.values(catalog.plans),
group: productPrefix,
}),
prefix: "",
createInStripe: false,
}),
],
@@ -448,6 +465,18 @@ const runKnowledgePlatformSeed = async () => {
const attachPlan = getArgValue("--attach-plan") as
| Extract<KnowledgePlatformPlanKey, "trial" | "enterprise">
| undefined;
if (!process.env.TESTS_ORG) {
throw new Error("TESTS_ORG is required to seed knowledge platform data");
}
if (!process.argv.includes("--skip-clear")) {
await clearOrg({
orgSlug: process.env.TESTS_ORG,
env: AppEnv.Sandbox,
skipStripeReset: true,
});
}
const ctx = await createTestContext();
const result = await seedKnowledgePlatformCustomers({