feat(studio): connect assistant to GitHub repositories

This commit is contained in:
Saxon Fletcher committed 2026-08-21 17:00:10 +10:00
1 parent 1c24315dbc
commit 8cebaa7533
24 files changed
+502 -72

No files matched your search

@@ -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 = ({
</RadioGroup>
)}
/>
<FormField
control={control}
name="hasRepoAccess"
render={({ field }) => (
<div className="flex items-start justify-between gap-6 border-t pt-4">
<div>
<p className="text-sm font-medium text-foreground">Repository access</p>
<p className="text-sm text-foreground-light">
Allow the Assistant to read the connected GitHub repository and propose changes
through pull requests.
</p>
</div>
<Switch
aria-label="Repository access"
checked={field.value}
onCheckedChange={field.onChange}
disabled={disabled || aiOptInLevel === 'disabled'}
/>
</div>
)}
/>
</div>
</FormItemLayout>
)
@@ -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 (
<Form {...form}>
@@ -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 (
<Dialog open={visible} onOpenChange={onOpenChange}>
@@ -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<string, [string, string]> = {
'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 (
<Tool
icon={
@@ -72,10 +82,14 @@ function MessagePartTool({ toolPart }: { toolPart: ToolUIPart }) {
)
}
label={
<div>
{toolPart.state === 'input-streaming' ? 'Running ' : 'Ran '}
<span className="text-foreground-lighter">{`${toolPart.type.replace('tool-', '')}`}</span>
</div>
repoLabel ? (
repoLabel[toolPart.state === 'input-streaming' ? 0 : 1]
) : (
<div>
{toolPart.state === 'input-streaming' ? 'Running ' : 'Ran '}
<span className="text-foreground-lighter">{`${toolPart.type.replace('tool-', '')}`}</span>
</div>
)
}
/>
)
@@ -268,6 +282,36 @@ function MessagePartNotebookProposal({
)
}
function MessagePartOpenPullRequest({ toolPart }: { toolPart: ToolUIPart }) {
const { state, input, output } = toolPart
const { addToolApprovalResponse } = useMessageActionsContext()
if (state === 'input-streaming') return <MessagePartTool toolPart={toolPart} />
if (state === 'output-error') {
return <p className="text-xs text-danger">Failed to open pull request.</p>
}
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 (
<PullRequestRenderer
{...parsedInput.data}
url={parsedOutput.success ? parsedOutput.data.url : undefined}
number={parsedOutput.success ? parsedOutput.data.number : undefined}
confirmState={confirmState}
onApprove={onApprove}
onDeny={onDeny}
/>
)
}
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 <MessagePart.Tool toolPart={part} />
}
case 'tool-search_repo':
case 'tool-read_repo_file':
case 'tool-write_repo_file': {
return <MessagePart.Tool toolPart={part} />
}
case 'reasoning':
return <MessagePart.Reasoning reasoningPart={part} />
case 'text':
@@ -310,6 +360,9 @@ export function MessagePartSwitcher({
case 'tool-update_notebook': {
return <MessagePart.NotebookProposal toolPart={part} mode="update" />
}
case 'tool-open_pull_request': {
return <MessagePart.OpenPullRequest toolPart={part} />
}
case 'source-url':
case 'source-document':
@@ -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',
@@ -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(
<PullRequestRenderer
title="Fix auth handling"
patch="diff --git a/auth.ts b/auth.ts"
confirmState="approval-requested"
onApprove={onApprove}
onDeny={onDeny}
/>
)
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(
<PullRequestRenderer
title="Fix auth handling"
patch="diff"
url="https://github.com/acme/repo/pull/12"
number={12}
/>
)
expect(screen.getByRole('link', { name: /View pull request #12/ })).toHaveAttribute(
'href',
'https://github.com/acme/repo/pull/12'
)
expect(screen.queryByText('diff')).not.toBeInTheDocument()
})
})
@@ -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 (
<Confirm
className="my-4"
state={confirmState}
message="Assistant wants to open this pull request"
cancelLabel="Skip"
confirmLabel="Open PR"
confirmLabelLoading="Opening..."
onCancel={onDeny}
onConfirm={onApprove}
>
<div className="space-y-2 bg-surface-100 p-4">
<div className="flex items-start gap-2">
<GitPullRequest size={16} className="mt-0.5 shrink-0 text-foreground-light" />
<div className="min-w-0">
<p className="text-sm font-medium text-foreground">{title}</p>
{body && (
<p className="mt-1 whitespace-pre-wrap text-sm text-foreground-light">{body}</p>
)}
</div>
</div>
{url && (
<a
href={url}
target="_blank"
rel="noreferrer"
className="inline-flex items-center gap-1 text-sm text-brand-link hover:underline"
>
View pull request{number ? ` #${number}` : ''} <ExternalLink size={12} />
</a>
)}
{!url && (
<pre className="max-h-64 overflow-auto rounded border bg-surface-200 p-3 text-xs text-foreground-light">
{patch}
</pre>
)}
</div>
</Confirm>
)
}
@@ -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)
@@ -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: {
+9 -2
View File
@@ -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<typeof AIOptInSchema>
@@ -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),
}
}
@@ -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)
})
})
@@ -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),
}
}
+30
View File
@@ -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<typeof vi.fn>
let mockGetProjectSettings: ReturnType<typeof vi.fn>
let mockGetAiOptInLevel: ReturnType<typeof vi.fn>
let mockGetAiRepoAccess: ReturnType<typeof vi.fn>
let mockSubscriptionHasHipaaAddon: ReturnType<typeof vi.fn>
let mockCheckEntitlement: ReturnType<typeof vi.fn>
@@ -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 () => {
+28 -11
View File
@@ -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,
}
}
@@ -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}
+9 -1
View File
@@ -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.
`
+27
View File
@@ -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()
})
})
+11
View File
@@ -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
}
+8 -1
View File
@@ -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
}
+24 -9
View File
@@ -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')
})
})
+53 -27
View File
@@ -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<GitHubPullRequest> => {
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<GitHubPullRequest> => {
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,
})
},
}),
}
: {}),
}
}
+1
View File
@@ -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
@@ -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<typeof createLazySandboxSession> | undefined
let repoTools: ReturnType<typeof getRepoTools> | 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({
+6 -6
View File
@@ -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/**"],
},
},
}