From 3d9d125bba5613ae03a548df53e0aa536b36a572 Mon Sep 17 00:00:00 2001 From: Waleed Latif Date: Thu, 6 Mar 2025 21:38:09 -0800 Subject: [PATCH] feat(oauth): simplified logic for oauth, tested with send email and it works --- app/api/auth/oauth/check-tool/route.ts | 127 ----------------- app/api/auth/oauth/credentials/route.ts | 10 +- app/api/auth/oauth/disconnect/route.ts | 1 - app/api/auth/oauth/token/route.ts | 129 ++++++++++++++++++ .../components/credential-selector.tsx | 10 +- app/w/[id]/workflow.tsx | 17 +-- blocks/blocks/gmail.ts | 9 +- blocks/types.ts | 1 + components/ui/oauth-required-modal.tsx | 6 +- lib/oauth-utils.ts | 28 ++++ tools/gmail/read.ts | 10 +- tools/gmail/search.ts | 9 +- tools/gmail/send.ts | 10 +- tools/index.ts | 80 +++-------- tools/types.ts | 2 +- 15 files changed, 229 insertions(+), 220 deletions(-) delete mode 100644 app/api/auth/oauth/check-tool/route.ts create mode 100644 app/api/auth/oauth/token/route.ts create mode 100644 lib/oauth-utils.ts diff --git a/app/api/auth/oauth/check-tool/route.ts b/app/api/auth/oauth/check-tool/route.ts deleted file mode 100644 index 39f3bcc1c7..0000000000 --- a/app/api/auth/oauth/check-tool/route.ts +++ /dev/null @@ -1,127 +0,0 @@ -import { NextRequest, NextResponse } from 'next/server' -import { and, eq } from 'drizzle-orm' -import { getSession } from '@/lib/auth' -import { db } from '@/db' -import { account } from '@/db/schema' -import { OAuthProvider } from '@/tools/types' - -/** - * Check if the user has authorized a specific OAuth provider - */ -async function hasAuthorizedProvider( - userId: string, - provider: OAuthProvider, - requiredScopes?: string[], - credentialId?: string -): Promise { - try { - // If a specific credential ID is provided, check if it exists and belongs to the user - if (credentialId) { - const credential = await db - .select() - .from(account) - .where(and(eq(account.id, credentialId), eq(account.userId, userId))) - .limit(1) - - return credential.length > 0 - } - - // Otherwise, determine the appropriate provider ID based on scopes - let featureType = 'default' - if (requiredScopes && requiredScopes.length > 0) { - if (requiredScopes.some((scope) => scope.includes('repo'))) { - featureType = 'repo' - } else if (requiredScopes.some((scope) => scope.includes('workflow'))) { - featureType = 'workflow' - } else if ( - requiredScopes.some((scope) => scope.includes('gmail') || scope.includes('mail')) - ) { - featureType = 'email' - } else if (requiredScopes.some((scope) => scope.includes('calendar'))) { - featureType = 'calendar' - } else if (requiredScopes.some((scope) => scope.includes('drive'))) { - featureType = 'drive' - } else if (requiredScopes.some((scope) => scope.includes('write'))) { - featureType = 'write' - } else if (requiredScopes.some((scope) => scope.includes('read'))) { - featureType = 'read' - } - } - - // Construct the provider ID based on the provider and feature type - const providerId = `${provider}-${featureType}` - - // Check if the user has this provider account - const accounts = await db - .select() - .from(account) - .where(and(eq(account.userId, userId), eq(account.providerId, providerId))) - .limit(1) - - return accounts.length > 0 - } catch (error) { - console.error('Error checking OAuth authorization:', error) - return false - } -} - -/** - * API route to check if a tool requires OAuth and if the user is authorized - */ -export async function POST(request: NextRequest) { - try { - // Get the session - const session = await getSession() - - // Check if the user is authenticated - if (!session?.user?.id) { - return NextResponse.json( - { requiresAuth: true, isAuthorized: false, error: 'User not authenticated' }, - { status: 401 } - ) - } - - // Get the tool and credential ID from the request body - const { tool, credentialId } = await request.json() - - // Check if the tool requires OAuth - if (!tool.oauth || !tool.oauth.required) { - return NextResponse.json({ requiresAuth: false, isAuthorized: true }, { status: 200 }) - } - - // Get the provider and required scopes - const provider = tool.oauth.provider - const requiredScopes = tool.oauth.additionalScopes || [] - - // Check if the user has authorized this provider - const isAuthorized = await hasAuthorizedProvider( - session.user.id, - provider, - requiredScopes, - credentialId - ) - - // Return the authorization status - if (isAuthorized) { - return NextResponse.json({ requiresAuth: true, isAuthorized: true }, { status: 200 }) - } else { - return NextResponse.json( - { - requiresAuth: true, - isAuthorized: false, - error: JSON.stringify({ - type: 'oauth_required', - provider, - toolId: tool.id, - toolName: tool.name, - requiredScopes, - }), - }, - { status: 200 } - ) - } - } catch (error) { - console.error('Error checking OAuth authorization:', error) - return NextResponse.json({ error: 'Internal server error' }, { status: 500 }) - } -} diff --git a/app/api/auth/oauth/credentials/route.ts b/app/api/auth/oauth/credentials/route.ts index b7bee98af1..7535e8e1f7 100644 --- a/app/api/auth/oauth/credentials/route.ts +++ b/app/api/auth/oauth/credentials/route.ts @@ -2,6 +2,7 @@ import { NextRequest, NextResponse } from 'next/server' import { and, eq, like } from 'drizzle-orm' import { jwtDecode } from 'jwt-decode' import { getSession } from '@/lib/auth' +import { parseProvider } from '@/lib/oauth-utils' import { db } from '@/db' import { account } from '@/db/schema' import { OAuthProvider } from '@/tools/types' @@ -32,11 +33,16 @@ export async function GET(request: NextRequest) { return NextResponse.json({ error: 'Provider is required' }, { status: 400 }) } + // Parse the provider to get base provider and feature type + const { baseProvider } = parseProvider(provider) + // Get all accounts for this user and provider const accounts = await db .select() .from(account) - .where(and(eq(account.userId, session.user.id), like(account.providerId, `${provider}-%`))) + .where( + and(eq(account.userId, session.user.id), like(account.providerId, `${baseProvider}-%`)) + ) // Transform accounts into credentials const credentials = await Promise.all( @@ -46,7 +52,7 @@ export async function GET(request: NextRequest) { // For Google accounts, try to get the email from the ID token let name = acc.accountId - if (provider === 'google' && acc.idToken) { + if (baseProvider === 'google' && acc.idToken) { try { const decoded = jwtDecode(acc.idToken) if (decoded.email) { diff --git a/app/api/auth/oauth/disconnect/route.ts b/app/api/auth/oauth/disconnect/route.ts index 8055f80041..0bc3f7f6d0 100644 --- a/app/api/auth/oauth/disconnect/route.ts +++ b/app/api/auth/oauth/disconnect/route.ts @@ -3,7 +3,6 @@ import { and, eq, like } from 'drizzle-orm' import { getSession } from '@/lib/auth' import { db } from '@/db' import { account } from '@/db/schema' -import { OAuthProvider } from '@/tools/types' /** * Disconnect an OAuth provider for the current user diff --git a/app/api/auth/oauth/token/route.ts b/app/api/auth/oauth/token/route.ts new file mode 100644 index 0000000000..e838ced248 --- /dev/null +++ b/app/api/auth/oauth/token/route.ts @@ -0,0 +1,129 @@ +import { NextRequest, NextResponse } from 'next/server' +import { and, eq } from 'drizzle-orm' +import { getSession } from '@/lib/auth' +import { db } from '@/db' +import { account } from '@/db/schema' + +/** + * Get an access token for a specific credential + */ +export async function POST(request: NextRequest) { + try { + // Get the session + const session = await getSession() + + // Check if the user is authenticated + if (!session?.user?.id) { + return NextResponse.json({ error: 'User not authenticated' }, { status: 401 }) + } + + // Get the credential ID from the request body + const { credentialId } = await request.json() + + if (!credentialId) { + return NextResponse.json({ error: 'Credential ID is required' }, { status: 400 }) + } + + // Get the credential from the database + const credentials = await db + .select() + .from(account) + .where(and(eq(account.id, credentialId), eq(account.userId, session.user.id))) + .limit(1) + + if (!credentials.length) { + return NextResponse.json({ error: 'Credential not found' }, { status: 404 }) + } + + const credential = credentials[0] + + // Check if we need to refresh the token + const expiresAt = credential.accessTokenExpiresAt + const now = new Date() + const needsRefresh = !expiresAt || expiresAt <= now + + if (needsRefresh && credential.refreshToken) { + try { + // Get the provider from the providerId (e.g., 'google-email' -> 'google') + const provider = credential.providerId.split('-')[0] + + // Determine the token endpoint based on the provider + let tokenEndpoint: string + let clientId: string | undefined + let clientSecret: string | undefined + + switch (provider) { + case 'google': + tokenEndpoint = 'https://oauth2.googleapis.com/token' + clientId = process.env.GOOGLE_CLIENT_ID + clientSecret = process.env.GOOGLE_CLIENT_SECRET + break + case 'github': + tokenEndpoint = 'https://github.com/login/oauth/access_token' + clientId = process.env.GITHUB_CLIENT_ID + clientSecret = process.env.GITHUB_CLIENT_SECRET + break + case 'twitter': + tokenEndpoint = 'https://api.twitter.com/2/oauth2/token' + clientId = process.env.TWITTER_CLIENT_ID + clientSecret = process.env.TWITTER_CLIENT_SECRET + break + default: + throw new Error(`Unsupported provider: ${provider}`) + } + + if (!clientId || !clientSecret) { + throw new Error(`Missing client credentials for provider: ${provider}`) + } + + // Refresh the token + const response = await fetch(tokenEndpoint, { + method: 'POST', + headers: { + 'Content-Type': 'application/x-www-form-urlencoded', + ...(provider === 'github' && { + Accept: 'application/json', + }), + }, + body: new URLSearchParams({ + client_id: clientId, + client_secret: clientSecret, + grant_type: 'refresh_token', + refresh_token: credential.refreshToken, + }).toString(), + }) + + if (!response.ok) { + throw new Error('Failed to refresh token') + } + + const data = await response.json() + + // Update the credential in the database + await db + .update(account) + .set({ + accessToken: data.access_token, + accessTokenExpiresAt: data.expires_in + ? new Date(Date.now() + data.expires_in * 1000) + : null, + refreshToken: data.refresh_token || credential.refreshToken, // Some providers don't return a new refresh token + updatedAt: new Date(), + }) + .where(eq(account.id, credentialId)) + + // Return the new access token + return NextResponse.json({ accessToken: data.access_token }, { status: 200 }) + } catch (error) { + console.error('Error refreshing token:', error) + return NextResponse.json({ error: 'Failed to refresh access token' }, { status: 500 }) + } + } + + // Return the current access token + return NextResponse.json({ accessToken: credential.accessToken }, { status: 200 }) + } catch (error) { + console.error('Error getting access token:', error) + return NextResponse.json({ error: 'Internal server error' }, { status: 500 }) + } +} diff --git a/app/w/[id]/components/workflow-block/components/sub-block/components/credential-selector.tsx b/app/w/[id]/components/workflow-block/components/sub-block/components/credential-selector.tsx index 309c1a30bc..27d1e0f8e4 100644 --- a/app/w/[id]/components/workflow-block/components/sub-block/components/credential-selector.tsx +++ b/app/w/[id]/components/workflow-block/components/sub-block/components/credential-selector.tsx @@ -2,7 +2,7 @@ import { useCallback, useEffect, useRef, useState } from 'react' import { Check, ChevronDown, ExternalLink, Key, RefreshCw } from 'lucide-react' -import { GoogleIcon } from '@/components/icons' +import { GithubIcon, GmailIcon, GoogleIcon, xIcon as TwitterIcon } from '@/components/icons' import { Button } from '@/components/ui/button' import { Command, @@ -176,6 +176,12 @@ export function CredentialSelector({ switch (provider) { case 'google': return + case 'google-email': + return + case 'github': + return + case 'twitter': + return default: return } @@ -186,6 +192,8 @@ export function CredentialSelector({ switch (provider) { case 'google': return 'Google' + case 'google-email': + return 'Gmail' case 'github': return 'GitHub' case 'twitter': diff --git a/app/w/[id]/workflow.tsx b/app/w/[id]/workflow.tsx index 396eef89e7..72b840df84 100644 --- a/app/w/[id]/workflow.tsx +++ b/app/w/[id]/workflow.tsx @@ -13,10 +13,9 @@ import ReactFlow, { } from 'reactflow' import 'reactflow/dist/style.css' import { OAuthRequiredModal } from '@/components/ui/oauth-required-modal' -import { useOAuthErrorHandler } from '@/lib/oauth' import { useNotificationStore } from '@/stores/notifications/store' import { useGeneralStore } from '@/stores/settings/general/store' -import { getSyncManagers, initializeSyncManagers, isSyncInitialized } from '@/stores/sync-registry' +import { initializeSyncManagers, isSyncInitialized } from '@/stores/sync-registry' import { useWorkflowRegistry } from '@/stores/workflows/registry/store' import { useSubBlockStore } from '@/stores/workflows/subblock/store' import { useWorkflowStore } from '@/stores/workflows/workflow/store' @@ -54,9 +53,6 @@ function WorkflowContent() { useWorkflowStore() const { setValue: setSubBlockValue } = useSubBlockStore() - // Add OAuth error handling - const { modalState, handleOAuthError, closeModal } = useOAuthErrorHandler() - // Initialize workflow useEffect(() => { if (typeof window !== 'undefined') { @@ -337,17 +333,6 @@ function WorkflowContent() { return ( <> - {/* Add the OAuth modal */} - {modalState.isOpen && modalState.provider && ( - - )} -
= { title: 'Gmail Account', type: 'oauth-input', layout: 'full', - provider: 'google', + provider: 'google-email', serviceId: 'gmail', requiredScopes: [ 'https://www.googleapis.com/auth/gmail.send', @@ -106,10 +106,11 @@ export const GmailBlock: BlockConfig = { } }, params: (params) => { - // Add the credential ID to the params + // Pass the credential directly from the credential field + const { credential, ...rest } = params return { - ...params, - _credentialId: params.credential, + ...rest, + credential, // Keep the credential parameter } }, }, diff --git a/blocks/types.ts b/blocks/types.ts index 4f0c35a362..f5d4e5c745 100644 --- a/blocks/types.ts +++ b/blocks/types.ts @@ -58,6 +58,7 @@ export type BlockOutput = export interface ParamConfig { type: ParamType required: boolean + requiredForToolCall?: boolean description?: string schema?: { type: string diff --git a/components/ui/oauth-required-modal.tsx b/components/ui/oauth-required-modal.tsx index 71a1da616c..2638ffce4e 100644 --- a/components/ui/oauth-required-modal.tsx +++ b/components/ui/oauth-required-modal.tsx @@ -1,7 +1,7 @@ 'use client' import { Check } from 'lucide-react' -import { GithubIcon, GoogleIcon, xIcon as XIcon } from '@/components/icons' +import { GithubIcon, GmailIcon, GoogleIcon, xIcon as XIcon } from '@/components/icons' import { Button } from '@/components/ui/button' import { Dialog, @@ -11,7 +11,7 @@ import { DialogHeader, DialogTitle, } from '@/components/ui/dialog' -import { loadFromStorage, saveToStorage } from '@/stores/workflows/persistence' +import { saveToStorage } from '@/stores/workflows/persistence' import { OAuthProvider } from '@/tools/types' export interface OAuthRequiredModalProps { @@ -27,6 +27,7 @@ export interface OAuthRequiredModalProps { const PROVIDER_NAMES: Record = { github: 'GitHub', google: 'Google', + 'google-email': 'Gmail', twitter: 'X (Twitter)', } @@ -34,6 +35,7 @@ const PROVIDER_NAMES: Record = { const PROVIDER_ICONS: Record>> = { github: GithubIcon, google: GoogleIcon, + 'google-email': GmailIcon, twitter: XIcon, } diff --git a/lib/oauth-utils.ts b/lib/oauth-utils.ts new file mode 100644 index 0000000000..a1b85875f5 --- /dev/null +++ b/lib/oauth-utils.ts @@ -0,0 +1,28 @@ +import { OAuthProvider } from '@/tools/types' + +interface ProviderConfig { + baseProvider: string + featureType: string +} + +/** + * Parse a provider string into its base provider and feature type + * This is a server-safe utility that can be used in both client and server code + */ +export function parseProvider(provider: OAuthProvider): ProviderConfig { + // Handle compound providers (e.g., 'google-email' -> { baseProvider: 'google', featureType: 'email' }) + const [base, feature] = provider.split('-') + + if (feature) { + return { + baseProvider: base, + featureType: feature, + } + } + + // For simple providers, use 'default' as feature type + return { + baseProvider: provider, + featureType: 'default', + } +} diff --git a/tools/gmail/read.ts b/tools/gmail/read.ts index 7f3f4e57ae..748ca13348 100644 --- a/tools/gmail/read.ts +++ b/tools/gmail/read.ts @@ -9,17 +9,21 @@ export const gmailReadTool: ToolConfig = { description: 'Read emails from Gmail', version: '1.0.0', + oauth: { + required: true, + provider: 'google-email', + additionalScopes: ['https://www.googleapis.com/auth/gmail.readonly'], + }, + params: { accessToken: { type: 'string', required: true, - requiredForToolCall: true, - description: 'OAuth access token for Gmail API', + description: 'Access token for Gmail API', }, messageId: { type: 'string', required: true, - requiredForToolCall: true, description: 'ID of the message to read', }, }, diff --git a/tools/gmail/search.ts b/tools/gmail/search.ts index 82db1f3c19..f6bfaf385a 100644 --- a/tools/gmail/search.ts +++ b/tools/gmail/search.ts @@ -9,12 +9,17 @@ export const gmailSearchTool: ToolConfig = description: 'Search emails in Gmail', version: '1.0.0', + oauth: { + required: true, + provider: 'google-email', + additionalScopes: ['https://www.googleapis.com/auth/gmail.readonly'], + }, + params: { accessToken: { type: 'string', required: true, - requiredForToolCall: true, - description: 'OAuth access token for Gmail API', + description: 'Access token for Gmail API', }, query: { type: 'string', diff --git a/tools/gmail/send.ts b/tools/gmail/send.ts index e3660b18d9..0624a14f6f 100644 --- a/tools/gmail/send.ts +++ b/tools/gmail/send.ts @@ -9,17 +9,21 @@ export const gmailSendTool: ToolConfig = { description: 'Send emails using Gmail', version: '1.0.0', + oauth: { + required: true, + provider: 'google-email', + additionalScopes: ['https://www.googleapis.com/auth/gmail.send'], + }, + params: { accessToken: { type: 'string', required: true, - requiredForToolCall: true, - description: 'OAuth access token for Gmail API', + description: 'Access token for Gmail API', }, to: { type: 'string', required: true, - requiredForToolCall: true, description: 'Recipient email address', }, subject: { diff --git a/tools/index.ts b/tools/index.ts index 66779cca9d..f436bac2db 100644 --- a/tools/index.ts +++ b/tools/index.ts @@ -1,4 +1,3 @@ -import { OAuthRequiredError } from '@/lib/oauth' import { useCustomToolsStore } from '@/stores/custom-tools/store' import { useEnvironmentStore } from '@/stores/settings/environment/store' import { visionTool as crewAIVision } from './crewai/vision' @@ -275,58 +274,6 @@ function getCustomTool(customToolId: string): ToolConfig | undefined { } } -// Function to check OAuth via API -async function checkOAuth(tool: any, params: Record): Promise { - if (!tool.oauth || !tool.oauth.required) { - return // No OAuth required - } - - // Check if a credential ID is provided - const credentialId = params._credentialId - - try { - // Call the API to check if the user is authorized - const response = await fetch('/api/auth/oauth/check-tool', { - method: 'POST', - headers: { - 'Content-Type': 'application/json', - }, - body: JSON.stringify({ - tool, - credentialId, // Pass the credential ID if provided - }), - }) - - if (!response.ok) { - throw new Error('Failed to check OAuth authorization') - } - - const data = await response.json() - - if (!data.isAuthorized) { - // Parse the error to get OAuth details - const errorDetails = JSON.parse(data.error || '{}') - - if (errorDetails.type === 'oauth_required') { - throw new Error( - JSON.stringify({ - type: 'oauth_required', - provider: errorDetails.provider, - toolId: errorDetails.toolId, - toolName: errorDetails.toolName, - requiredScopes: errorDetails.requiredScopes, - }) - ) - } else { - throw new Error('OAuth authorization required') - } - } - } catch (error) { - console.error('Error checking OAuth authorization:', error) - throw error - } -} - // Execute a tool by calling either the proxy for external APIs or directly for internal routes export async function executeTool( toolId: string, @@ -344,11 +291,6 @@ export async function executeTool( throw new Error(`Tool not found: ${toolId}`) } - // Check OAuth requirements before executing the tool - if (tool.oauth?.required && !isBrowser()) { - await checkOAuth(tool, params) - } - // For custom tools, try direct execution in browser first if available if (toolId.startsWith('custom_') && tool.directExecution) { const directResult = await tool.directExecution(params) @@ -516,6 +458,28 @@ async function handleProxyRequest( throw new Error('NEXT_PUBLIC_APP_URL environment variable is not set') } + // If we have a credential parameter, fetch the access token + if (params.credential) { + try { + const response = await fetch(`${baseUrl}/api/auth/oauth/token`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ credentialId: params.credential }), + }) + + if (!response.ok) { + throw new Error('Failed to fetch access token') + } + + const data = await response.json() + params.accessToken = data.accessToken + delete params.credential + } catch (error) { + console.error('Error fetching access token:', error) + throw error + } + } + const proxyUrl = new URL('/api/proxy', baseUrl).toString() const response = await fetch(proxyUrl, { method: 'POST', diff --git a/tools/types.ts b/tools/types.ts index 758ccfdf4d..a8c96d7578 100644 --- a/tools/types.ts +++ b/tools/types.ts @@ -1,5 +1,5 @@ export type HttpMethod = 'GET' | 'POST' | 'PUT' | 'DELETE' | 'PATCH' -export type OAuthProvider = 'google' | 'github' | 'twitter' +export type OAuthProvider = 'google' | 'google-email' | 'github' | 'twitter' export interface ToolResponse { success: boolean // Whether the tool execution was successful