From 3df326f8bac99b1d4a03cc72e317906444a953db Mon Sep 17 00:00:00 2001 From: Waleed Latif Date: Wed, 5 Mar 2025 15:36:14 -0800 Subject: [PATCH] feat(oauth): add oauth as param for tool, if required & user doesn't have access, modal will popup that allows user to grant access for that group of scopes --- app/api/auth/oauth/check/route.ts | 47 +++ components/ui/oauth-required-modal.tsx | 128 ++++++++ lib/auth-client.ts | 5 +- lib/auth.ts | 92 +++++- lib/oauth.ts | 399 +++++++++++++++++++++++++ tools/index.ts | 7 + tools/types.ts | 10 + 7 files changed, 686 insertions(+), 2 deletions(-) create mode 100644 app/api/auth/oauth/check/route.ts create mode 100644 components/ui/oauth-required-modal.tsx create mode 100644 lib/oauth.ts diff --git a/app/api/auth/oauth/check/route.ts b/app/api/auth/oauth/check/route.ts new file mode 100644 index 0000000000..69bd8d392b --- /dev/null +++ b/app/api/auth/oauth/check/route.ts @@ -0,0 +1,47 @@ +import { NextRequest, NextResponse } from 'next/server' +import { getSession } from '@/lib/auth' +import { hasAuthorizedProviderServer } from '@/lib/oauth' +import { OAuthProvider } from '@/tools/types' + +/** + * API endpoint to check if a user has authorized a specific OAuth provider + * + * @param request - The request object with provider and optional scopes + * @returns JSON response with authorization status + */ +export async function GET(request: NextRequest) { + try { + // Get the session + const session = await getSession() + + if (!session?.user?.id) { + return NextResponse.json({ isAuthorized: false, error: 'Not authenticated' }, { status: 401 }) + } + + // Get the provider from the query string + const url = new URL(request.url) + const provider = url.searchParams.get('provider') as OAuthProvider | null + + if (!provider) { + return NextResponse.json( + { isAuthorized: false, error: 'Provider is required' }, + { status: 400 } + ) + } + + // Get optional scopes from the query string + const scopesParam = url.searchParams.get('scopes') + const requiredScopes = scopesParam ? scopesParam.split(',') : undefined + + // Check if the user has authorized this provider with the required scopes + const isAuthorized = await hasAuthorizedProviderServer(provider, requiredScopes) + + return NextResponse.json({ isAuthorized }) + } catch (error) { + console.error('Error checking OAuth authorization:', error) + return NextResponse.json( + { isAuthorized: false, error: 'Internal server error' }, + { status: 500 } + ) + } +} diff --git a/components/ui/oauth-required-modal.tsx b/components/ui/oauth-required-modal.tsx new file mode 100644 index 0000000000..b598380bc8 --- /dev/null +++ b/components/ui/oauth-required-modal.tsx @@ -0,0 +1,128 @@ +'use client' + +import { GithubIcon, GoogleIcon, xIcon as XIcon } from '@/components/icons' +import { Button } from '@/components/ui/button' +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from '@/components/ui/dialog' +import { client } from '@/lib/auth-client' +import { OAuthProvider } from '@/tools/types' + +export interface OAuthRequiredModalProps { + isOpen: boolean + onClose: () => void + provider: OAuthProvider + toolName: string + requiredScopes?: string[] +} + +// Map of provider names to friendly display names +const PROVIDER_NAMES: Record = { + github: 'GitHub', + google: 'Google', + twitter: 'X (Twitter)', +} + +// Map of provider to icons +const PROVIDER_ICONS: Record>> = { + github: GithubIcon, + google: GoogleIcon, + twitter: XIcon, +} + +export function OAuthRequiredModal({ + isOpen, + onClose, + provider, + toolName, + requiredScopes = [], +}: OAuthRequiredModalProps) { + const providerName = PROVIDER_NAMES[provider] || provider + const ProviderIcon = PROVIDER_ICONS[provider] + + const handleAuth = async () => { + try { + // Determine the appropriate providerId based on the provider and required scopes + let featureType = 'default' + + // Simple scope-based feature detection (expand as needed) + 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 providerId based on the provider and feature type + const providerId = `${provider}-${featureType}` + + // Begin OAuth flow with the appropriate provider + await client.signIn.oauth2({ + providerId, + callbackURL: window.location.href, // Return to the current page after auth + }) + } catch (error) { + console.error('OAuth login error:', error) + } + } + + return ( + !open && onClose()}> + + + Additional Access Required + + The "{toolName}" tool requires access to your {providerName} account to function + properly. + + +
+
+
+ +
+
+

