From 8cebaa7533fd01334c01ce456f10654f4e17daa4 Mon Sep 17 00:00:00 2001 From: Saxon Fletcher Date: Fri, 21 Aug 2026 17:00:10 +1000 Subject: [PATCH] feat(studio): connect assistant to GitHub repositories --- .../GeneralSettings/AIOptInLevelSelector.tsx | 26 +++++- .../GeneralSettings/DataPrivacyForm.tsx | 6 +- .../ui/AIAssistantPanel/AIOptInModal.tsx | 7 +- .../ui/AIAssistantPanel/Message.Parts.tsx | 61 +++++++++++++- .../ui/AIAssistantPanel/Message.utils.ts | 13 +++ .../PullRequestRenderer.test.tsx | 47 +++++++++++ .../AIAssistantPanel/PullRequestRenderer.tsx | 64 +++++++++++++++ .../integrations/github-connections-query.ts | 4 +- .../data/permissions/permissions-query.ts | 4 +- apps/studio/hooks/forms/useAIOptInForm.ts | 11 ++- .../hooks/misc/useOrgOptedIntoAi.test.ts | 10 +++ apps/studio/hooks/misc/useOrgOptedIntoAi.ts | 6 ++ apps/studio/lib/ai/ai-details.test.ts | 30 +++++++ apps/studio/lib/ai/ai-details.ts | 39 ++++++--- .../lib/ai/generate-assistant-response.ts | 2 + apps/studio/lib/ai/prompts.ts | 10 ++- apps/studio/lib/ai/repo-ref.test.ts | 27 +++++++ apps/studio/lib/ai/repo-ref.ts | 11 +++ apps/studio/lib/ai/tools/index.ts | 9 ++- apps/studio/lib/ai/tools/repo-tools.test.ts | 33 +++++--- apps/studio/lib/ai/tools/repo-tools.ts | 80 ++++++++++++------- apps/studio/lib/constants/index.ts | 1 + apps/studio/pages/api/ai/sql/generate-v4.ts | 61 ++++++++++++++ apps/studio/turbo.jsonc | 12 +-- 24 files changed, 502 insertions(+), 72 deletions(-) create mode 100644 apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.test.tsx create mode 100644 apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.tsx create mode 100644 apps/studio/hooks/misc/useOrgOptedIntoAi.test.ts create mode 100644 apps/studio/lib/ai/repo-ref.test.ts create mode 100644 apps/studio/lib/ai/repo-ref.ts diff --git a/apps/studio/components/interfaces/Organization/GeneralSettings/AIOptInLevelSelector.tsx b/apps/studio/components/interfaces/Organization/GeneralSettings/AIOptInLevelSelector.tsx index 47207ec0a1f..2c8b388c436 100644 --- a/apps/studio/components/interfaces/Organization/GeneralSettings/AIOptInLevelSelector.tsx +++ b/apps/studio/components/interfaces/Organization/GeneralSettings/AIOptInLevelSelector.tsx @@ -1,6 +1,6 @@ import { ReactNode } from 'react' -import { Control } from 'react-hook-form' -import { FormField, RadioGroup, RadioGroupItem } from 'ui' +import { Control, useWatch } from 'react-hook-form' +import { FormField, RadioGroup, RadioGroupItem, Switch } from 'ui' import { FormItemLayout } from 'ui-patterns/form/FormItemLayout/FormItemLayout' import { OptInToOpenAIToggle } from './OptInToOpenAIToggle' @@ -20,6 +20,7 @@ export const AIOptInLevelSelector = ({ label, layout = 'vertical', }: AIOptInLevelSelectorProps) => { + const aiOptInLevel = useWatch({ control, name: 'aiOptInLevel' }) const { aiOptInLevelDisabled, aiOptInLevelSchema, @@ -125,6 +126,27 @@ export const AIOptInLevelSelector = ({ )} /> + ( +
+
+

Repository access

+

+ Allow the Assistant to read the connected GitHub repository and propose changes + through pull requests. +

+
+ +
+ )} + /> ) diff --git a/apps/studio/components/interfaces/Organization/GeneralSettings/DataPrivacyForm.tsx b/apps/studio/components/interfaces/Organization/GeneralSettings/DataPrivacyForm.tsx index 38d191a19b3..1c0e9e24d45 100644 --- a/apps/studio/components/interfaces/Organization/GeneralSettings/DataPrivacyForm.tsx +++ b/apps/studio/components/interfaces/Organization/GeneralSettings/DataPrivacyForm.tsx @@ -8,7 +8,7 @@ import { useAIOptInForm } from '@/hooks/forms/useAIOptInForm' import { useAsyncCheckPermissions } from '@/hooks/misc/useCheckPermissions' export const DataPrivacyForm = () => { - const { form, onSubmit, isUpdating, currentOptInLevel } = useAIOptInForm() + const { form, onSubmit, isUpdating, currentOptInLevel, currentRepoAccess } = useAIOptInForm() const { can: canUpdateOrganization } = useAsyncCheckPermissions( PermissionAction.UPDATE, 'organizations' @@ -19,8 +19,8 @@ export const DataPrivacyForm = () => { : undefined useEffect(() => { - form.reset({ aiOptInLevel: currentOptInLevel }) - }, [currentOptInLevel, form]) + form.reset({ aiOptInLevel: currentOptInLevel, hasRepoAccess: currentRepoAccess }) + }, [currentOptInLevel, currentRepoAccess, form]) return (
diff --git a/apps/studio/components/ui/AIAssistantPanel/AIOptInModal.tsx b/apps/studio/components/ui/AIAssistantPanel/AIOptInModal.tsx index 489870732e8..4a945b769e5 100644 --- a/apps/studio/components/ui/AIAssistantPanel/AIOptInModal.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/AIOptInModal.tsx @@ -23,7 +23,8 @@ interface AIOptInModalProps { } export const AIOptInModal = ({ visible, onCancel }: AIOptInModalProps) => { - const { form, onSubmit, isUpdating, currentOptInLevel } = useAIOptInForm(onCancel) + const { form, onSubmit, isUpdating, currentOptInLevel, currentRepoAccess } = + useAIOptInForm(onCancel) const { can: canUpdateOrganization } = useAsyncCheckPermissions( PermissionAction.UPDATE, 'organizations' @@ -37,9 +38,9 @@ export const AIOptInModal = ({ visible, onCancel }: AIOptInModalProps) => { useEffect(() => { if (visible) { - form.reset({ aiOptInLevel: currentOptInLevel }) + form.reset({ aiOptInLevel: currentOptInLevel, hasRepoAccess: currentRepoAccess }) } - }, [visible, currentOptInLevel, form]) + }, [visible, currentOptInLevel, currentRepoAccess, form]) return ( diff --git a/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx b/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx index e6ea99e7623..019378a68d5 100644 --- a/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx +++ b/apps/studio/components/ui/AIAssistantPanel/Message.Parts.tsx @@ -11,10 +11,13 @@ import { useMessageActionsContext, useMessageInfoContext } from './Message.Conte import { deployEdgeFunctionInputSchema, deployEdgeFunctionOutputSchema, + openPullRequestInputSchema, + openPullRequestOutputSchema, parseExecuteSqlChartResult, } from './Message.utils' import { MessageMarkdown } from './MessageMarkdown' import { NotebookProposalRenderer, type NotebookProposalMode } from './NotebookProposalRenderer' +import { PullRequestRenderer } from './PullRequestRenderer' import { parseSupportRequestMessage, SupportRequestMessage } from './SupportRequestMessage' function MessagePartText({ textPart }: { textPart: TextUIPart }) { @@ -62,6 +65,13 @@ function MessagePartDynamicTool({ toolPart }: { toolPart: DynamicToolUIPart }) { } function MessagePartTool({ toolPart }: { toolPart: ToolUIPart }) { + const repoLabels: Record = { + 'tool-search_repo': ['Searching repository...', 'Searched repository'], + 'tool-read_repo_file': ['Reading repository file...', 'Read repository file'], + 'tool-write_repo_file': ['Updating repository...', 'Updated repository working copy'], + } + const repoLabel = repoLabels[toolPart.type] + return ( - {toolPart.state === 'input-streaming' ? 'Running ' : 'Ran '} - {`${toolPart.type.replace('tool-', '')}`} - + repoLabel ? ( + repoLabel[toolPart.state === 'input-streaming' ? 0 : 1] + ) : ( +
+ {toolPart.state === 'input-streaming' ? 'Running ' : 'Ran '} + {`${toolPart.type.replace('tool-', '')}`} +
+ ) } /> ) @@ -268,6 +282,36 @@ function MessagePartNotebookProposal({ ) } +function MessagePartOpenPullRequest({ toolPart }: { toolPart: ToolUIPart }) { + const { state, input, output } = toolPart + const { addToolApprovalResponse } = useMessageActionsContext() + + if (state === 'input-streaming') return + if (state === 'output-error') { + return

Failed to open pull request.

+ } + + const parsedInput = openPullRequestInputSchema.safeParse(input) + if (!parsedInput.success) return null + const parsedOutput = openPullRequestOutputSchema.safeParse(output) + const { confirmState, onApprove, onDeny } = getManualToolApprovalHandlers({ + state, + approval: toolPart.approval, + addToolApprovalResponse, + }) + + return ( + + ) +} + const MessagePart = { Text: MessagePartText, Dynamic: MessagePartDynamicTool, @@ -276,6 +320,7 @@ const MessagePart = { ExecuteSql: MessagePartExecuteSql, DeployEdgeFunction: MessagePartDeployEdgeFunction, NotebookProposal: MessagePartNotebookProposal, + OpenPullRequest: MessagePartOpenPullRequest, } as const export function MessagePartSwitcher({ @@ -293,6 +338,11 @@ export function MessagePartSwitcher({ case 'tool-load_knowledge': { return } + case 'tool-search_repo': + case 'tool-read_repo_file': + case 'tool-write_repo_file': { + return + } case 'reasoning': return case 'text': @@ -310,6 +360,9 @@ export function MessagePartSwitcher({ case 'tool-update_notebook': { return } + case 'tool-open_pull_request': { + return + } case 'source-url': case 'source-document': diff --git a/apps/studio/components/ui/AIAssistantPanel/Message.utils.ts b/apps/studio/components/ui/AIAssistantPanel/Message.utils.ts index 9867cc27511..5f6a2f01094 100644 --- a/apps/studio/components/ui/AIAssistantPanel/Message.utils.ts +++ b/apps/studio/components/ui/AIAssistantPanel/Message.utils.ts @@ -143,6 +143,19 @@ export const updateNotebookInputSchema = z.object({ export const notebookToolOutputSchema = z.object({ id: z.string(), name: z.string() }) +export const openPullRequestInputSchema = z.object({ + title: z.string(), + body: z.string().optional(), + patch: z.string(), +}) + +export const openPullRequestOutputSchema = z.object({ + url: z.string().url(), + number: z.number(), + branch: z.string(), + sha: z.string(), +}) + export const rateMessageResponseSchema = z.object({ category: z.enum([ 'sql_generation', diff --git a/apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.test.tsx b/apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.test.tsx new file mode 100644 index 00000000000..678c928610e --- /dev/null +++ b/apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.test.tsx @@ -0,0 +1,47 @@ +import { screen } from '@testing-library/react' +import userEvent from '@testing-library/user-event' +import { describe, expect, it, vi } from 'vitest' + +import { PullRequestRenderer } from './PullRequestRenderer' +import { render } from '@/tests/helpers' + +describe('PullRequestRenderer', () => { + it('shows the exact patch and supports approval or denial', async () => { + const user = userEvent.setup() + const onApprove = vi.fn() + const onDeny = vi.fn() + + render( + + ) + + expect(screen.getByText('diff --git a/auth.ts b/auth.ts')).toBeVisible() + await user.click(screen.getByRole('button', { name: 'Open PR' })) + await user.click(screen.getByRole('button', { name: 'Skip' })) + expect(onApprove).toHaveBeenCalledOnce() + expect(onDeny).toHaveBeenCalledOnce() + }) + + it('links to the created pull request', () => { + render( + + ) + + expect(screen.getByRole('link', { name: /View pull request #12/ })).toHaveAttribute( + 'href', + 'https://github.com/acme/repo/pull/12' + ) + expect(screen.queryByText('diff')).not.toBeInTheDocument() + }) +}) diff --git a/apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.tsx b/apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.tsx new file mode 100644 index 00000000000..98b93dd1a9b --- /dev/null +++ b/apps/studio/components/ui/AIAssistantPanel/PullRequestRenderer.tsx @@ -0,0 +1,64 @@ +import { ExternalLink, GitPullRequest } from 'lucide-react' + +import { Confirm } from './Confirm' +import type { ConfirmFooterApprovalState } from './Confirm.utils' + +export function PullRequestRenderer({ + title, + body, + patch, + url, + number, + confirmState, + onApprove, + onDeny, +}: { + title: string + body?: string + patch: string + url?: string + number?: number + confirmState?: ConfirmFooterApprovalState + onApprove?: () => void + onDeny?: () => void +}) { + return ( + +
+
+ +
+

{title}

+ {body && ( +

{body}

+ )} +
+
+ {url && ( + + View pull request{number ? ` #${number}` : ''} + + )} + {!url && ( +
+            {patch}
+          
+ )} +
+
+ ) +} diff --git a/apps/studio/data/integrations/github-connections-query.ts b/apps/studio/data/integrations/github-connections-query.ts index cb22cf1916b..142ecced588 100644 --- a/apps/studio/data/integrations/github-connections-query.ts +++ b/apps/studio/data/integrations/github-connections-query.ts @@ -12,7 +12,8 @@ export type GitHubConnectionsVariables = { export async function getGitHubConnections( { organizationId }: GitHubConnectionsVariables, - signal?: AbortSignal + signal?: AbortSignal, + headers?: HeadersInit ) { if (!organizationId) throw new Error('organizationId is required') @@ -23,6 +24,7 @@ export async function getGitHubConnections( }, }, signal, + headers, }) if (error) handleError(error) diff --git a/apps/studio/data/permissions/permissions-query.ts b/apps/studio/data/permissions/permissions-query.ts index 29805386523..cd8a55009bf 100644 --- a/apps/studio/data/permissions/permissions-query.ts +++ b/apps/studio/data/permissions/permissions-query.ts @@ -8,8 +8,8 @@ import type { Permission, ResponseError, UseCustomQueryOptions } from '@/types' export type PermissionsResponse = Permission[] -export async function getPermissions(signal?: AbortSignal) { - const { data, error } = await get('/platform/profile/permissions', { signal }) +export async function getPermissions(signal?: AbortSignal, headers?: HeadersInit) { + const { data, error } = await get('/platform/profile/permissions', { signal, headers }) if (error) { handleError(error, { sentryContext: { diff --git a/apps/studio/hooks/forms/useAIOptInForm.ts b/apps/studio/hooks/forms/useAIOptInForm.ts index 57361a16f26..99f52b47b22 100644 --- a/apps/studio/hooks/forms/useAIOptInForm.ts +++ b/apps/studio/hooks/forms/useAIOptInForm.ts @@ -10,7 +10,7 @@ import { useOrganizationUpdateMutation } from '@/data/organizations/organization import { invalidateOrganizationsQuery } from '@/data/organizations/organizations-query' import { useAsyncCheckPermissions } from '@/hooks/misc/useCheckPermissions' import { useLocalStorageQuery } from '@/hooks/misc/useLocalStorage' -import { getAiOptInLevel } from '@/hooks/misc/useOrgOptedIntoAi' +import { getAiOptInLevel, getAiRepoAccess } from '@/hooks/misc/useOrgOptedIntoAi' import { useSelectedOrganizationQuery } from '@/hooks/misc/useSelectedOrganization' import { OPT_IN_TAGS } from '@/lib/constants' import type { ResponseError } from '@/types' @@ -20,6 +20,7 @@ export const AIOptInSchema = z.object({ aiOptInLevel: z.enum(['disabled', 'schema', 'schema_and_log', 'schema_and_log_and_data'], { required_error: 'AI Opt-in level selection is required', }), + hasRepoAccess: z.boolean(), }) export type AIOptInFormValues = z.infer @@ -47,6 +48,7 @@ export const useAIOptInForm = (onSuccessCallback?: () => void) => { resolver: zodResolver(AIOptInSchema), defaultValues: { aiOptInLevel: getAiOptInLevel(selectedOrganization?.opt_in_tags), + hasRepoAccess: getAiRepoAccess(selectedOrganization?.opt_in_tags), }, }) @@ -64,7 +66,8 @@ export const useAIOptInForm = (onSuccessCallback?: () => void) => { (tag: string) => tag !== OPT_IN_TAGS.AI_SQL && tag !== (OPT_IN_TAGS.AI_DATA ?? 'AI_DATA') && - tag !== (OPT_IN_TAGS.AI_LOG ?? 'AI_LOG') + tag !== (OPT_IN_TAGS.AI_LOG ?? 'AI_LOG') && + tag !== OPT_IN_TAGS.AI_REPO ) if ( @@ -83,6 +86,9 @@ export const useAIOptInForm = (onSuccessCallback?: () => void) => { if (values.aiOptInLevel === 'schema_and_log_and_data') { updatedOptInTags.push(OPT_IN_TAGS.AI_DATA) } + if (values.aiOptInLevel !== 'disabled' && values.hasRepoAccess) { + updatedOptInTags.push(OPT_IN_TAGS.AI_REPO) + } updatedOptInTags = [...new Set(updatedOptInTags)] @@ -107,5 +113,6 @@ export const useAIOptInForm = (onSuccessCallback?: () => void) => { onSubmit, isUpdating, currentOptInLevel: getAiOptInLevel(selectedOrganization?.opt_in_tags), + currentRepoAccess: getAiRepoAccess(selectedOrganization?.opt_in_tags), } } diff --git a/apps/studio/hooks/misc/useOrgOptedIntoAi.test.ts b/apps/studio/hooks/misc/useOrgOptedIntoAi.test.ts new file mode 100644 index 00000000000..da6563353d9 --- /dev/null +++ b/apps/studio/hooks/misc/useOrgOptedIntoAi.test.ts @@ -0,0 +1,10 @@ +import { describe, expect, it } from 'vitest' + +import { getAiRepoAccess } from './useOrgOptedIntoAi' + +describe('getAiRepoAccess', () => { + it('requires the independent repository opt-in tag', () => { + expect(getAiRepoAccess(['AI_SQL_GENERATOR_OPT_IN'])).toBe(false) + expect(getAiRepoAccess(['AI_SQL_GENERATOR_OPT_IN', 'AI_REPO_ACCESS_OPT_IN'])).toBe(true) + }) +}) diff --git a/apps/studio/hooks/misc/useOrgOptedIntoAi.ts b/apps/studio/hooks/misc/useOrgOptedIntoAi.ts index 00ba0a91862..0f28bbcec94 100644 --- a/apps/studio/hooks/misc/useOrgOptedIntoAi.ts +++ b/apps/studio/hooks/misc/useOrgOptedIntoAi.ts @@ -32,6 +32,9 @@ export const getAiOptInLevel = (tags: string[] | undefined): AiOptInLevel => { } } +export const getAiRepoAccess = (tags: string[] | undefined) => + tags?.includes(OPT_IN_TAGS.AI_REPO) ?? false + /** * Determines if the organization has opted into *any* level of AI features (schema or schema_and_log or schema_and_log_and_data). * This is primarily for backward compatibility. @@ -50,6 +53,7 @@ export function useOrgAiOptInLevel(): { aiOptInLevel: AiOptInLevel includeSchemaMetadata: boolean isHipaaProjectDisallowed: boolean + hasRepoAccess: boolean } { const { data: selectedProject } = useSelectedProjectQuery() const { data: selectedOrganization } = useSelectedOrganizationQuery() @@ -80,5 +84,7 @@ export function useOrgAiOptInLevel(): { aiOptInLevel, includeSchemaMetadata, isHipaaProjectDisallowed: preventProjectFromUsingAI, + hasRepoAccess: + !preventProjectFromUsingAI && aiOptInLevel !== 'disabled' && getAiRepoAccess(optInTags), } } diff --git a/apps/studio/lib/ai/ai-details.test.ts b/apps/studio/lib/ai/ai-details.test.ts index 56e4b61b3a2..7a70dce0b15 100644 --- a/apps/studio/lib/ai/ai-details.test.ts +++ b/apps/studio/lib/ai/ai-details.test.ts @@ -20,6 +20,7 @@ vi.mock('@/data/config/project-settings-v2-query', () => ({ vi.mock('@/hooks/misc/useOrgOptedIntoAi', () => ({ getAiOptInLevel: vi.fn(), + getAiRepoAccess: vi.fn(), })) vi.mock('@/components/interfaces/Billing/Subscription/Subscription.utils', () => ({ @@ -41,6 +42,7 @@ describe('getAIDetails', () => { let mockGetProjectDetail: ReturnType let mockGetProjectSettings: ReturnType let mockGetAiOptInLevel: ReturnType + let mockGetAiRepoAccess: ReturnType let mockSubscriptionHasHipaaAddon: ReturnType let mockCheckEntitlement: ReturnType @@ -61,6 +63,7 @@ describe('getAIDetails', () => { mockGetProjectDetail = vi.mocked(projectQuery.getProjectDetail) mockGetProjectSettings = vi.mocked(settingsQuery.getProjectSettings) mockGetAiOptInLevel = vi.mocked(aiHook.getAiOptInLevel) + mockGetAiRepoAccess = vi.mocked(aiHook.getAiRepoAccess) mockSubscriptionHasHipaaAddon = vi.mocked(subscriptionUtils.subscriptionHasHipaaAddon) mockCheckEntitlement = vi.mocked(entitlementsQuery.checkEntitlement) @@ -77,6 +80,7 @@ describe('getAIDetails', () => { mockSubscriptionHasHipaaAddon.mockReturnValue(false) mockCheckEntitlement.mockResolvedValue({ hasAccess: false }) mockGetAiOptInLevel.mockReturnValue('schema') + mockGetAiRepoAccess.mockReturnValue(false) }) it('returns the resolved posture when the project belongs to the org', async () => { @@ -94,6 +98,8 @@ describe('getAIDetails', () => { planId: 'pro', region: 'us-east-1', isSensitive: false, + parentProjectRef: undefined, + hasRepoAccess: false, }) }) @@ -152,6 +158,12 @@ describe('getAIDetails', () => { undefined, HEADERS ) + expect(mockCheckEntitlement).toHaveBeenCalledWith( + ORG_SLUG, + 'assistant.repo_access', + undefined, + HEADERS + ) expect(mockGetProjectDetail).toHaveBeenCalledWith( { ref: PROJECT_REF, skipWake: true }, undefined, @@ -164,6 +176,21 @@ describe('getAIDetails', () => { ) }) + it('requires the repository tag and entitlement', async () => { + mockGetAiRepoAccess.mockReturnValue(true) + mockCheckEntitlement.mockImplementation(async (_slug, feature) => ({ + hasAccess: feature === 'assistant.repo_access', + })) + + const result = await getAIDetails({ + orgSlug: ORG_SLUG, + projectRef: PROJECT_REF, + authorization: AUTH, + }) + + expect(result.hasRepoAccess).toBe(true) + }) + describe('when the project does not belong to the org', () => { beforeEach(() => { mockGetOrganizations.mockResolvedValue([ @@ -239,6 +266,8 @@ describe('getAIDetails', () => { it('disables the opt-in level for a sensitive project', async () => { mockGetProjectSettings.mockResolvedValue({ is_sensitive: true }) + mockGetAiRepoAccess.mockReturnValue(true) + mockCheckEntitlement.mockResolvedValue({ hasAccess: true }) const result = await getAIDetails({ orgSlug: ORG_SLUG, @@ -248,6 +277,7 @@ describe('getAIDetails', () => { expect(result.aiOptInLevel).toBe('disabled') expect(result.hasHipaaAddon).toBe(true) + expect(result.hasRepoAccess).toBe(false) }) it('disables the opt-in level when project sensitivity is unknown', async () => { diff --git a/apps/studio/lib/ai/ai-details.ts b/apps/studio/lib/ai/ai-details.ts index be61660368d..103b214d462 100644 --- a/apps/studio/lib/ai/ai-details.ts +++ b/apps/studio/lib/ai/ai-details.ts @@ -4,7 +4,7 @@ import { checkEntitlement } from '@/data/entitlements/entitlements-query' import { getOrganizations } from '@/data/organizations/organizations-query' import { getProjectDetail } from '@/data/projects/project-detail-query' import { getOrgSubscription } from '@/data/subscriptions/org-subscription-query' -import { getAiOptInLevel, type AiOptInLevel } from '@/hooks/misc/useOrgOptedIntoAi' +import { getAiOptInLevel, getAiRepoAccess, type AiOptInLevel } from '@/hooks/misc/useOrgOptedIntoAi' export type AIDetails = { aiOptInLevel: AiOptInLevel @@ -14,6 +14,8 @@ export type AIDetails = { planId: string | undefined region: string | undefined isSensitive: boolean | null | undefined + parentProjectRef: string | undefined + hasRepoAccess: boolean } // Resolves the AI opt-in level, model access and tracing inputs for one org/project pair. @@ -33,15 +35,22 @@ export const getAIDetails = async ({ ...(authorization && { Authorization: authorization }), } - const [organizations, subscription, advanceModelAccess, project, projectSettings] = - await Promise.all([ - getOrganizations({ headers }), - getOrgSubscription({ orgSlug }, undefined, headers), - checkEntitlement(orgSlug, 'assistant.advance_model', undefined, headers), - // skipWake: only organization_id and region are needed, neither requires a running project - getProjectDetail({ ref: projectRef, skipWake: true }, undefined, headers), - getProjectSettings({ projectRef }, undefined, headers), - ]) + const [ + organizations, + subscription, + advanceModelAccess, + repoAccessEntitlement, + project, + projectSettings, + ] = await Promise.all([ + getOrganizations({ headers }), + getOrgSubscription({ orgSlug }, undefined, headers), + checkEntitlement(orgSlug, 'assistant.advance_model', undefined, headers), + checkEntitlement(orgSlug, 'assistant.repo_access', undefined, headers), + // skipWake: only organization_id and region are needed, neither requires a running project + getProjectDetail({ ref: projectRef, skipWake: true }, undefined, headers), + getProjectSettings({ projectRef }, undefined, headers), + ]) const selectedOrg = organizations.find((org) => org.slug === orgSlug) const region = project?.region @@ -59,6 +68,8 @@ export const getAIDetails = async ({ planId: undefined, region, isSensitive, + parentProjectRef: project?.parent_project_ref, + hasRepoAccess: false, } } @@ -66,14 +77,20 @@ export const getAIDetails = async ({ // Mirrors the client-side gate in useOrgAiOptInLevel, which had no server-side equivalent const isRestrictedByHipaa = hasHipaaAddon && isSensitive !== false + const aiOptInLevel = isRestrictedByHipaa ? 'disabled' : getAiOptInLevel(selectedOrg.opt_in_tags) return { - aiOptInLevel: isRestrictedByHipaa ? 'disabled' : getAiOptInLevel(selectedOrg.opt_in_tags), + aiOptInLevel, hasAccessToAdvanceModel: advanceModelAccess.hasAccess, hasHipaaAddon, orgId: selectedOrg.id, planId: selectedOrg.plan.id, region, isSensitive, + parentProjectRef: project.parent_project_ref, + hasRepoAccess: + aiOptInLevel !== 'disabled' && + getAiRepoAccess(selectedOrg.opt_in_tags) && + repoAccessEntitlement.hasAccess, } } diff --git a/apps/studio/lib/ai/generate-assistant-response.ts b/apps/studio/lib/ai/generate-assistant-response.ts index 25608e7e8b2..f2c5eb76f50 100644 --- a/apps/studio/lib/ai/generate-assistant-response.ts +++ b/apps/studio/lib/ai/generate-assistant-response.ts @@ -22,6 +22,7 @@ import { GENERAL_PROMPT, LIMITATIONS_PROMPT, NOTEBOOKS_PROMPT, + REPO_PROMPT, SECURITY_PROMPT, } from '@/lib/ai/prompts' import { sanitizeMessagePart } from '@/lib/ai/tools/tool-sanitizer' @@ -122,6 +123,7 @@ export async function generateAssistantResponse({ ${GENERAL_PROMPT} ${CHAT_PROMPT} ${isExplorerEnabled ? NOTEBOOKS_PROMPT : ''} + ${sandbox ? REPO_PROMPT : ''} ${SECURITY_PROMPT} ${LIMITATIONS_PROMPT} diff --git a/apps/studio/lib/ai/prompts.ts b/apps/studio/lib/ai/prompts.ts index 8629718e8fc..f49079087d5 100644 --- a/apps/studio/lib/ai/prompts.ts +++ b/apps/studio/lib/ai/prompts.ts @@ -770,6 +770,14 @@ export const NOTEBOOKS_PROMPT = ` - When describing an existing notebook, report each query cell's configuration that changes what it returns — a log cell's time range, a database cell's row limit — and don't count markdown cells as queries. ` +export const REPO_PROMPT = ` +## Connected repository +- Use \`search_repo\` and \`read_repo_file\` to ground code-related answers in the connected repository. +- Use \`write_repo_file\` only when the user asks for a code change. +- After making the complete requested change, copy the exact \`patch\` from the final \`write_repo_file\` result into one \`open_pull_request\` call. The user reviews and approves that patch before anything is written to GitHub. +- Never claim a change is in GitHub until \`open_pull_request\` returns successfully. +` + export const OUTPUT_ONLY_PROMPT = ` # Output-Only Mode @@ -805,7 +813,7 @@ export const LIMITATIONS_PROMPT = ` - For questions about plan, billing or usage limitations, refer to the user to Supabase documentation - Always search_docs before providing any links to Supabase documentation or dashboard pages ## Destructive Operations -- Do not help with local filesystem or git operations (e.g. \`git reset --hard\`, \`git clean\`, \`rm -rf\`). These are outside your scope — politely decline and direct the user to git documentation or a developer peer. +- Do not use or suggest destructive filesystem or git operations (e.g. \`git reset --hard\`, \`git clean\`, \`rm -rf\`). Connected repository tools, when available, may read files, edit files, and open pull requests. - For irreversible database operations (DROP TABLE, TRUNCATE, DELETE without a WHERE clause, dropping columns or schemas), always lead with an explicit warning that the operation cannot be undone before proceeding. - When a user appears non-technical based on their language or questions, explain consequences of destructive actions in plain terms before suggesting anything irreversible. ` diff --git a/apps/studio/lib/ai/repo-ref.test.ts b/apps/studio/lib/ai/repo-ref.test.ts new file mode 100644 index 00000000000..2e58699eafd --- /dev/null +++ b/apps/studio/lib/ai/repo-ref.test.ts @@ -0,0 +1,27 @@ +import { describe, expect, it } from 'vitest' + +import { resolveRepoRef } from './repo-ref' + +describe('resolveRepoRef', () => { + it('uses the Git branch attached to the current preview', () => { + expect( + resolveRepoRef({ + currentBranch: { git_branch: 'feature/current' }, + branches: [{ is_default: true, git_branch: 'main' }], + }) + ).toBe('feature/current') + }) + + it('falls back to the production branch Git ref', () => { + expect( + resolveRepoRef({ + currentBranch: {}, + branches: [{ is_default: true, git_branch: 'production' }], + }) + ).toBe('production') + }) + + it('leaves the final default-branch fallback to the archive endpoint', () => { + expect(resolveRepoRef({ currentBranch: {}, branches: [] })).toBeUndefined() + }) +}) diff --git a/apps/studio/lib/ai/repo-ref.ts b/apps/studio/lib/ai/repo-ref.ts new file mode 100644 index 00000000000..9ec9ceabe5e --- /dev/null +++ b/apps/studio/lib/ai/repo-ref.ts @@ -0,0 +1,11 @@ +type GitBranch = { git_branch?: string | null; is_default?: boolean } + +export function resolveRepoRef({ + currentBranch, + branches, +}: { + currentBranch?: GitBranch + branches?: GitBranch[] +}) { + return currentBranch?.git_branch || branches?.find((branch) => branch.is_default)?.git_branch +} diff --git a/apps/studio/lib/ai/tools/index.ts b/apps/studio/lib/ai/tools/index.ts index 6b487d0c8a7..fbdf93936dd 100644 --- a/apps/studio/lib/ai/tools/index.ts +++ b/apps/studio/lib/ai/tools/index.ts @@ -22,6 +22,8 @@ export const getTools = async ({ supportMode, isExplorerEnabled, signal, + repoTools, + hasRepoAccess = false, }: { projectRef: string connectionString: string @@ -37,6 +39,8 @@ export const getTools = async ({ // Required: tools fetched from the remote MCP server hold an HTTP connection // that is closed when this signal aborts (i.e. when the request ends). signal: AbortSignal + repoTools?: ToolSet + hasRepoAccess?: boolean }) => { // Always include studio tools let tools: ToolSet = getStudioTools({ projectRef, connectionString, authorization, aiOptInLevel }) @@ -80,12 +84,15 @@ export const getTools = async ({ ...getReportTools({ projectRef, authorization }), ...(isExplorerEnabled ? getNotebookTools({ projectRef, authorization }) : {}), ...(baseUrl ? getIncidentTools({ baseUrl }) : {}), + ...repoTools, } } // Filter all tools based on the (potentially modified) AI opt-in level const toolsWithSupport = supportMode ? { ...tools, ...getSupportLifecycleTools() } : tools - const filteredTools: ToolSet = filterToolsByOptInLevel(toolsWithSupport, aiOptInLevel) + const filteredTools: ToolSet = filterToolsByOptInLevel(toolsWithSupport, aiOptInLevel, { + hasRepoAccess, + }) return filteredTools } diff --git a/apps/studio/lib/ai/tools/repo-tools.test.ts b/apps/studio/lib/ai/tools/repo-tools.test.ts index 5975c8eed89..7036e64e22c 100644 --- a/apps/studio/lib/ai/tools/repo-tools.test.ts +++ b/apps/studio/lib/ai/tools/repo-tools.test.ts @@ -3,13 +3,12 @@ import { describe, expect, it, vi } from 'vitest' import { getRepoTools } from './repo-tools' -const executionOptions = (sandbox: Experimental_SandboxSession) => - ({ - toolCallId: 'call', - messages: [], - context: undefined, - experimental_sandbox: sandbox, - }) as any +const executionOptions = (sandbox: Experimental_SandboxSession) => ({ + toolCallId: 'call', + messages: [], + context: undefined as never, + experimental_sandbox: sandbox, +}) function mockSandbox(): Experimental_SandboxSession { return { @@ -45,7 +44,7 @@ function setup() { } describe('getRepoTools', () => { - it('quotes search input before running ripgrep', async () => { + it('quotes search input before running grep', async () => { const { tools, sandbox } = setup() await tools.search_repo.execute?.({ query: "user's query" }, executionOptions(sandbox)) @@ -90,7 +89,10 @@ describe('getRepoTools', () => { stderr: '', }) - await tools.open_pull_request.execute?.({ title: 'Fix the issue' }, executionOptions(sandbox)) + await tools.open_pull_request!.execute?.( + { title: 'Fix the issue', patch: 'diff --git a/a b/a' }, + executionOptions(sandbox) + ) expect(openPullRequest).toHaveBeenCalledWith( expect.objectContaining({ @@ -101,4 +103,17 @@ describe('getRepoTools', () => { }) ) }) + + it('omits PR creation without GitHub connection update permission', () => { + const tools = getRepoTools({ + connectionId: 42, + authorization: 'Bearer token', + baseRef: 'main', + headBranch: 'assistant/chat', + canOpenPullRequest: false, + }) + + expect(tools).not.toHaveProperty('open_pull_request') + expect(tools).toHaveProperty('read_repo_file') + }) }) diff --git a/apps/studio/lib/ai/tools/repo-tools.ts b/apps/studio/lib/ai/tools/repo-tools.ts index ed313c6067a..43f0bb1f99f 100644 --- a/apps/studio/lib/ai/tools/repo-tools.ts +++ b/apps/studio/lib/ai/tools/repo-tools.ts @@ -38,12 +38,14 @@ export function getRepoTools({ authorization, baseRef, headBranch, + canOpenPullRequest = true, openPullRequest = createGitHubPullRequest, }: { connectionId: number authorization: string baseRef: string headBranch: string + canOpenPullRequest?: boolean openPullRequest?: typeof createGitHubPullRequest }) { return { @@ -55,9 +57,9 @@ export function getRepoTools({ }), execute: async ({ query, glob }, options) => { const sandbox = requireSandbox(options) - const globArgs = glob ? `--glob ${shellQuote(glob)}` : '' + const globArgs = glob ? `--include=${shellQuote(glob)}` : '' const result = await sandbox.run({ - command: `rg --line-number --color never --max-count 100 ${globArgs} -- ${shellQuote(query)} . || [ $? -eq 1 ]`, + command: `grep -RIn -m 100 --exclude-dir=.git --binary-files=without-match ${globArgs} -- ${shellQuote(query)} . || [ $? -eq 1 ]`, abortSignal: options.abortSignal, }) return { matches: truncate(result.stdout) } @@ -92,36 +94,60 @@ export function getRepoTools({ content, abortSignal: options.abortSignal, }) - return { path, bytes: Buffer.byteLength(content) } - }, - }), - open_pull_request: tool({ - description: - 'Ask the user to open a pull request containing every repository change made in this chat.', - inputSchema: z.object({ - title: z.string().min(1).max(120), - body: z.string().max(20_000).optional(), - }), - execute: async ({ title, body }, options): Promise => { - const sandbox = requireSandbox(options) const { stdout: patch } = await sandbox.run({ - command: 'git diff --binary --no-ext-diff', + command: 'git add -N . && git diff --binary --no-ext-diff', abortSignal: options.abortSignal, }) - if (!patch.trim()) throw new Error('There are no repository changes to open') - - return openPullRequest({ - connectionId, - authorization, - baseRef, - headBranch, - title, - body, - patch, - signal: options.abortSignal, - }) + if (Buffer.byteLength(patch) > MAX_TOOL_OUTPUT) { + throw new Error('Repository change is too large to review in one pull request') + } + return { path, bytes: Buffer.byteLength(content), patch } }, }), + ...(canOpenPullRequest + ? { + open_pull_request: tool({ + description: + 'Ask the user to open a pull request containing every repository change made in this chat.', + inputSchema: z.object({ + title: z.string().min(1).max(120), + body: z.string().max(20_000).optional(), + patch: z + .string() + .min(1) + .max(MAX_TOOL_OUTPUT) + .describe('The exact patch returned by the final write_repo_file call.'), + }), + execute: async ( + { title, body, patch: proposedPatch }, + options + ): Promise => { + const sandbox = requireSandbox(options) + const { stdout: patch } = await sandbox.run({ + command: 'git add -N . && git diff --binary --no-ext-diff', + abortSignal: options.abortSignal, + }) + if (!patch.trim()) throw new Error('There are no repository changes to open') + if (patch !== proposedPatch) { + throw new Error( + 'Repository changes changed after the pull request preview was prepared' + ) + } + + return openPullRequest({ + connectionId, + authorization, + baseRef, + headBranch, + title, + body, + patch, + signal: options.abortSignal, + }) + }, + }), + } + : {}), } } diff --git a/apps/studio/lib/constants/index.ts b/apps/studio/lib/constants/index.ts index 2dc1022e152..2090f620608 100644 --- a/apps/studio/lib/constants/index.ts +++ b/apps/studio/lib/constants/index.ts @@ -74,6 +74,7 @@ export const OPT_IN_TAGS = { AI_SQL: 'AI_SQL_GENERATOR_OPT_IN', AI_DATA: 'AI_DATA_GENERATOR_OPT_IN', AI_LOG: 'AI_LOG_GENERATOR_OPT_IN', + AI_REPO: 'AI_REPO_ACCESS_OPT_IN', } export const GB = 1024 * 1024 * 1024 diff --git a/apps/studio/pages/api/ai/sql/generate-v4.ts b/apps/studio/pages/api/ai/sql/generate-v4.ts index d8900ab88db..9a4c8671b1a 100644 --- a/apps/studio/pages/api/ai/sql/generate-v4.ts +++ b/apps/studio/pages/api/ai/sql/generate-v4.ts @@ -1,11 +1,17 @@ import pgMeta from '@supabase/pg-meta' +import { PermissionAction } from '@supabase/shared-types/out/constants' import type { JwtPayload } from '@supabase/supabase-js' import { pipeUIMessageStreamToResponse, safeValidateUIMessages, toUIMessageStream } from 'ai' import { IS_PLATFORM } from 'common' import type { NextApiRequest, NextApiResponse } from 'next' import z from 'zod' +import { getBranches } from '@/data/branches/branches-query' +import { createGitHubRepoArchive } from '@/data/integrations/github-connection-repo' +import { getGitHubConnections } from '@/data/integrations/github-connections-query' +import { getPermissions } from '@/data/permissions/permissions-query' import { executeSql } from '@/data/sql/execute-sql-mutation' +import { doPermissionsCheck } from '@/hooks/misc/useCheckPermissions' import type { AiOptInLevel } from '@/hooks/misc/useOrgOptedIntoAi' import { getAIDetails } from '@/lib/ai/ai-details' import { NO_SCHEMA_ACCESS_MESSAGE } from '@/lib/ai/assistant-context' @@ -25,7 +31,11 @@ import { isKnownAssistantModelId, type AssistantModelId, } from '@/lib/ai/model.utils' +import { resolveRepoRef } from '@/lib/ai/repo-ref' +import { isSandboxConfigured } from '@/lib/ai/sandbox/sandbox-config' +import { createLazySandboxSession } from '@/lib/ai/sandbox/vercel-sandbox-session' import { getTools } from '@/lib/ai/tools' +import { getRepoTools } from '@/lib/ai/tools/repo-tools' import { apiWrapper } from '@/lib/api/apiWrapper' import { executeQuery } from '@/lib/api/self-hosted/query' import { getURL } from '@/lib/helpers' @@ -126,6 +136,8 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw let projectRegion: string | undefined let orgId: number | undefined let planId: string | undefined + let parentProjectRef: string | undefined + let hasRepoAccess = false if (!IS_PLATFORM) { aiOptInLevel = 'schema' @@ -143,6 +155,8 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw planId = aiDetails.planId projectIsSensitive = aiDetails.isSensitive projectRegion = aiDetails.region + parentProjectRef = aiDetails.parentProjectRef + hasRepoAccess = aiDetails.hasRepoAccess } catch (error) { return res.status(400).json({ error: 'There was an error fetching your organization details', @@ -180,6 +194,50 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw // is what tears down the remote MCP connection opened in getTools. res.on('close', () => abortController.abort()) + let sandbox: ReturnType | undefined + let repoTools: ReturnType | undefined + + if (IS_PLATFORM && hasRepoAccess && isSandboxConfigured() && orgId && chatId && authorization) { + try { + const headers = { 'Content-Type': 'application/json', Authorization: authorization } + const rootProjectRef = parentProjectRef ?? projectRef + const [connections, branches, permissions] = await Promise.all([ + getGitHubConnections({ organizationId: orgId }, abortController.signal, headers), + getBranches({ projectRef: rootProjectRef }, abortController.signal).catch(() => []), + getPermissions(abortController.signal, headers), + ]) + const connection = connections.find(({ project }) => project.ref === rootProjectRef) + + if (connection) { + const currentBranch = branches?.find((branch) => branch.project_ref === projectRef) + const ref = resolveRepoRef({ currentBranch, branches }) + const archive = await createGitHubRepoArchive({ + connectionId: connection.id, + ref: ref ?? undefined, + authorization, + signal: abortController.signal, + }) + sandbox = createLazySandboxSession({ projectRef, chatId, archive }) + repoTools = getRepoTools({ + connectionId: connection.id, + authorization, + baseRef: archive.ref, + headBranch: `supabase-assistant/${chatId.replace(/[^a-zA-Z0-9_-]/g, '-').slice(0, 48)}`, + canOpenPullRequest: doPermissionsCheck( + permissions, + PermissionAction.UPDATE, + 'integrations.github_connections', + undefined, + orgSlug, + rootProjectRef + ), + }) + } + } catch (error) { + console.error('Failed to prepare connected repository:', error) + } + } + const tools = await getTools({ projectRef, connectionString, @@ -190,6 +248,8 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw supportMode, isExplorerEnabled: explorerEnabled, signal: abortController.signal, + repoTools, + hasRepoAccess: Boolean(repoTools), }) // Get a list of all schemas to add to context @@ -242,6 +302,7 @@ async function handlePost(req: NextApiRequest, res: NextApiResponse, claims?: Jw onSpanCreated: (spanId) => { res.setHeader('x-braintrust-span-id', spanId) }, + sandbox, }) const stream = toUIMessageStream({ diff --git a/apps/studio/turbo.jsonc b/apps/studio/turbo.jsonc index 3a085d47704..7d12eb76070 100644 --- a/apps/studio/turbo.jsonc +++ b/apps/studio/turbo.jsonc @@ -71,6 +71,11 @@ "OPENAI_API_KEY", "BRAINTRUST_API_KEY", "BRAINTRUST_PROJECT_ID", + "TOOL_APPROVAL_SECRET", + "VERCEL_OIDC_TOKEN", + "VERCEL_TEAM_ID", + "VERCEL_PROJECT_ID", + "VERCEL_TOKEN", // Gates the dashboard assistant between the remote MCP server and the // legacy in-process one (see lib/ai/tools/mcp-tools.ts). "USE_REMOTE_MCP", @@ -130,12 +135,7 @@ "S3_PROTOCOL_ACCESS_KEY_ID", "S3_PROTOCOL_ACCESS_KEY_SECRET", ], - "outputs": [ - ".next/**", - "!.next/cache/**", - "!.next/dev/**/*", - "dist/**", - ], + "outputs": [".next/**", "!.next/cache/**", "!.next/dev/**/*", "dist/**"], }, }, }