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

This commit is contained in:
Waleed Latif
2025-03-06 17:22:18 -08:00
parent 4d702cfdbf
commit 3df326f8ba
7 changed files with 686 additions and 2 deletions
+47
View File
@@ -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 }
)
}
}
+128
View File
@@ -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<OAuthProvider, string> = {
github: 'GitHub',
google: 'Google',
twitter: 'X (Twitter)',
}
// Map of provider to icons
const PROVIDER_ICONS: Record<OAuthProvider, React.FC<React.SVGProps<SVGSVGElement>>> = {
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 (
<Dialog open={isOpen} onOpenChange={(open) => !open && onClose()}>
<DialogContent className="sm:max-w-md">
<DialogHeader>
<DialogTitle>Additional Access Required</DialogTitle>
<DialogDescription>
The "{toolName}" tool requires access to your {providerName} account to function
properly.
</DialogDescription>
</DialogHeader>
<div className="flex flex-col gap-4 py-4">
<div className="flex items-center gap-4">
<div className="rounded-full bg-muted p-2">
<ProviderIcon className="h-5 w-5" />
</div>
<div className="flex-1">
<p className="text-sm font-medium">Connect {providerName}</p>
<p className="text-sm text-muted-foreground">Authorize access to use this tool</p>
</div>
</div>
{requiredScopes.length > 0 && (
<details className="text-sm text-muted-foreground rounded-md border p-2">
<summary className="cursor-pointer font-medium">Permissions requested</summary>
<ul className="mt-2 pl-4 list-disc space-y-1">
{requiredScopes.map((scope) => (
<li key={scope}>{scope}</li>
))}
</ul>
</details>
)}
</div>
<DialogFooter className="flex space-x-2 sm:justify-end">
<Button variant="outline" onClick={onClose}>
Cancel
</Button>
<Button type="button" onClick={handleAuth}>
Connect {providerName}
</Button>
</DialogFooter>
</DialogContent>
</Dialog>
)
}
+4 -1
View File
@@ -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
+91 -1
View File
@@ -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',
+399
View File
@@ -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<boolean> {
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<void> {
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<boolean>(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<typeof useToolOAuthRequirement>,
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,
}
}
+7
View File
@@ -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)
+10
View File
@@ -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<P = any, R extends ToolResponse = ToolResponse> {
// Basic tool identification
id: string
@@ -25,6 +32,9 @@ export interface ToolConfig<P = any, R extends ToolResponse = ToolResponse> {
}
>
// OAuth configuration for this tool (if it requires authentication)
oauth?: OAuthConfig
// Request configuration
request: {
url: string | ((params: P) => string)