feat(oauth): simplified logic for oauth, tested with send email and it works

This commit is contained in:
Waleed Latif
2025-03-06 21:38:11 -08:00
parent f244c96f9a
commit 3d9d125bba
15 changed files with 229 additions and 220 deletions
-127
View File
@@ -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<boolean> {
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 })
}
}
+8 -2
View File
@@ -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<GoogleIdToken>(acc.idToken)
if (decoded.email) {
-1
View File
@@ -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
+129
View File
@@ -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 })
}
}
@@ -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 <GoogleIcon className="h-4 w-4" />
case 'google-email':
return <GmailIcon className="h-4 w-4" />
case 'github':
return <GithubIcon className="h-4 w-4" />
case 'twitter':
return <TwitterIcon className="h-4 w-4" />
default:
return <ExternalLink className="h-4 w-4" />
}
@@ -186,6 +192,8 @@ export function CredentialSelector({
switch (provider) {
case 'google':
return 'Google'
case 'google-email':
return 'Gmail'
case 'github':
return 'GitHub'
case 'twitter':
+1 -16
View File
@@ -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 && (
<OAuthRequiredModal
isOpen={modalState.isOpen}
onClose={closeModal}
provider={modalState.provider}
toolName={modalState.toolName}
requiredScopes={modalState.requiredScopes}
/>
)}
<div className="relative w-full h-[calc(100vh-4rem)]">
<NotificationList />
<ReactFlow
+5 -4
View File
@@ -30,7 +30,7 @@ export const GmailBlock: BlockConfig<GmailToolResponse> = {
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<GmailToolResponse> = {
}
},
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
}
},
},
+1
View File
@@ -58,6 +58,7 @@ export type BlockOutput =
export interface ParamConfig {
type: ParamType
required: boolean
requiredForToolCall?: boolean
description?: string
schema?: {
type: string
+4 -2
View File
@@ -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<OAuthProvider, string> = {
github: 'GitHub',
google: 'Google',
'google-email': 'Gmail',
twitter: 'X (Twitter)',
}
@@ -34,6 +35,7 @@ const PROVIDER_NAMES: Record<OAuthProvider, string> = {
const PROVIDER_ICONS: Record<OAuthProvider, React.FC<React.SVGProps<SVGSVGElement>>> = {
github: GithubIcon,
google: GoogleIcon,
'google-email': GmailIcon,
twitter: XIcon,
}
+28
View File
@@ -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',
}
}
+7 -3
View File
@@ -9,17 +9,21 @@ export const gmailReadTool: ToolConfig<GmailReadParams, GmailToolResponse> = {
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',
},
},
+7 -2
View File
@@ -9,12 +9,17 @@ export const gmailSearchTool: ToolConfig<GmailSearchParams, GmailToolResponse> =
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',
+7 -3
View File
@@ -9,17 +9,21 @@ export const gmailSendTool: ToolConfig<GmailSendParams, GmailToolResponse> = {
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: {
+22 -58
View File
@@ -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<string, any>): Promise<void> {
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',
+1 -1
View File
@@ -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