Connect {providerName}

+

Authorize access to use this tool

+
+
+ + {requiredScopes.length > 0 && ( +
+ Permissions requested +
    + {requiredScopes.map((scope) => ( +
  • {scope}
  • + ))} +
+
+ )} +
+ + + + +
+
+ ) +} diff --git a/lib/auth-client.ts b/lib/auth-client.ts index efb1fe0dda..451e4c47a4 100644 --- a/lib/auth-client.ts +++ b/lib/auth-client.ts @@ -1,6 +1,9 @@ +import { genericOAuthClient } from 'better-auth/client/plugins' import { createAuthClient } from 'better-auth/react' -export const client = createAuthClient() +export const client = createAuthClient({ + plugins: [genericOAuthClient()], +}) export const { useSession } = client // Export commonly used hooks and methods diff --git a/lib/auth.ts b/lib/auth.ts index baf0647609..453da58767 100644 --- a/lib/auth.ts +++ b/lib/auth.ts @@ -2,6 +2,7 @@ import { headers } from 'next/headers' import { betterAuth } from 'better-auth' import { drizzleAdapter } from 'better-auth/adapters/drizzle' import { nextCookies } from 'better-auth/next-js' +import { genericOAuth } from 'better-auth/plugins' import { Resend } from 'resend' import { db } from '@/db' import * as schema from '@/db/schema' @@ -27,10 +28,20 @@ export const auth = betterAuth({ github: { clientId: process.env.GITHUB_CLIENT_ID as string, clientSecret: process.env.GITHUB_CLIENT_SECRET as string, + scopes: ['user:email', 'repo'], }, google: { clientId: process.env.GOOGLE_CLIENT_ID as string, clientSecret: process.env.GOOGLE_CLIENT_SECRET as string, + scopes: [ + 'https://www.googleapis.com/auth/userinfo.email', + 'https://www.googleapis.com/auth/userinfo.profile', + ], + }, + twitter: { + clientId: process.env.TWITTER_CLIENT_ID as string, + clientSecret: process.env.TWITTER_CLIENT_SECRET as string, + scopes: ['tweet.read', 'users.read'], }, }, emailAndPassword: { @@ -89,7 +100,86 @@ export const auth = betterAuth({ } }, }, - plugins: [nextCookies()], + plugins: [ + nextCookies(), + genericOAuth({ + config: [ + { + providerId: 'github-repo', + clientId: process.env.GITHUB_CLIENT_ID as string, + clientSecret: process.env.GITHUB_CLIENT_SECRET as string, + authorizationUrl: 'https://github.com/login/oauth/authorize', + tokenUrl: 'https://github.com/login/oauth/access_token', + userInfoUrl: 'https://api.github.com/user', + scopes: ['user:email', 'repo'], + }, + { + providerId: 'github-workflow', + clientId: process.env.GITHUB_CLIENT_ID as string, + clientSecret: process.env.GITHUB_CLIENT_SECRET as string, + authorizationUrl: 'https://github.com/login/oauth/authorize', + tokenUrl: 'https://github.com/login/oauth/access_token', + userInfoUrl: 'https://api.github.com/user', + scopes: ['workflow', 'repo'], + }, + + // Google providers for different purposes + { + providerId: 'google-email', + clientId: process.env.GOOGLE_CLIENT_ID as string, + clientSecret: process.env.GOOGLE_CLIENT_SECRET as string, + discoveryUrl: 'https://accounts.google.com/.well-known/openid-configuration', + scopes: [ + 'https://www.googleapis.com/auth/userinfo.email', + 'https://www.googleapis.com/auth/userinfo.profile', + 'https://www.googleapis.com/auth/gmail.send', + ], + }, + { + providerId: 'google-calendar', + clientId: process.env.GOOGLE_CLIENT_ID as string, + clientSecret: process.env.GOOGLE_CLIENT_SECRET as string, + discoveryUrl: 'https://accounts.google.com/.well-known/openid-configuration', + scopes: [ + 'https://www.googleapis.com/auth/userinfo.email', + 'https://www.googleapis.com/auth/userinfo.profile', + 'https://www.googleapis.com/auth/calendar', + ], + }, + { + providerId: 'google-drive', + clientId: process.env.GOOGLE_CLIENT_ID as string, + clientSecret: process.env.GOOGLE_CLIENT_SECRET as string, + discoveryUrl: 'https://accounts.google.com/.well-known/openid-configuration', + scopes: [ + 'https://www.googleapis.com/auth/userinfo.email', + 'https://www.googleapis.com/auth/userinfo.profile', + 'https://www.googleapis.com/auth/drive', + ], + }, + + // Twitter providers + { + providerId: 'twitter-read', + clientId: process.env.TWITTER_CLIENT_ID as string, + clientSecret: process.env.TWITTER_CLIENT_SECRET as string, + authorizationUrl: 'https://twitter.com/i/oauth2/authorize', + tokenUrl: 'https://api.twitter.com/2/oauth2/token', + userInfoUrl: 'https://api.twitter.com/2/users/me', + scopes: ['tweet.read', 'users.read'], + }, + { + providerId: 'twitter-write', + clientId: process.env.TWITTER_CLIENT_ID as string, + clientSecret: process.env.TWITTER_CLIENT_SECRET as string, + authorizationUrl: 'https://twitter.com/i/oauth2/authorize', + tokenUrl: 'https://api.twitter.com/2/oauth2/token', + userInfoUrl: 'https://api.twitter.com/2/users/me', + scopes: ['tweet.read', 'tweet.write', 'users.read', 'offline.access'], + }, + ], + }), + ], pages: { signIn: '/login', signUp: '/signup', diff --git a/lib/oauth.ts b/lib/oauth.ts new file mode 100644 index 0000000000..035f7456cd --- /dev/null +++ b/lib/oauth.ts @@ -0,0 +1,399 @@ +import { useCallback, useEffect, useState } from 'react' +import { and, eq } from 'drizzle-orm' +import { getSession } from '@/lib/auth' +import { useSession } from '@/lib/auth-client' +import { db } from '@/db' +import { account } from '@/db/schema' +import { OAuthProvider } from '@/tools/types' + +/** + * Interface for the OAuth error structure + */ +export interface OAuthRequiredError { + type: 'oauth_required' + provider: OAuthProvider + toolId: string + toolName: string + requiredScopes?: string[] +} + +/** + * Check if the user has authorized the required OAuth provider with necessary scopes (server-side) + * + * @param provider - The OAuth provider to check + * @param requiredScopes - Optional scopes to check + * @returns Boolean indicating if the provider is authorized with required scopes + */ +export async function hasAuthorizedProviderServer( + provider: OAuthProvider, + requiredScopes?: string[] +): Promise { + try { + // Get the session + const session = await getSession() + + // If not authenticated, return false + if (!session?.user?.id) { + return false + } + + // Determine the appropriate feature type based on the 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' + } + } + + // We check the appropriate provider ID based on the feature type + const providerId = `${provider}-${featureType}` + + // Check if there's an account for this provider for the user + const accounts = await db + .select() + .from(account) + .where(and(eq(account.userId, session.user.id), eq(account.providerId, providerId))) + + return accounts.length > 0 + } catch (error) { + console.error('Error checking provider authorization:', error) + return false + } +} + +/** + * Check if a tool requires OAuth and if the user has authorized it (server-side) + * + * @param tool - The tool configuration + * @returns Object indicating if OAuth is required and if the user has authorized it + */ +export async function checkOAuthRequirementServer(tool: any): Promise<{ + requiresAuth: boolean + isAuthorized: boolean + provider?: OAuthProvider + requiredScopes?: string[] +}> { + // Skip if no OAuth config or not required + if (!tool.oauth || !tool.oauth.required) { + return { requiresAuth: false, isAuthorized: false } + } + + const provider = tool.oauth.provider + const additionalScopes = tool.oauth.additionalScopes || [] + + // Check if the user has authorized this provider with required scopes + const isAuthorized = await hasAuthorizedProviderServer(provider, additionalScopes) + + return { + requiresAuth: true, + isAuthorized, + provider, + requiredScopes: additionalScopes, + } +} + +/** + * Verify OAuth requirements before executing a tool (server-side) + * Throws an error if OAuth is required but not authorized + * + * @param tool - The tool configuration + * @throws Error with JSON.stringify(OAuthRequiredError) + */ +export async function verifyOAuthBeforeExecutionServer(tool: any): Promise { + const { requiresAuth, isAuthorized, provider, requiredScopes } = + await checkOAuthRequirementServer(tool) + + if (requiresAuth && !isAuthorized && provider) { + // Throw a structured error that can be caught and handled + throw new Error( + JSON.stringify({ + type: 'oauth_required', + provider, + toolId: tool.id, + toolName: tool.name, + requiredScopes, + }) + ) + } +} + +/** + * Get OAuth tokens for a provider if the user has authorized it + * + * @param userId - The user's ID + * @param provider - The OAuth provider to get tokens for + * @returns The OAuth tokens or null if not authorized + */ +export async function getOAuthTokens( + userId: string, + provider: OAuthProvider +): Promise<{ + accessToken: string + refreshToken?: string + expiresAt?: Date +} | null> { + try { + // Query the account table for this user and provider + const accounts = await db + .select() + .from(account) + .where(and(eq(account.userId, userId), eq(account.providerId, provider))) + .limit(1) + + if (!accounts.length || !accounts[0].accessToken) { + return null + } + + const userAccount = accounts[0] + + // Check if the token is expired + if ( + userAccount.accessTokenExpiresAt && + new Date(userAccount.accessTokenExpiresAt) < new Date() + ) { + // In a production app, we would use the refresh token to get a new access token here + // But for simplicity, we'll just return null for expired tokens + console.warn(`Token for ${provider} is expired and needs refresh`) + return null + } + + // Ensure accessToken is not null using the type guard we did earlier + const accessToken = userAccount.accessToken as string + + return { + accessToken, + refreshToken: userAccount.refreshToken || undefined, + expiresAt: userAccount.accessTokenExpiresAt + ? new Date(userAccount.accessTokenExpiresAt) + : undefined, + } + } catch (error) { + console.error('Error getting OAuth tokens:', error) + return null + } +} + +/** + * Get OAuth tokens for a specific tool if the user has authorized it + * + * @param userId - The user's ID + * @param tool - The tool configuration + * @returns The OAuth tokens or null if not required or not authorized + */ +export async function getOAuthTokensForTool( + userId: string, + tool: any +): Promise<{ + accessToken: string + refreshToken?: string + expiresAt?: Date +} | null> { + // Skip if no OAuth config or not required + if (!tool.oauth || !tool.oauth.required) { + return null + } + + // Get tokens for the provider + return getOAuthTokens(userId, tool.oauth.provider) +} + +/** + * Custom hook to check if a user has authorized an OAuth provider + * + * @param provider - The OAuth provider to check + * @param requiredScopes - Optional array of scopes required for the operation + * @returns An object with authorization status and loading state + */ +export function useProviderAuthorization(provider: OAuthProvider, requiredScopes?: string[]) { + const { data: session, isPending } = useSession() + const [isAuthorized, setIsAuthorized] = useState(false) + + useEffect(() => { + if (isPending || !session?.user) { + setIsAuthorized(false) + return + } + + // Check if the user has provider accounts in their session + // This is a client-side check, so it may not be as accurate as the server-side check + const checkAuthorization = async () => { + try { + // We'll use an API endpoint to check authorization status + const response = await fetch( + `/api/auth/oauth/check?provider=${provider}${ + requiredScopes ? `&scopes=${requiredScopes.join(',')}` : '' + }` + ) + + if (response.ok) { + const data = await response.json() + setIsAuthorized(data.isAuthorized || false) + } else { + setIsAuthorized(false) + } + } catch (error) { + console.error('Error checking OAuth authorization:', error) + setIsAuthorized(false) + } + } + + checkAuthorization() + }, [session, isPending, provider, requiredScopes]) + + return { + isAuthorized, + isLoading: isPending, + isLoggedIn: !!session?.user, + } +} + +/** + * Check if a tool requires OAuth and if the user has authorized it + * This function must be used in a client component + * + * @param tool - The tool configuration + * @returns Object indicating if OAuth is required and if the user has the necessary authorization + */ +export function useToolOAuthRequirement(tool: any) { + // Skip if no OAuth config or not required + if (!tool.oauth || !tool.oauth.required) { + return { + requiresAuth: false, + isAuthorized: true, + isLoading: false, + } + } + + const provider = tool.oauth.provider + const additionalScopes = tool.oauth.additionalScopes || [] + + // Use the provider authorization hook + const { isAuthorized, isLoading, isLoggedIn } = useProviderAuthorization( + provider, + additionalScopes + ) + + return { + requiresAuth: true, + isAuthorized: isAuthorized, + isLoading, + isLoggedIn, + provider, + requiredScopes: additionalScopes, + } +} + +/** + * Verify OAuth requirements before executing a tool + * Throws an error if OAuth is required but not authorized + * This function must be used in a client component + * + * @param toolOAuthStatus - The result from useToolOAuthRequirement + * @param tool - The tool configuration + * @throws Error with JSON stringified OAuthRequiredError + */ +export function verifyOAuthBeforeExecution( + toolOAuthStatus: ReturnType, + tool: any +): void { + const { requiresAuth, isAuthorized, isLoading, provider, requiredScopes } = toolOAuthStatus + + // Don't verify while loading + if (isLoading) { + return + } + + if (requiresAuth && !isAuthorized && provider) { + // Throw a structured error that the frontend can catch and handle + throw new Error( + JSON.stringify({ + type: 'oauth_required', + provider, + toolId: tool.id, + toolName: tool.name, + requiredScopes, + } as OAuthRequiredError) + ) + } +} + +/** + * Hook for handling OAuth errors during tool execution + * Provides a modal state and error handler function + */ +export function useOAuthErrorHandler() { + const [oauthModalState, setOAuthModalState] = useState<{ + isOpen: boolean + provider: OAuthProvider + toolName: string + }>({ + isOpen: false, + provider: 'github', + toolName: '', + }) + + /** + * Handle an error that might be an OAuth required error + * Returns true if it was handled as an OAuth error, false otherwise + */ + const handleError = useCallback((error: any): boolean => { + if (!error) return false + + try { + // Try to parse error message as JSON + let errorObj + + if (typeof error === 'string') { + errorObj = JSON.parse(error) + } else if (error instanceof Error && error.message) { + try { + errorObj = JSON.parse(error.message) + } catch { + return false + } + } else { + return false + } + + // Check if it's an OAuth required error + if (errorObj?.type === 'oauth_required') { + setOAuthModalState({ + isOpen: true, + provider: errorObj.provider, + toolName: errorObj.toolName, + }) + return true + } + } catch (parseError) { + // Not a JSON error or not an OAuth error + return false + } + + return false + }, []) + + const closeModal = useCallback(() => { + setOAuthModalState((prev) => ({ ...prev, isOpen: false })) + }, []) + + return { + oauthModalState, + handleOAuthError: handleError, + closeOAuthModal: closeModal, + } +} diff --git a/tools/index.ts b/tools/index.ts index 8afb4f7bc0..8550f7cecc 100644 --- a/tools/index.ts +++ b/tools/index.ts @@ -1,3 +1,4 @@ +import { verifyOAuthBeforeExecutionServer } from '@/lib/oauth' import { useCustomToolsStore } from '@/stores/custom-tools/store' import { useEnvironmentStore } from '@/stores/settings/environment/store' import { visionTool as crewAIVision } from './crewai/vision' @@ -291,6 +292,12 @@ export async function executeTool( throw new Error(`Tool not found: ${toolId}`) } + // Check OAuth requirements before executing the tool + // This will throw an OAuthRequiredError if the tool requires OAuth but the user hasn't authorized it + if (tool.oauth?.required && !isBrowser()) { + await verifyOAuthBeforeExecutionServer(tool) + } + // For custom tools, try direct execution in browser first if available if (toolId.startsWith('custom_') && tool.directExecution) { const directResult = await tool.directExecution(params) diff --git a/tools/types.ts b/tools/types.ts index 6bf2d7a593..758ccfdf4d 100644 --- a/tools/types.ts +++ b/tools/types.ts @@ -1,4 +1,5 @@ export type HttpMethod = 'GET' | 'POST' | 'PUT' | 'DELETE' | 'PATCH' +export type OAuthProvider = 'google' | 'github' | 'twitter' export interface ToolResponse { success: boolean // Whether the tool execution was successful @@ -6,6 +7,12 @@ export interface ToolResponse { error?: string // Error message if success is false } +export interface OAuthConfig { + required: boolean // Whether this tool requires OAuth authentication + provider: OAuthProvider // The provider that needs to be authorized + additionalScopes?: string[] // Additional scopes required for the tool +} + export interface ToolConfig

{ // Basic tool identification id: string @@ -25,6 +32,9 @@ export interface ToolConfig

{ } > + // OAuth configuration for this tool (if it requires authentication) + oauth?: OAuthConfig + // Request configuration request: { url: string | ((params: P) => string)