feat(stripe): added stripe integration, keys for anthropic + openai models (#300)

This commit is contained in:
Waleed Latif
2025-04-24 20:21:19 -07:00
committed by GitHub
parent fb0d3d5b50
commit 53641868d4
44 changed files with 9511 additions and 1856 deletions
+7
View File
@@ -40,6 +40,13 @@ jobs:
- name: Build application
working-directory: ./sim
env:
NODE_OPTIONS: "--no-warnings"
NEXT_PUBLIC_APP_URL: "https://www.simstudio.ai"
STRIPE_SECRET_KEY: "dummy_key_for_ci_only"
STRIPE_WEBHOOK_SECRET: "dummy_secret_for_ci_only"
RESEND_API_KEY: "dummy_key_for_ci_only"
AWS_REGION: "us-west-2"
run: npm run build
- name: Upload coverage to Codecov
@@ -1,5 +1,7 @@
'use server'
import { isProd } from '@/lib/environment'
export async function getOAuthProviderStatus() {
const githubAvailable = !!(
process.env.GITHUB_CLIENT_ID &&
@@ -15,7 +17,5 @@ export async function getOAuthProviderStatus() {
process.env.GOOGLE_CLIENT_SECRET !== 'placeholder'
)
const isProduction = process.env.NODE_ENV === 'production'
return { githubAvailable, googleAvailable, isProduction }
return { githubAvailable, googleAvailable, isProduction: isProd }
}
+3 -3
View File
@@ -1,20 +1,20 @@
import { isProd } from '@/lib/environment'
import { VerifyContent } from './verify-content'
export default function VerifyPage() {
const protocol = process.env.NODE_ENV === 'development' ? 'http' : 'https'
const protocol = isProd ? 'https' : 'http'
const appUrl = process.env.NEXT_PUBLIC_APP_URL || 'localhost:3000'
const baseUrl = `${protocol}://${appUrl}`
const hasResendKey = Boolean(
process.env.RESEND_API_KEY && process.env.RESEND_API_KEY !== 'placeholder'
)
const isProduction = process.env.NODE_ENV === 'production'
return (
<main className="flex min-h-screen flex-col items-center justify-center bg-gray-50">
<div className="sm:mx-auto sm:w-full sm:max-w-md">
<h1 className="text-2xl font-bold text-center mb-8">Sim Studio</h1>
<VerifyContent hasResendKey={hasResendKey} baseUrl={baseUrl} isProduction={isProduction} />
<VerifyContent hasResendKey={hasResendKey} baseUrl={baseUrl} isProduction={isProd} />
</div>
</main>
)
-1
View File
@@ -49,7 +49,6 @@ describe('File Delete API Route', () => {
S3_CONFIG: {
bucket: 'test-bucket',
region: 'test-region',
baseUrl: 'https://test-bucket.s3.test-region.amazonaws.com',
},
}))
+3 -3
View File
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server'
import * as binExt from 'binary-extensions'
import binaryExtensionsList from 'binary-extensions'
import { Buffer } from 'buffer'
import { createHash } from 'crypto'
import fsPromises, { readFile, unlink, writeFile } from 'fs/promises'
@@ -480,7 +480,7 @@ function handleGenericBuffer(
extension: string,
fileType?: string
): ParseResult {
const isBinary = binExt.includes(extension)
const isBinary = binaryExtensionsList.includes(extension)
const content = isBinary
? `[Binary ${extension.toUpperCase()} file - ${fileBuffer.length} bytes]`
: fileBuffer.toString('utf-8')
@@ -686,7 +686,7 @@ async function handleGenericFile(
const fileSize = fileBuffer.length
// Determine if file should be treated as binary
const isBinary = binExt.includes(extension)
const isBinary = binaryExtensionsList.includes(extension)
// Parse content based on binary status
let content: string
@@ -50,7 +50,6 @@ describe('File Serve API Route', () => {
S3_CONFIG: {
bucket: 'test-bucket',
region: 'test-region',
baseUrl: 'https://test-bucket.s3.test-region.amazonaws.com',
},
}))
@@ -73,7 +72,7 @@ describe('File Serve API Route', () => {
const { GET } = await import('./route')
// Call the handler
const response = await GET(req, { params })
const response = await GET(req, { params: Promise.resolve(params) })
// Verify response
expect(response.status).toBe(200)
@@ -101,7 +100,7 @@ describe('File Serve API Route', () => {
const { GET } = await import('./route')
// Call the handler
const response = await GET(req, { params })
const response = await GET(req, { params: Promise.resolve(params) })
// Verify file was read with correct path
expect(mockReadFile).toHaveBeenCalledWith('/test/uploads/nested/path/file.txt')
@@ -124,7 +123,7 @@ describe('File Serve API Route', () => {
const { GET } = await import('./route')
// Call the handler
const response = await GET(req, { params })
const response = await GET(req, { params: Promise.resolve(params) })
// Verify redirect to presigned URL
expect(response.status).toBe(307) // Temporary redirect
@@ -154,7 +153,7 @@ describe('File Serve API Route', () => {
const { GET } = await import('./route')
// Call the handler
const response = await GET(req, { params })
const response = await GET(req, { params: Promise.resolve(params) })
// Verify response falls back to downloading and proxying the file
expect(response.status).toBe(200)
@@ -176,7 +175,7 @@ describe('File Serve API Route', () => {
const { GET } = await import('./route')
// Call the handler
const response = await GET(req, { params })
const response = await GET(req, { params: Promise.resolve(params) })
// Verify 404 response
expect(response.status).toBe(404)
@@ -233,7 +232,6 @@ describe('File Serve API Route', () => {
S3_CONFIG: {
bucket: 'test-bucket',
region: 'test-region',
baseUrl: 'https://test-bucket.s3.test-region.amazonaws.com',
},
}))
@@ -265,7 +263,7 @@ describe('File Serve API Route', () => {
const { GET } = await import('./route')
// Call the handler
const response = await GET(req, { params })
const response = await GET(req, { params: Promise.resolve(params) })
// Verify correct content type
expect(response.headers.get('Content-Type')).toBe(test.contentType)
+2 -3
View File
@@ -1,10 +1,9 @@
import { NextRequest, NextResponse } from 'next/server'
import { readFile } from 'fs/promises'
import { join } from 'path'
import { createLogger } from '@/lib/logs/console-logger'
import { downloadFromS3, getPresignedUrl } from '@/lib/uploads/s3-client'
import { UPLOAD_DIR, USE_S3_STORAGE } from '@/lib/uploads/setup'
// Import to ensure the uploads directory is created
import { USE_S3_STORAGE } from '@/lib/uploads/setup'
import '@/lib/uploads/setup.server'
import {
createErrorResponse,
-1
View File
@@ -74,7 +74,6 @@ describe('File Upload API Route', () => {
S3_CONFIG: {
bucket: 'test-bucket',
region: 'test-region',
baseUrl: 'https://test-bucket.s3.test-region.amazonaws.com',
},
}))
-56
View File
@@ -1,56 +0,0 @@
import { NextRequest, NextResponse } from 'next/server';
/**
* Direct HTTP request handler that fetches external URLs server-side
* This avoids CORS and other browser restrictions
*/
export async function GET(request: NextRequest) {
// Get URL from query parameter
const searchParams = request.nextUrl.searchParams;
const url = searchParams.get('url');
if (!url) {
return NextResponse.json(
{ error: 'Missing URL parameter' },
{ status: 400 }
);
}
try {
// Direct fetch from server side
const response = await fetch(url, {
method: 'GET',
headers: {
'User-Agent': 'Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/135.0.0.0 Safari/537.36',
'Accept': '*/*',
'Accept-Encoding': 'gzip, deflate, br',
'Cache-Control': 'no-cache',
'Connection': 'keep-alive',
},
});
// Get the response data
const contentType = response.headers.get('content-type') || '';
let data;
if (contentType.includes('application/json')) {
data = await response.json();
} else {
data = await response.text();
}
// Return full response information
return NextResponse.json({
success: response.ok,
status: response.status,
statusText: response.statusText,
headers: Object.fromEntries(response.headers.entries()),
data
});
} catch (error: any) {
return NextResponse.json(
{ error: error.message || 'Failed to fetch URL' },
{ status: 500 }
);
}
}
+34
View File
@@ -21,6 +21,7 @@ import { db } from '@/db'
import { environment, userStats, workflow, workflowSchedule } from '@/db/schema'
import { Executor } from '@/executor'
import { Serializer } from '@/serializer'
import { checkServerSideUsageLimits } from '@/lib/usage-monitor'
// Add dynamic export to prevent caching
export const dynamic = 'force-dynamic'
@@ -109,6 +110,39 @@ export async function GET(req: NextRequest) {
runningExecutions.delete(schedule.workflowId)
continue
}
// Check if the user has exceeded their usage limits
const usageCheck = await checkServerSideUsageLimits(workflowRecord.userId)
if (usageCheck.isExceeded) {
logger.warn(`[${requestId}] User ${workflowRecord.userId} has exceeded usage limits. Skipping scheduled execution.`, {
currentUsage: usageCheck.currentUsage,
limit: usageCheck.limit,
workflowId: schedule.workflowId
})
// Log an execution error for the user to see why their schedule was skipped
await persistExecutionError(
schedule.workflowId,
executionId,
new Error(usageCheck.message || 'Usage limit exceeded. Please upgrade your plan to continue running scheduled workflows.'),
'schedule'
)
// Update the next run time to avoid constant retries
const retryDelay = 24 * 60 * 60 * 1000 // 24 hour delay for exceeded limits
const nextRetryAt = new Date(now.getTime() + retryDelay)
await db
.update(workflowSchedule)
.set({
updatedAt: now,
nextRunAt: nextRetryAt,
})
.where(eq(workflowSchedule.id, schedule.id))
runningExecutions.delete(schedule.workflowId)
continue
}
// The state in the database is exactly what we store in localStorage
const state = workflowRecord.state as WorkflowState
+32
View File
@@ -0,0 +1,32 @@
import { NextRequest, NextResponse } from 'next/server'
import { isProPlan } from '@/lib/subscription'
import { getSession } from '@/lib/auth'
import { createLogger } from '@/lib/logs/console-logger'
const logger = createLogger('UserSubscriptionAPI')
export async function GET(request: NextRequest) {
try {
// Get the authenticated user
const session = await getSession()
if (!session?.user?.id) {
logger.warn('Unauthorized subscription access attempt')
return NextResponse.json(
{ error: 'Unauthorized' },
{ status: 401 }
)
}
// Check if the user is on the Pro plan
const isPro = await isProPlan(session.user.id)
return NextResponse.json({ isPro })
} catch (error) {
logger.error('Error checking subscription status:', error)
return NextResponse.json(
{ error: 'Failed to check subscription status' },
{ status: 500 }
)
}
}
+40
View File
@@ -0,0 +1,40 @@
import { NextRequest, NextResponse } from 'next/server'
import { getSession } from '@/lib/auth'
import { createLogger } from '@/lib/logs/console-logger'
import { checkUsageStatus } from '@/lib/usage-monitor'
const logger = createLogger('UserUsageAPI')
export async function GET(request: NextRequest) {
try {
// Get the authenticated user
const session = await getSession()
if (!session?.user?.id) {
logger.warn('Unauthorized usage data access attempt')
return NextResponse.json(
{ error: 'Unauthorized' },
{ status: 401 }
)
}
// Get usage data using our monitor utility
const usageData = await checkUsageStatus(session.user.id)
// Set appropriate caching headers
const response = NextResponse.json(usageData)
// Cache for 5 minutes, private (user-specific data), must revalidate
response.headers.set('Cache-Control', 'private, max-age=300, must-revalidate')
// Add date header for age calculation
response.headers.set('Date', new Date().toUTCString())
return response
} catch (error) {
logger.error('Error checking usage data:', error)
return NextResponse.json(
{ error: 'Failed to check usage data' },
{ status: 500 }
)
}
}
+65 -45
View File
@@ -13,6 +13,7 @@ import {
processWebhook,
fetchAndProcessAirtablePayloads
} from '@/lib/webhooks/utils'
import { checkServerSideUsageLimits } from '@/lib/usage-monitor'
const logger = createLogger('WebhookTriggerAPI')
@@ -21,7 +22,7 @@ export const dynamic = 'force-dynamic' // Ensure dynamic rendering
export const maxDuration = 300 // 5 minutes max execution time
// Storage for active processing tasks to prevent garbage collection
const activeProcessingTasks = new Map<string, Promise<any>>();
const activeProcessingTasks = new Map<string, Promise<any>>()
/**
* Webhook Verification Handler (GET)
@@ -183,19 +184,19 @@ export async function POST(
foundWorkflow = webhooks[0].workflow
// Detect provider type
const isAirtableWebhook = foundWebhook.provider === 'airtable';
const isAirtableWebhook = foundWebhook.provider === 'airtable'
// Handle Slack challenge verification (must be done before timeout)
const slackChallengeResponse = body?.type === 'url_verification' ? handleSlackChallenge(body) : null;
const slackChallengeResponse = body?.type === 'url_verification' ? handleSlackChallenge(body) : null
if (slackChallengeResponse) {
logger.info(`[${requestId}] Responding to Slack URL verification challenge`);
return slackChallengeResponse;
logger.info(`[${requestId}] Responding to Slack URL verification challenge`)
return slackChallengeResponse
}
// Skip processing if another instance is already handling this request
if (!hasExecutionLock) {
logger.info(`[${requestId}] Skipping execution as lock was not acquired`);
return new NextResponse('Request is being processed by another instance', { status: 200 });
logger.info(`[${requestId}] Skipping execution as lock was not acquired`)
return new NextResponse('Request is being processed by another instance', { status: 200 })
}
// --- PHASE 5: Provider-specific processing ---
@@ -203,13 +204,13 @@ export async function POST(
// For Airtable: Process synchronously without timeouts
if (isAirtableWebhook) {
try {
logger.info(`[${requestId}] Airtable webhook ping received for webhook: ${foundWebhook.id}`);
logger.info(`[${requestId}] Airtable webhook ping received for webhook: ${foundWebhook.id}`)
// Handle Airtable deduplication
const notificationId = body.notificationId || null;
const notificationId = body.notificationId || null
if (notificationId) {
try {
const processedKey = `airtable-webhook-${foundWebhook.id}-${notificationId}`;
const processedKey = `airtable-webhook-${foundWebhook.id}-${notificationId}`
// Check if notification was already processed
const alreadyProcessed = await db
@@ -221,20 +222,20 @@ export async function POST(
sql`(webhook.provider_config->>'processedNotifications')::jsonb ? ${processedKey}`
)
)
.limit(1);
.limit(1)
if (alreadyProcessed.length > 0) {
logger.info(`[${requestId}] Duplicate Airtable notification detected: ${notificationId}`);
return new NextResponse('Notification already processed', { status: 200 });
logger.info(`[${requestId}] Duplicate Airtable notification detected: ${notificationId}`)
return new NextResponse('Notification already processed', { status: 200 })
}
// Store notification ID for deduplication
const providerConfig = foundWebhook.providerConfig || {};
const processedNotifications = providerConfig.processedNotifications || [];
processedNotifications.push(processedKey);
const providerConfig = foundWebhook.providerConfig || {}
const processedNotifications = providerConfig.processedNotifications || []
processedNotifications.push(processedKey)
// Keep only the last 100 notifications to prevent unlimited growth
const limitedNotifications = processedNotifications.slice(-100);
const limitedNotifications = processedNotifications.slice(-100)
// Update the webhook record
await db
@@ -246,74 +247,93 @@ export async function POST(
},
updatedAt: new Date(),
})
.where(eq(webhook.id, foundWebhook.id));
.where(eq(webhook.id, foundWebhook.id))
} catch (error) {
logger.warn(`[${requestId}] Airtable deduplication check failed, continuing`, {
error: error instanceof Error ? error.message : String(error)
});
})
}
}
// Process Airtable payloads synchronously
try {
logger.info(`[${requestId}] Starting Airtable payload processing`);
await fetchAndProcessAirtablePayloads(foundWebhook, foundWorkflow, requestId);
return new NextResponse('Airtable ping processed successfully', { status: 200 });
logger.info(`[${requestId}] Starting Airtable payload processing`)
await fetchAndProcessAirtablePayloads(foundWebhook, foundWorkflow, requestId)
return new NextResponse('Airtable ping processed successfully', { status: 200 })
} catch (error: any) {
logger.error(`[${requestId}] Error during Airtable processing`, {
error: error.message
});
})
return new NextResponse(`Error processing Airtable webhook: ${error.message}`, {
status: 500,
});
})
}
} catch (error: any) {
logger.error(`[${requestId}] Error in Airtable processing`, error);
return new NextResponse(`Internal server error: ${error.message}`, { status: 500 });
logger.error(`[${requestId}] Error in Airtable processing`, error)
return new NextResponse(`Internal server error: ${error.message}`, { status: 500 })
}
}
// --- For all other webhook types: Use async processing with timeout ---
// Create timeout promise for fast initial response (2.5 seconds)
const timeoutDuration = 25000;
const timeoutDuration = 25000
const timeoutPromise = new Promise<NextResponse>((resolve) => {
setTimeout(() => {
logger.info(`[${requestId}] Fast response timeout activated`);
resolve(new NextResponse('Request received', { status: 200 }));
}, timeoutDuration);
});
logger.info(`[${requestId}] Fast response timeout activated`)
resolve(new NextResponse('Request received', { status: 200 }))
}, timeoutDuration)
})
// Create the processing promise for asynchronous execution
const processingPromise = (async () => {
try {
// Provider-specific deduplication
if (foundWebhook.provider === 'whatsapp') {
const data = body?.entry?.[0]?.changes?.[0]?.value;
const messages = data?.messages || [];
const data = body?.entry?.[0]?.changes?.[0]?.value
const messages = data?.messages || []
const whatsappDuplicateResponse = await processWhatsAppDeduplication(requestId, messages);
const whatsappDuplicateResponse = await processWhatsAppDeduplication(requestId, messages)
if (whatsappDuplicateResponse) {
return whatsappDuplicateResponse;
return whatsappDuplicateResponse
}
} else if (foundWebhook.provider !== 'slack') {
const genericDuplicateResponse = await processGenericDeduplication(requestId, path, body);
const genericDuplicateResponse = await processGenericDeduplication(requestId, path, body)
if (genericDuplicateResponse) {
return genericDuplicateResponse;
return genericDuplicateResponse
}
}
// Execute workflow for the webhook event
logger.info(`[${requestId}] Executing workflow for ${foundWebhook.provider} webhook`);
// Check if the user has exceeded their usage limits
const usageCheck = await checkServerSideUsageLimits(foundWorkflow.userId)
if (usageCheck.isExceeded) {
logger.warn(`[${requestId}] User ${foundWorkflow.userId} has exceeded usage limits. Skipping webhook execution.`, {
currentUsage: usageCheck.currentUsage,
limit: usageCheck.limit,
workflowId: foundWorkflow.id
})
// Return a successful response to avoid webhook retries, but don't execute the workflow
return new NextResponse(JSON.stringify({
status: 'error',
message: usageCheck.message || 'Usage limit exceeded. Please upgrade your plan to continue using webhooks.'
}), {
status: 200, // Use 200 to prevent webhook provider retries
headers: { 'Content-Type': 'application/json' }
})
}
const executionId = uuidv4();
return await processWebhook(foundWebhook, foundWorkflow, body, request, executionId, requestId);
// Execute workflow for the webhook event
logger.info(`[${requestId}] Executing workflow for ${foundWebhook.provider} webhook`)
const executionId = uuidv4()
return await processWebhook(foundWebhook, foundWorkflow, body, request, executionId, requestId)
} catch (error: any) {
logger.error(`[${requestId}] Error processing webhook:`, error);
return new NextResponse(`Internal server error: ${error.message}`, { status: 500 });
logger.error(`[${requestId}] Error processing webhook:`, error)
return new NextResponse(`Internal server error: ${error.message}`, { status: 500 })
}
})();
})()
// Race processing against timeout to ensure fast response
return Promise.race([timeoutPromise, processingPromise]);
return Promise.race([timeoutPromise, processingPromise])
}
@@ -88,7 +88,7 @@ describe('Workflow Execution API Route', () => {
decryptSecret: vi.fn().mockResolvedValue({
decrypted: 'decrypted-secret-value',
}),
isHostedVersion: vi.fn().mockReturnValue(false),
isHosted: vi.fn().mockReturnValue(false),
getRotatingApiKey: vi.fn().mockReturnValue('rotated-api-key'),
}))
@@ -15,6 +15,7 @@ import { Executor } from '@/executor'
import { Serializer } from '@/serializer'
import { validateWorkflowAccess } from '../../middleware'
import { createErrorResponse, createSuccessResponse } from '../../utils'
import { checkServerSideUsageLimits } from '@/lib/usage-monitor'
const logger = createLogger('WorkflowExecuteAPI')
@@ -27,6 +28,17 @@ const EnvVarsSchema = z.record(z.string())
// Keep track of running executions to prevent overlap
const runningExecutions = new Set<string>()
// Custom error class for usage limit exceeded
class UsageLimitError extends Error {
statusCode: number
constructor(message: string) {
super(message)
this.name = 'UsageLimitError'
this.statusCode = 402 // Payment Required status code
}
}
async function executeWorkflow(workflow: any, requestId: string, input?: any) {
const workflowId = workflow.id
const executionId = uuidv4()
@@ -37,6 +49,16 @@ async function executeWorkflow(workflow: any, requestId: string, input?: any) {
throw new Error('Workflow is already running')
}
// Check if the user has exceeded their usage limits
const usageCheck = await checkServerSideUsageLimits(workflow.userId)
if (usageCheck.isExceeded) {
logger.warn(`[${requestId}] User ${workflow.userId} has exceeded usage limits`, {
currentUsage: usageCheck.currentUsage,
limit: usageCheck.limit
})
throw new UsageLimitError(usageCheck.message || 'Usage limit exceeded. Please upgrade your plan to continue.')
}
// Log input to help debug
logger.info(
`[${requestId}] Executing workflow with input:`,
@@ -273,6 +295,16 @@ export async function GET(request: NextRequest, { params }: { params: Promise<{
return createSuccessResponse(result)
} catch (error: any) {
logger.error(`[${requestId}] Error executing workflow: ${id}`, error)
// Check if this is a usage limit error
if (error instanceof UsageLimitError) {
return createErrorResponse(
error.message,
error.statusCode,
'USAGE_LIMIT_EXCEEDED'
)
}
return createErrorResponse(
error.message || 'Failed to execute workflow',
500,
@@ -320,6 +352,16 @@ export async function POST(request: NextRequest, { params }: { params: Promise<{
return createSuccessResponse(result)
} catch (error: any) {
logger.error(`[${requestId}] Error executing workflow: ${id}`, error)
// Check if this is a usage limit error
if (error instanceof UsageLimitError) {
return createErrorResponse(
error.message,
error.statusCode,
'USAGE_LIMIT_EXCEEDED'
)
}
return createErrorResponse(
error.message || 'Failed to execute workflow',
500,
@@ -16,6 +16,7 @@ import {
Store,
Trash2,
X,
CreditCard,
} from 'lucide-react'
import {
AlertDialog,
@@ -51,9 +52,18 @@ import { DeploymentControls } from './components/deployment-controls/deployment-
import { HistoryDropdownItem } from './components/history-dropdown-item/history-dropdown-item'
import { MarketplaceModal } from './components/marketplace-modal/marketplace-modal'
import { NotificationDropdownItem } from './components/notification-dropdown-item/notification-dropdown-item'
import { useSession } from '@/lib/auth-client'
const logger = createLogger('ControlBar')
// Cache for usage data to prevent excessive API calls
let usageDataCache = {
data: null,
timestamp: 0,
// Cache expires after 1 minute
expirationMs: 60 * 1000
}
// Predefined run count options
const RUN_COUNT_OPTIONS = [1, 5, 10, 25, 50, 100]
@@ -63,6 +73,7 @@ const RUN_COUNT_OPTIONS = [1, 5, 10, 25, 50, 100]
*/
export function ControlBar() {
const router = useRouter()
const { data: session } = useSession()
// Store hooks
const {
@@ -111,6 +122,16 @@ export function ControlBar() {
const [isCancelling, setIsCancelling] = useState(false)
const cancelFlagRef = useRef(false)
// Usage limit state
const [usageExceeded, setUsageExceeded] = useState(false)
const [usageData, setUsageData] = useState<{
percentUsed: number
isWarning: boolean
isExceeded: boolean
currentUsage: number
limit: number
} | null>(null)
// Register keyboard shortcut for running workflow
useKeyboardShortcuts(
() => {
@@ -337,6 +358,56 @@ export function ControlBar() {
}
}, [needsRedeployment, activeWorkflowId, notifications, removeNotification, addNotification])
// Check usage limits when component mounts and when user executes a workflow
useEffect(() => {
if (session?.user?.id) {
checkUserUsage(session.user.id).then(usage => {
if (usage) {
setUsageExceeded(usage.isExceeded)
setUsageData(usage)
}
})
}
}, [session?.user?.id, completedRuns])
/**
* Check user usage data with caching to prevent excessive API calls
* @param userId User ID to check usage for
* @param forceRefresh Whether to force a fresh API call ignoring cache
* @returns Usage data or null if error
*/
async function checkUserUsage(userId: string, forceRefresh = false): Promise<any | null> {
const now = Date.now()
const cacheAge = now - usageDataCache.timestamp
// Use cache if available and not expired
if (!forceRefresh && usageDataCache.data && cacheAge < usageDataCache.expirationMs) {
logger.info('Using cached usage data', { cacheAge: `${Math.round(cacheAge/1000)}s` })
return usageDataCache.data
}
try {
const response = await fetch('/api/user/usage')
if (!response.ok) {
throw new Error('Failed to fetch usage data')
}
const usage = await response.json()
// Update cache
usageDataCache = {
data: usage,
timestamp: now,
expirationMs: usageDataCache.expirationMs
}
return usage
} catch (error) {
logger.error('Error checking usage limits:', { error })
return null
}
}
/**
* Workflow name handlers
*/
@@ -407,6 +478,12 @@ export function ControlBar() {
*/
const handleMultipleRuns = async () => {
if (isExecuting || isMultiRunning || runCount <= 0) return
// Check if usage is exceeded before allowing execution
if (usageExceeded) {
openSubscriptionSettings()
return
}
// Reset state and ref for a new batch of runs
setCompletedRuns(0)
@@ -417,6 +494,8 @@ export function ControlBar() {
let workflowError = null
let wasCancelled = false
let runCounter = 0
let shouldCheckUsage = false
try {
// Run the workflow multiple times sequentially
@@ -430,14 +509,38 @@ export function ControlBar() {
// Run the workflow and immediately increment counter for visual feedback
await handleRunWorkflow()
setCompletedRuns(i + 1)
runCounter = i + 1
setCompletedRuns(runCounter)
// Only check usage periodically to avoid excessive API calls
// Check on first run, every 5 runs, and on last run
shouldCheckUsage = i === 0 || (i + 1) % 5 === 0 || i === runCount - 1
// Check usage if needed
if (shouldCheckUsage && session?.user?.id) {
const usage = await checkUserUsage(session.user.id, i === 0)
if (usage?.isExceeded) {
setUsageExceeded(true)
setUsageData(usage)
// Stop execution if we've exceeded the limit during this batch
if (i < runCount - 1) {
addNotification(
'info',
`Usage limit reached after ${runCounter} runs. Execution stopped.`,
activeWorkflowId
)
break
}
}
}
}
// Update workflow stats only if the run wasn't cancelled and completed normally
if (!wasCancelled && activeWorkflowId) {
try {
// Don't block UI on stats update
fetch(`/api/workflows/${activeWorkflowId}/stats?runs=${runCount}`, {
fetch(`/api/workflows/${activeWorkflowId}/stats?runs=${runCounter}`, {
method: 'POST',
}).catch((error) => {
logger.error(`Failed to update workflow stats: ${error.message}`)
@@ -845,6 +948,16 @@ export function ControlBar() {
)
}
// Helper function to open subscription settings
const openSubscriptionSettings = () => {
// Dispatch custom event to open settings modal with subscription tab
if (typeof window !== 'undefined') {
window.dispatchEvent(new CustomEvent('open-settings', {
detail: { tab: 'subscription' }
}))
}
}
/**
* Render run workflow button with multi-run dropdown and cancel button
*/
@@ -888,7 +1001,7 @@ export function ControlBar() {
? 'rounded py-2 px-4 h-10'
: 'rounded-r-none border-r border-r-[#6420cc] py-2 px-4 h-10'
)}
onClick={isDebugModeEnabled ? handleRunWorkflow : handleMultipleRuns}
onClick={usageExceeded ? openSubscriptionSettings : (isDebugModeEnabled ? handleRunWorkflow : handleMultipleRuns)}
disabled={isExecuting || isMultiRunning || isCancelling}
>
{isCancelling ? (
@@ -914,14 +1027,26 @@ export function ControlBar() {
</Button>
</TooltipTrigger>
<TooltipContent>
{isDebugModeEnabled
? 'Debug Workflow'
: runCount === 1
? 'Run Workflow'
: `Run Workflow ${runCount} times`}
<span className="text-xs text-muted-foreground ml-1">
{getKeyboardShortcutText('Enter', true)}
</span>
{usageExceeded ? (
<div className="text-center">
<p className="font-medium text-destructive">Usage Limit Exceeded</p>
<p className="text-xs">
You've used {usageData?.currentUsage.toFixed(2)}$ of {usageData?.limit}$.
Upgrade your plan to continue.
</p>
</div>
) : (
<>
{isDebugModeEnabled
? 'Debug Workflow'
: runCount === 1
? 'Run Workflow'
: `Run Workflow ${runCount} times`}
<span className="text-xs text-muted-foreground ml-1">
{getKeyboardShortcutText('Enter', true)}
</span>
</>
)}
</TooltipContent>
</Tooltip>
@@ -1,14 +1,22 @@
import { Key, KeyRound, KeySquare, Settings, UserCircle } from 'lucide-react'
import { Key, KeyRound, KeySquare, Settings, UserCircle, CreditCard } from 'lucide-react'
import { cn } from '@/lib/utils'
import { isDev } from '@/lib/environment'
interface SettingsNavigationProps {
activeSection: string
onSectionChange: (
section: 'general' | 'environment' | 'account' | 'credentials' | 'apikeys'
section: 'general' | 'environment' | 'account' | 'credentials' | 'apikeys' | 'subscription'
) => void
}
const navigationItems = [
type NavigationItem = {
id: 'general' | 'environment' | 'account' | 'credentials' | 'apikeys' | 'subscription'
label: string
icon: React.ComponentType<{ className?: string }>
hideInDev?: boolean
}
const allNavigationItems: NavigationItem[] = [
{
id: 'general',
label: 'General',
@@ -34,9 +42,22 @@ const navigationItems = [
label: 'API Keys',
icon: KeySquare,
},
] as const
{
id: 'subscription',
label: 'Subscription',
icon: CreditCard,
hideInDev: true,
},
]
export function SettingsNavigation({ activeSection, onSectionChange }: SettingsNavigationProps) {
const navigationItems = allNavigationItems.filter(item => {
if (item.hideInDev && isDev) {
return false
}
return true
})
return (
<div className="py-4">
{navigationItems.map((item) => (
@@ -0,0 +1,321 @@
import { useState, useEffect } from 'react'
import { client, useSession } from '@/lib/auth-client'
import { Alert, AlertDescription, AlertTitle } from '@/components/ui/alert'
import { AlertCircle } from 'lucide-react'
import { Button } from '@/components/ui/button'
import { LoadingAgent } from '@/components/ui/loading-agent'
import { Progress } from '@/components/ui/progress'
interface SubscriptionProps {
onOpenChange: (open: boolean) => void
}
export function Subscription({ onOpenChange }: SubscriptionProps) {
const { data: session } = useSession()
const [isPro, setIsPro] = useState<boolean>(false)
const [usageData, setUsageData] = useState<{
percentUsed: number;
isWarning: boolean;
isExceeded: boolean;
currentUsage: number;
limit: number;
}>({
percentUsed: 0,
isWarning: false,
isExceeded: false,
currentUsage: 0,
limit: 0
})
const [loading, setLoading] = useState<boolean>(true)
const [subscriptionData, setSubscriptionData] = useState<any>(null)
const [isCanceling, setIsCanceling] = useState<boolean>(false)
const [error, setError] = useState<string | null>(null)
useEffect(() => {
async function checkSubscriptionStatus() {
if (session?.user?.id) {
try {
setLoading(true)
setError(null)
// Fetch subscription status from API
const proStatusResponse = await fetch('/api/user/subscription')
if (!proStatusResponse.ok) {
throw new Error('Failed to fetch subscription status')
}
const proStatusData = await proStatusResponse.json()
setIsPro(proStatusData.isPro)
// Fetch usage data from API
const usageResponse = await fetch('/api/user/usage')
if (!usageResponse.ok) {
throw new Error('Failed to fetch usage data')
}
const usageData = await usageResponse.json()
setUsageData(usageData)
// Fetch detailed subscription data
const { data, error: subError } = await client.subscription.list()
if (subError) {
console.error('Error fetching subscription details', subError)
// Continue with basic subscription info we already have
} else {
// Find active subscription
const activeSubscription = data?.find(
sub => sub.status === 'active'
)
setSubscriptionData(activeSubscription)
}
} catch (error) {
console.error('Error checking subscription status:', error)
} finally {
setLoading(false)
}
}
}
checkSubscriptionStatus()
}, [session?.user?.id])
const handleUpgrade = async () => {
if (!session?.user) {
setError('You need to be logged in to upgrade your subscription')
return
}
try {
const { error } = await client.subscription.upgrade({
plan: 'pro',
successUrl: window.location.href,
cancelUrl: window.location.href,
})
if (error) {
setError(error.message || 'There was an error upgrading your subscription')
}
} catch (error: any) {
setError(error.message || 'There was an error upgrading your subscription')
}
}
const handleCancel = async () => {
if (!session?.user) {
setError('You need to be logged in to cancel your subscription')
return
}
setIsCanceling(true)
try {
const { error } = await client.subscription.cancel({
returnUrl: window.location.href,
})
if (error) {
setError(error.message || 'There was an error canceling your subscription')
}
} catch (error: any) {
setError(error.message || 'There was an error canceling your subscription')
} finally {
setIsCanceling(false)
}
}
return (
<div className="p-6 space-y-6">
<h3 className="text-lg font-medium">Subscription Plans</h3>
{error && (
<Alert variant="destructive" className="mb-4">
<AlertCircle className="h-4 w-4" />
<AlertDescription>{error}</AlertDescription>
</Alert>
)}
{(usageData.isWarning || usageData.isExceeded) && !isPro && (
<Alert variant="destructive" className="mb-4">
<AlertCircle className="h-4 w-4" />
<AlertTitle>{usageData.isExceeded ? 'Usage Limit Exceeded' : 'Usage Warning'}</AlertTitle>
<AlertDescription>
You've used {usageData.percentUsed}% of your free tier limit
({usageData.currentUsage.toFixed(2)}$ of {usageData.limit}$).
{usageData.isExceeded
? ' You have exceeded your limit. Upgrade to Pro to continue using all features.'
: ' Upgrade to Pro to avoid any service interruptions.'}
</AlertDescription>
</Alert>
)}
{loading ? (
<div className="flex items-center justify-center py-8">
<LoadingAgent size="sm" />
<span className="ml-2">Loading subscription details...</span>
</div>
) : (
<>
<div className="grid gap-6 md:grid-cols-2">
{/* Free Tier */}
<div className={`border rounded-lg p-4 ${!isPro ? 'border-primary' : ''}`}>
<h4 className="text-md font-semibold">Free Tier</h4>
<p className="text-sm text-muted-foreground mt-1">For individual users and small projects</p>
<ul className="mt-3 space-y-2 text-sm">
<li>• ${!isPro ? 5 : usageData.limit} of inference credits</li>
<li>• Basic features</li>
<li>• No sharing capabilities</li>
</ul>
{!isPro && (
<div className="mt-4 space-y-2">
<div className="flex justify-between text-xs">
<span>Usage</span>
<span>
{usageData.currentUsage.toFixed(2)}$ / {usageData.limit}$
</span>
</div>
<Progress
value={usageData.percentUsed}
className={`h-2 ${
usageData.isExceeded
? 'bg-muted [&>*]:bg-destructive'
: usageData.isWarning
? 'bg-muted [&>*]:bg-amber-500'
: ''
}`}
/>
</div>
)}
<div className="mt-4">
{!isPro ? (
<div className="text-sm bg-secondary/50 text-secondary-foreground py-1 px-2 rounded inline-block">
Current Plan
</div>
) : (
<Button
variant="outline"
size="sm"
onClick={handleCancel}
disabled={isCanceling}
>
{isCanceling && <LoadingAgent size="sm" />}
<span className={isCanceling ? "ml-2" : ""}>Downgrade</span>
</Button>
)}
</div>
</div>
{/* Pro Tier */}
<div className={`border rounded-lg p-4 ${isPro ? 'border-primary' : ''}`}>
<h4 className="text-md font-semibold">Pro Tier</h4>
<p className="text-sm text-muted-foreground mt-1">For professional users and teams</p>
<ul className="mt-3 space-y-2 text-sm">
<li>• ${isPro ? usageData.limit : 20} of inference credits</li>
<li>• All features included</li>
<li>• Workflow sharing capabilities</li>
</ul>
{isPro && (
<div className="mt-4 space-y-2">
<div className="flex justify-between text-xs">
<span>Usage</span>
<span>
{usageData.currentUsage.toFixed(2)}$ / {usageData.limit}$
</span>
</div>
<Progress
value={usageData.percentUsed}
className={`h-2 ${
usageData.isExceeded
? 'bg-muted [&>*]:bg-destructive'
: usageData.isWarning
? 'bg-muted [&>*]:bg-amber-500'
: ''
}`}
/>
</div>
)}
<div className="mt-4">
{isPro ? (
<div className="text-sm bg-secondary/50 text-secondary-foreground py-1 px-2 rounded inline-block">
Current Plan
</div>
) : (
<Button
variant="default"
size="sm"
onClick={handleUpgrade}
>
Upgrade
</Button>
)}
</div>
</div>
{/* Enterprise Tier */}
<div className="border rounded-lg p-4 col-span-full">
<h4 className="text-md font-semibold">Enterprise</h4>
<p className="text-sm text-muted-foreground mt-1">For larger teams and organizations</p>
<ul className="mt-3 space-y-2 text-sm">
<li>• Custom cost limits</li>
<li>• Priority support</li>
<li>• Custom integrations</li>
<li>• Dedicated account manager</li>
</ul>
<div className="mt-4">
<Button
variant="outline"
size="sm"
onClick={() => {
window.open(
'https://calendly.com/emir-simstudio/15min',
'_blank',
'noopener,noreferrer'
)
}}
>
Contact Us
</Button>
</div>
</div>
</div>
{subscriptionData && (
<div className="mt-8 border-t pt-6">
<h4 className="text-md font-medium mb-4">Subscription Details</h4>
<div className="text-sm space-y-2">
<p>
<span className="font-medium">Status:</span>{' '}
<span className="capitalize">{subscriptionData.status}</span>
</p>
{subscriptionData.periodEnd && (
<p>
<span className="font-medium">Next billing date:</span>{' '}
{new Date(subscriptionData.periodEnd).toLocaleDateString()}
</p>
)}
{isPro && (
<div className="mt-4">
<Button
variant="outline"
onClick={handleCancel}
disabled={isCanceling}
>
{isCanceling && <LoadingAgent size="sm" />}
<span className={isCanceling ? "ml-2" : ""}>Manage Subscription</span>
</Button>
</div>
)}
</div>
</div>
)}
</>
)}
</div>
)
}
@@ -5,11 +5,13 @@ import { X } from 'lucide-react'
import { Button } from '@/components/ui/button'
import { Dialog, DialogContent, DialogHeader, DialogTitle } from '@/components/ui/dialog'
import { cn } from '@/lib/utils'
import { client } from '@/lib/auth-client'
import { Account } from './components/account/account'
import { ApiKeys } from './components/api-keys/api-keys'
import { Credentials } from './components/credentials/credentials'
import { EnvironmentVariables } from './components/environment/environment'
import { General } from './components/general/general'
import { Subscription } from './components/subscription/subscription'
import { SettingsNavigation } from './components/settings-navigation/settings-navigation'
interface SettingsModalProps {
@@ -17,7 +19,7 @@ interface SettingsModalProps {
onOpenChange: (open: boolean) => void
}
type SettingsSection = 'general' | 'environment' | 'account' | 'credentials' | 'apikeys'
type SettingsSection = 'general' | 'environment' | 'account' | 'credentials' | 'apikeys' | 'subscription'
export function SettingsModal({ open, onOpenChange }: SettingsModalProps) {
const [activeSection, setActiveSection] = useState<SettingsSection>('general')
@@ -38,6 +40,9 @@ export function SettingsModal({ open, onOpenChange }: SettingsModalProps) {
}
}, [onOpenChange])
// Check if subscriptions are enabled
const isSubscriptionEnabled = !!client.subscription
return (
<Dialog open={open} onOpenChange={onOpenChange}>
<DialogContent className="sm:max-w-[700px] h-[64vh] flex flex-col p-0 gap-0" hideCloseButton>
@@ -79,6 +84,11 @@ export function SettingsModal({ open, onOpenChange }: SettingsModalProps) {
<div className={cn('h-full', activeSection === 'apikeys' ? 'block' : 'hidden')}>
<ApiKeys onOpenChange={onOpenChange} />
</div>
{isSubscriptionEnabled && (
<div className={cn('h-full', activeSection === 'subscription' ? 'block' : 'hidden')}>
<Subscription onOpenChange={onOpenChange} />
</div>
)}
</div>
</div>
</DialogContent>
+14 -5
View File
@@ -1,6 +1,6 @@
import { AgentIcon } from '@/components/icons'
import { createLogger } from '@/lib/logs/console-logger'
import { isHostedVersion } from '@/lib/utils'
import { isHosted } from '@/lib/environment'
import { useOllamaStore } from '@/stores/ollama/store'
import { getAllBlocks } from '@/blocks'
import { MODELS_TEMP_RANGE_0_1, MODELS_TEMP_RANGE_0_2 } from '@/providers/model-capabilities'
@@ -8,7 +8,6 @@ import { getAllModelProviders, getBaseModelProviders } from '@/providers/utils'
import { ToolResponse } from '@/tools/types'
import { BlockConfig } from '../types'
const isHosted = isHostedVersion()
const logger = createLogger('AgentBlock')
interface AgentResponse extends ToolResponse {
@@ -104,12 +103,22 @@ export const AgentBlock: BlockConfig<AgentResponse> = {
placeholder: 'Enter your API key',
password: true,
connectionDroppable: false,
// Hide API key for GPT-4o models when running on hosted version
// Hide API key for all OpenAI and Claude models when running on hosted version
condition: isHosted
? {
field: 'model',
value: 'gpt-4o',
not: true, // Show for all models EXCEPT GPT-4o models
// Include all OpenAI models and Claude models for which we don't show the API key field
value: [
// OpenAI models
'gpt-4o',
'o1', 'o1-mini', 'o1-preview',
'o3', 'o3-preview',
'o4-mini',
// Claude models
'claude-3-5-sonnet-20240620',
'claude-3-7-sonnet-20250219'
],
not: true, // Show for all models EXCEPT those listed
}
: undefined, // Show for all models in non-hosted environments
},
+2 -2
View File
@@ -1,13 +1,13 @@
import { DocumentIcon } from '@/components/icons'
import { createLogger } from '@/lib/logs/console-logger'
import { isProd } from '@/lib/environment'
import { FileParserOutput } from '@/tools/file/types'
import { BlockConfig, SubBlockConfig, SubBlockLayout, SubBlockType } from '../types'
const logger = createLogger('FileBlock')
const isProduction = process.env.NODE_ENV === 'production'
const isS3Enabled = process.env.USE_S3 === 'true'
const shouldEnableURLInput = isProduction || isS3Enabled
const shouldEnableURLInput = isProd || isS3Enabled
// Define sub-blocks conditionally
const inputMethodBlock: SubBlockConfig = {
+2 -2
View File
@@ -1,10 +1,10 @@
import { MistralIcon } from '@/components/icons'
import { isProd } from '@/lib/environment'
import { MistralParserOutput } from '@/tools/mistral/types'
import { BlockConfig, SubBlockConfig, SubBlockLayout, SubBlockType } from '../types'
const isProduction = process.env.NODE_ENV === 'production'
const isS3Enabled = process.env.USE_S3 === 'true'
const shouldEnableFileUpload = isProduction || isS3Enabled
const shouldEnableFileUpload = isProd || isS3Enabled
// Define the input method selector block when needed
const inputMethodBlock: SubBlockConfig = {
+35
View File
@@ -0,0 +1,35 @@
CREATE TABLE "chat" (
"id" text PRIMARY KEY NOT NULL,
"workflow_id" text NOT NULL,
"user_id" text NOT NULL,
"subdomain" text NOT NULL,
"title" text NOT NULL,
"description" text,
"is_active" boolean DEFAULT true NOT NULL,
"customizations" json DEFAULT '{}',
"auth_type" text DEFAULT 'public' NOT NULL,
"password" text,
"allowed_emails" json DEFAULT '[]',
"output_block_id" text,
"output_path" text,
"created_at" timestamp DEFAULT now() NOT NULL,
"updated_at" timestamp DEFAULT now() NOT NULL
);
--> statement-breakpoint
CREATE TABLE "subscription" (
"id" text PRIMARY KEY NOT NULL,
"plan" text NOT NULL,
"reference_id" text NOT NULL,
"stripe_customer_id" text,
"stripe_subscription_id" text,
"status" text,
"period_start" timestamp,
"period_end" timestamp,
"cancel_at_period_end" boolean,
"seats" integer
);
--> statement-breakpoint
ALTER TABLE "user" ADD COLUMN "stripe_customer_id" text;--> statement-breakpoint
ALTER TABLE "chat" ADD CONSTRAINT "chat_workflow_id_workflow_id_fk" FOREIGN KEY ("workflow_id") REFERENCES "public"."workflow"("id") ON DELETE cascade ON UPDATE no action;--> statement-breakpoint
ALTER TABLE "chat" ADD CONSTRAINT "chat_user_id_user_id_fk" FOREIGN KEY ("user_id") REFERENCES "public"."user"("id") ON DELETE cascade ON UPDATE no action;--> statement-breakpoint
CREATE UNIQUE INDEX "subdomain_idx" ON "chat" USING btree ("subdomain");
File diff suppressed because it is too large Load Diff
+7
View File
@@ -211,6 +211,13 @@
"when": 1745211620858,
"tag": "0029_grey_barracuda",
"breakpoints": true
},
{
"idx": 30,
"version": "7",
"when": 1745519847269,
"tag": "0030_happy_joseph",
"breakpoints": true
}
]
}
+50 -2
View File
@@ -17,7 +17,8 @@ export const user = pgTable('user', {
image: text('image'),
createdAt: timestamp('created_at').notNull(),
updatedAt: timestamp('updated_at').notNull(),
})
stripeCustomerId: text('stripe_customer_id')
});
export const session = pgTable('session', {
id: text('id').primaryKey(),
@@ -220,4 +221,51 @@ export const customTools = pgTable('custom_tools', {
code: text('code').notNull(),
createdAt: timestamp('created_at').notNull().defaultNow(),
updatedAt: timestamp('updated_at').notNull().defaultNow(),
})
})
export const subscription = pgTable("subscription", {
id: text('id').primaryKey(),
plan: text('plan').notNull(),
referenceId: text('reference_id').notNull(),
stripeCustomerId: text('stripe_customer_id'),
stripeSubscriptionId: text('stripe_subscription_id'),
status: text('status'),
periodStart: timestamp('period_start'),
periodEnd: timestamp('period_end'),
cancelAtPeriodEnd: boolean('cancel_at_period_end'),
seats: integer('seats')
});
export const chat = pgTable('chat', {
id: text('id').primaryKey(),
workflowId: text('workflow_id')
.notNull()
.references(() => workflow.id, { onDelete: 'cascade' }),
userId: text('user_id')
.notNull()
.references(() => user.id, { onDelete: 'cascade' }),
subdomain: text('subdomain').notNull(),
title: text('title').notNull(),
description: text('description'),
isActive: boolean('is_active').notNull().default(true),
customizations: json('customizations').default('{}'), // For UI customization options
// Authentication options
authType: text('auth_type').notNull().default('public'), // 'public', 'password', 'email'
password: text('password'), // Stored hashed, populated when authType is 'password'
allowedEmails: json('allowed_emails').default('[]'), // Array of allowed emails or domains when authType is 'email'
// Output configuration
outputBlockId: text('output_block_id'), // Stores the selected output block ID
outputPath: text('output_path'), // Stores the output path within the block
createdAt: timestamp('created_at').notNull().defaultNow(),
updatedAt: timestamp('updated_at').notNull().defaultNow(),
},
(table) => {
return {
// Ensure subdomains are unique
subdomainIdx: uniqueIndex('subdomain_idx').on(table.subdomain),
}
}
)
@@ -26,7 +26,7 @@ vi.mock('@/tools/utils', () => ({
// Utils
vi.mock('@/lib/utils', () => ({
isHostedVersion: vi.fn().mockReturnValue(false),
isHosted: vi.fn().mockReturnValue(false),
getRotatingApiKey: vi.fn(),
}))
@@ -1,6 +1,5 @@
import '../../__test-utils__/mock-dependencies'
import { beforeEach, describe, expect, it, Mock, vi } from 'vitest'
import { isHostedVersion } from '@/lib/utils'
import { isHosted } from '@/lib/environment'
import { getAllBlocks } from '@/blocks'
import { getProviderFromModel, transformBlockTool } from '@/providers/utils'
import { SerializedBlock, SerializedWorkflow } from '@/serializer/types'
@@ -8,9 +7,35 @@ import { executeTool } from '@/tools'
import { ExecutionContext } from '../../types'
import { AgentBlockHandler } from './agent-handler'
process.env.NEXT_PUBLIC_APP_URL = 'http://localhost:3000'
vi.mock('@/lib/environment', () => ({
isHosted: vi.fn().mockReturnValue(false),
isProd: vi.fn().mockReturnValue(false),
isDev: vi.fn().mockReturnValue(true),
isTest: vi.fn().mockReturnValue(false),
getCostMultiplier: vi.fn().mockReturnValue(1)
}))
vi.mock('@/providers/utils', () => ({
getProviderFromModel: vi.fn().mockReturnValue('mock-provider'),
transformBlockTool: vi.fn(),
getBaseModelProviders: vi.fn().mockReturnValue({ openai: {}, anthropic: {} })
}))
vi.mock('@/blocks', () => ({
getAllBlocks: vi.fn().mockReturnValue([])
}))
vi.mock('@/tools', () => ({
executeTool: vi.fn()
}))
global.fetch = vi.fn()
const mockGetAllBlocks = getAllBlocks as Mock
const mockExecuteTool = executeTool as Mock
const mockIsHostedVersion = isHostedVersion as Mock
const mockIsHosted = isHosted as unknown as Mock
const mockGetProviderFromModel = getProviderFromModel as Mock
const mockTransformBlockTool = transformBlockTool as Mock
const mockFetch = global.fetch as Mock
@@ -60,7 +85,7 @@ describe('AgentBlockHandler', () => {
loops: {},
} as SerializedWorkflow,
}
mockIsHostedVersion.mockReturnValue(false) // Default to non-hosted env for tests
mockIsHosted.mockReturnValue(false) // Default to non-hosted env for tests
mockGetProviderFromModel.mockReturnValue('mock-provider')
// Set up fetch mock to return a successful response
@@ -511,7 +536,7 @@ describe('AgentBlockHandler', () => {
it('should not require API key for gpt-4o on hosted version', async () => {
// Mock hosted environment
mockIsHostedVersion.mockReturnValue(true)
mockIsHosted.mockReturnValue(true)
const inputs = {
model: 'gpt-4o',
+33 -2
View File
@@ -1,5 +1,7 @@
import { emailOTPClient, genericOAuthClient } from 'better-auth/client/plugins'
import { stripeClient } from '@better-auth/stripe/client'
import { createAuthClient } from 'better-auth/react'
import { isProd } from '@/lib/environment'
export function getBaseURL() {
let baseURL
@@ -13,14 +15,43 @@ export function getBaseURL() {
} else if (process.env.NODE_ENV === 'development') {
baseURL = process.env.BETTER_AUTH_URL
}
return baseURL
}
export const client = createAuthClient({
baseURL: getBaseURL(),
plugins: [genericOAuthClient(), emailOTPClient()],
plugins: [
genericOAuthClient(),
emailOTPClient(),
// Only include Stripe client in production
...(isProd ? [
stripeClient({
subscription: true // Enable subscription management
})
] : []),
],
})
export const { useSession } = client
export const useSubscription = () => {
// In development, provide mock implementations
if (!isProd) {
return {
list: async () => ({ data: [] }),
upgrade: async () => ({ error: { message: "Subscriptions are disabled in development mode" } }),
cancel: async () => ({ data: null }),
restore: async () => ({ data: null })
}
}
// In production, use the real implementation
return {
list: client.subscription?.list,
upgrade: client.subscription?.upgrade,
cancel: client.subscription?.cancel,
restore: client.subscription?.restore
}
}
export const { signIn, signUp, signOut } = client
+105
View File
@@ -3,6 +3,8 @@ import { betterAuth } from 'better-auth'
import { drizzleAdapter } from 'better-auth/adapters/drizzle'
import { nextCookies } from 'better-auth/next-js'
import { emailOTP, genericOAuth } from 'better-auth/plugins'
import { stripe } from '@better-auth/stripe'
import Stripe from 'stripe'
import { Resend } from 'resend'
import {
getEmailSubject,
@@ -15,6 +17,12 @@ import * as schema from '@/db/schema'
const logger = createLogger('Auth')
const isProd = process.env.NODE_ENV === 'production'
const stripeClient = new Stripe(process.env.STRIPE_SECRET_KEY || '', {
apiVersion: "2025-02-24.acacia",
})
// If there is no resend key, it might be a local dev environment
// In that case, we don't want to send emails and just log them
@@ -612,6 +620,103 @@ export const auth = betterAuth({
},
],
}),
// Only include the Stripe plugin in production
...(isProd && stripeClient ? [
stripe({
stripeClient,
stripeWebhookSecret: process.env.STRIPE_WEBHOOK_SECRET || '',
createCustomerOnSignUp: true,
onCustomerCreate: async ({ customer, stripeCustomer, user }, request) => {
logger.info('Stripe customer created', {
customerId: customer.id,
userId: user.id
})
},
subscription: {
enabled: true,
plans: [
{
name: 'free',
priceId: process.env.STRIPE_FREE_PRICE_ID || '',
limits: {
cost: process.env.FREE_TIER_COST_LIMIT ? parseInt(process.env.FREE_TIER_COST_LIMIT) : 5,
sharingEnabled: 0,
}
},
{
name: 'pro',
priceId: process.env.STRIPE_PRO_PRICE_ID || '',
limits: {
cost: process.env.PRO_TIER_COST_LIMIT ? parseInt(process.env.PRO_TIER_COST_LIMIT) : 20,
sharingEnabled: 1,
}
}
],
onSubscriptionCreate: async ({
event,
stripeSubscription,
subscription
}: {
event: Stripe.Event
stripeSubscription: Stripe.Subscription
subscription: any
}) => {
logger.info('Subscription created', {
subscriptionId: subscription.id,
referenceId: subscription.referenceId,
plan: subscription.plan,
status: subscription.status
})
},
onSubscriptionUpdated: async ({
subscription,
previousStatus,
user
}: {
subscription: any
previousStatus: string
user: any
}, request?: any) => {
logger.info('Subscription updated', {
subscriptionId: subscription.id,
userId: user.id,
previousStatus,
newStatus: subscription.status
})
},
onSubscriptionDeleted: async ({
event,
stripeSubscription,
subscription
}: {
event: Stripe.Event
stripeSubscription: Stripe.Subscription
subscription: any
}) => {
logger.info('Subscription deleted', {
subscriptionId: subscription.id,
referenceId: subscription.referenceId
})
},
onEvent: async (event: any) => {
logger.info("Stripe webhook hit")
logger.info('Stripe webhook event received', {
type: event.type,
id: event.id
})
switch (event.type) {
case 'customer.subscription.created':
logger.info('Subscription creation event details', {
subscription: event.data.object,
customerId: event.data.object.customer
})
break
}
},
},
})
] : []),
],
pages: {
signIn: '/login',
+32
View File
@@ -0,0 +1,32 @@
/**
* Environment utility functions for consistent environment detection across the application
*/
/**
* Is the application running in production mode
*/
export const isProd = process.env.NODE_ENV === 'production'
/**
* Is the application running in development mode
*/
export const isDev = process.env.NODE_ENV === 'development'
/**
* Is the application running in test mode
*/
export const isTest = process.env.NODE_ENV === 'test'
/**
* Is this the hosted version of the application
*/
export const isHosted = process.env.NEXT_PUBLIC_APP_URL === 'https://www.simstudio.ai'
/**
* Get cost multiplier based on environment
*/
export function getCostMultiplier(): number {
return isProd
? parseFloat(process.env.COST_MULTIPLIER!) || 1
: 1
}
+6 -2
View File
@@ -5,6 +5,7 @@ import { db } from '@/db'
import { userStats, workflow, workflowLogs } from '@/db/schema'
import { ExecutionResult as ExecutorResult } from '@/executor/types'
import { stripCustomToolPrefix } from '../workflows/utils'
import { getCostMultiplier } from '@/lib/environment'
const logger = createLogger('ExecutionLogger')
@@ -545,6 +546,9 @@ export async function persistExecutionLogs(
.from(userStats)
.where(eq(userStats.userId, userId))
const costMultiplier = getCostMultiplier()
const costToStore = totalCost * costMultiplier
if (userStatsRecords.length === 0) {
await db.insert(userStats).values({
id: crypto.randomUUID(),
@@ -554,7 +558,7 @@ export async function persistExecutionLogs(
totalWebhookTriggers: 0,
totalScheduledExecutions: 0,
totalTokensUsed: totalTokens,
totalCost: totalCost.toString(),
totalCost: costToStore.toString(),
lastActive: new Date(),
})
} else {
@@ -562,7 +566,7 @@ export async function persistExecutionLogs(
.update(userStats)
.set({
totalTokensUsed: sql`total_tokens_used + ${totalTokens}`,
totalCost: sql`total_cost + ${totalCost}`,
totalCost: sql`total_cost + ${costToStore}`,
lastActive: new Date(),
})
.where(eq(userStats.userId, userId))
+119
View File
@@ -0,0 +1,119 @@
import { eq } from 'drizzle-orm'
import { db } from '@/db'
import * as schema from '@/db/schema'
import { client } from './auth-client'
import { createLogger } from './logs/console-logger'
import { isProd } from '@/lib/environment'
const logger = createLogger('Subscription')
/**
* Check if the user is on the Pro plan
*/
export async function isProPlan(userId: string): Promise<boolean> {
try {
// In development, enable Pro features for easier testing
if (!isProd) {
return true
}
const { data: subscriptions } = await client.subscription.list({
query: { referenceId: userId }
})
const activeSubscription = subscriptions?.find(
sub => sub.status === 'active' && sub.plan === 'pro'
)
return !!activeSubscription
} catch (error) {
logger.error('Error checking pro plan status', { error, userId })
return false
}
}
/**
* Check if a user has exceeded their cost limit based on their subscription plan
*/
export async function hasExceededCostLimit(userId: string): Promise<boolean> {
try {
// In development, users never exceed their limit
if (!isProd) {
return false
}
logger.info('Checking cost limit for user', { userId })
// Get user's subscription
const { data: subscriptions } = await client.subscription.list({
query: { referenceId: userId }
})
// Find active subscription
const activeSubscription = subscriptions?.find(
sub => sub.status === 'active'
)
// Get configured limits from environment variables or subscription
let costLimit: number
if (activeSubscription && typeof activeSubscription.limits?.cost === 'number') {
// Use the limit from the subscription
costLimit = activeSubscription.limits.cost
} else {
// Use default free tier limit
costLimit = process.env.FREE_TIER_COST_LIMIT
? parseFloat(process.env.FREE_TIER_COST_LIMIT)
: 5
}
logger.info('User cost limit from subscription', { userId, costLimit })
// Get user's actual usage from the database
const statsRecords = await db.select().from(schema.userStats).where(eq(schema.userStats.userId, userId))
if (statsRecords.length === 0) {
// No usage yet, so they haven't exceeded the limit
return false
}
// Get the current cost and compare with the limit
const currentCost = parseFloat(statsRecords[0].totalCost.toString())
return currentCost >= costLimit
} catch (error) {
logger.error('Error checking cost limit', { error, userId })
return false // Be conservative in case of error
}
}
/**
* Check if a user is allowed to share workflows based on their subscription plan
*/
export async function canShareWorkflows(userId: string): Promise<boolean> {
try {
// In development, always allow sharing
if (!isProd) {
return true
}
const { data: subscriptions } = await client.subscription.list({
query: { referenceId: userId }
})
const activeSubscription = subscriptions?.find(
sub => sub.status === 'active'
)
// If no active subscription or subscription is free tier, sharing is not allowed
if (!activeSubscription || activeSubscription.plan === 'free') {
return false
}
// Check if the plan's limits include sharing
return !!activeSubscription.limits?.sharingEnabled
} catch (error) {
logger.error('Error checking sharing permission', { error, userId })
return false // Be conservative in case of error
}
}
-1
View File
@@ -44,7 +44,6 @@ vi.mock('./setup', () => ({
S3_CONFIG: {
bucket: 'test-bucket',
region: 'test-region',
baseUrl: 'https://test-bucket.s3.test-region.amazonaws.com'
}
}))
+1 -4
View File
@@ -1,13 +1,10 @@
import { S3Client, PutObjectCommand, GetObjectCommand, DeleteObjectCommand } from '@aws-sdk/client-s3'
import { getSignedUrl } from '@aws-sdk/s3-request-presigner'
import { createLogger } from '@/lib/logs/console-logger'
import { S3_CONFIG } from './setup'
const logger = createLogger('S3Client')
// Create an S3 client
export const s3Client = new S3Client({
region: S3_CONFIG.region,
region: S3_CONFIG.region || '',
credentials: {
accessKeyId: process.env.AWS_ACCESS_KEY_ID || '',
secretAccessKey: process.env.AWS_SECRET_ACCESS_KEY || ''
+1 -5
View File
@@ -1,4 +1,4 @@
import { ensureUploadsDirectory, USE_S3_STORAGE, S3_CONFIG } from './setup'
import { ensureUploadsDirectory, USE_S3_STORAGE } from './setup'
import { createLogger } from '@/lib/logs/console-logger'
const logger = createLogger('UploadsSetup')
@@ -9,10 +9,6 @@ if (typeof process !== 'undefined') {
logger.info(`Storage mode: ${USE_S3_STORAGE ? 'S3' : 'Local'}`)
if (USE_S3_STORAGE) {
logger.info('Using S3 storage mode with configuration:')
logger.info(`- Bucket: ${S3_CONFIG.bucket}`)
logger.info(`- Region: ${S3_CONFIG.region}`)
// Verify AWS credentials
if (!process.env.AWS_ACCESS_KEY_ID || !process.env.AWS_SECRET_ACCESS_KEY) {
logger.warn('AWS credentials are not set in environment variables.')
+2 -5
View File
@@ -15,9 +15,8 @@ export const UPLOAD_DIR = join(PROJECT_ROOT, 'uploads')
export const USE_S3_STORAGE = process.env.NODE_ENV === 'production' || process.env.USE_S3 === 'true'
export const S3_CONFIG = {
bucket: process.env.S3_BUCKET_NAME || 'sim-studio-files',
region: process.env.AWS_REGION || 'us-east-1',
baseUrl: process.env.S3_BASE_URL || `https://${process.env.S3_BUCKET_NAME || 'sim-studio-files'}.s3.${process.env.AWS_REGION || 'us-east-1'}.amazonaws.com`
bucket: process.env.S3_BUCKET_NAME || '',
region: process.env.AWS_REGION || '',
}
/**
@@ -31,9 +30,7 @@ export async function ensureUploadsDirectory() {
try {
if (!existsSync(UPLOAD_DIR)) {
logger.info(`Creating uploads directory at ${UPLOAD_DIR}`)
await mkdir(UPLOAD_DIR, { recursive: true })
logger.info(`Created uploads directory at ${UPLOAD_DIR}`)
} else {
logger.info(`Uploads directory already exists at ${UPLOAD_DIR}`)
}
+267
View File
@@ -0,0 +1,267 @@
import { isProPlan } from './subscription'
import { createLogger } from './logs/console-logger'
import { db } from '@/db'
import { eq } from 'drizzle-orm'
import { userStats } from '@/db/schema'
import { client } from './auth-client'
import { isProd } from '@/lib/environment'
const logger = createLogger('UsageMonitor')
// Percentage threshold for showing warning
const WARNING_THRESHOLD = 80
interface UsageData {
percentUsed: number
isWarning: boolean
isExceeded: boolean
currentUsage: number
limit: number
}
/**
* Checks a user's cost usage against their subscription plan limit
* and returns usage information including whether they're approaching the limit
*/
export async function checkUsageStatus(userId: string): Promise<UsageData> {
try {
// In development, always return permissive limits
if (!isProd) {
// Get actual usage from the database for display purposes
const statsRecords = await db.select().from(userStats).where(eq(userStats.userId, userId))
const currentUsage = statsRecords.length > 0
? parseFloat(statsRecords[0].totalCost.toString())
: 0
// In development, set a very high limit to avoid restrictions
const devLimit = 1000
return {
percentUsed: Math.min(Math.round((currentUsage / devLimit) * 100), 100),
isWarning: false,
isExceeded: false,
currentUsage,
limit: devLimit
}
}
// Production environment - check real subscription limits
// Get user's subscription details
const isPro = await isProPlan(userId)
// Get the subscription limits
const { data: subscriptions } = await client.subscription.list({
query: { referenceId: userId }
})
// Find active subscription
const activeSubscription = subscriptions?.find(
sub => sub.status === 'active' || sub.status === 'trialing'
)
// Get configured limits from environment variables or subscription
let limit: number
if (activeSubscription && typeof activeSubscription.limits?.cost === 'number') {
// Use the limit from the subscription if available
limit = activeSubscription.limits.cost
} else {
// Fallback to environment variables
const freeLimit = process.env.FREE_TIER_COST_LIMIT
? parseFloat(process.env.FREE_TIER_COST_LIMIT)
: 5
const proLimit = process.env.PRO_TIER_COST_LIMIT
? parseFloat(process.env.PRO_TIER_COST_LIMIT)
: 50
// Set the appropriate limit based on subscription
limit = isPro ? proLimit : freeLimit
}
// Get actual usage from the database
const statsRecords = await db.select().from(userStats).where(eq(userStats.userId, userId))
// If no stats record exists, create a default one
if (statsRecords.length === 0) {
return {
percentUsed: 0,
isWarning: false,
isExceeded: false,
currentUsage: 0,
limit
}
}
// Get the current cost from the user stats
const currentUsage = parseFloat(statsRecords[0].totalCost.toString())
// Calculate percentage used
const percentUsed = Math.min(Math.round((currentUsage / limit) * 100), 100)
// Check if usage exceeds threshold or limit
const isWarning = percentUsed >= WARNING_THRESHOLD && percentUsed < 100
const isExceeded = currentUsage >= limit
return {
percentUsed,
isWarning,
isExceeded,
currentUsage,
limit
}
} catch (error) {
logger.error('Error checking usage status', { error, userId })
// Return default values in case of error
return {
percentUsed: 0,
isWarning: false,
isExceeded: false,
currentUsage: 0,
limit: 0
}
}
}
/**
* Displays a notification to the user when they're approaching their usage limit
* Can be called on app startup or before executing actions that might incur costs
*/
export async function checkAndNotifyUsage(userId: string): Promise<void> {
try {
// Skip usage notifications in development
if (!isProd) {
return
}
const usageData = await checkUsageStatus(userId)
if (usageData.isExceeded) {
// User has exceeded their limit
logger.warn('User has exceeded usage limits', {
userId,
usage: usageData.currentUsage,
limit: usageData.limit
})
// Dispatch event to show a UI notification
if (typeof window !== 'undefined') {
window.dispatchEvent(new CustomEvent('usage-exceeded', {
detail: { usageData }
}))
}
} else if (usageData.isWarning) {
// User is approaching their limit
logger.info('User approaching usage limits', {
userId,
usage: usageData.currentUsage,
limit: usageData.limit,
percent: usageData.percentUsed
})
// Dispatch event to show a UI notification
if (typeof window !== 'undefined') {
window.dispatchEvent(new CustomEvent('usage-warning', {
detail: { usageData }
}))
// Optionally open the subscription tab in settings
window.dispatchEvent(new CustomEvent('open-settings', {
detail: { tab: 'subscription' }
}))
}
}
} catch (error) {
logger.error('Error in usage notification system', { error, userId })
}
}
// Add this function to check usage limits on the server-side for API routes
/**
* Server-side function to check if a user has exceeded their usage limits
* For use in API routes, webhooks, and scheduled executions
*
* @param userId The ID of the user to check
* @returns An object containing the exceeded status and usage details
*/
export async function checkServerSideUsageLimits(userId: string): Promise<{
isExceeded: boolean;
currentUsage: number;
limit: number;
message?: string;
}> {
try {
// In development, always allow execution
if (!isProd) {
return {
isExceeded: false,
currentUsage: 0,
limit: 1000,
}
}
logger.info('Server-side checking usage limits for user', { userId })
// Get the user's subscription
const { data: subscriptions } = await client.subscription.list({
query: { referenceId: userId }
})
// Find active subscription
const activeSubscription = subscriptions?.find(
sub => sub.status === 'active' || sub.status === 'trialing'
)
// Get configured limits from environment variables or subscription
let costLimit: number
if (activeSubscription && typeof activeSubscription.limits?.cost === 'number') {
// Use the limit from the subscription
costLimit = activeSubscription.limits.cost
} else {
// Use default free tier limit
costLimit = process.env.FREE_TIER_COST_LIMIT
? parseFloat(process.env.FREE_TIER_COST_LIMIT)
: 5
}
logger.info('Server-side user cost limit from subscription', { userId, costLimit })
// Get user's actual usage from the database
const statsRecords = await db.select().from(userStats).where(eq(userStats.userId, userId))
if (statsRecords.length === 0) {
// No usage yet, so they haven't exceeded the limit
return {
isExceeded: false,
currentUsage: 0,
limit: costLimit
}
}
// Get the current cost and compare with the limit
const currentUsage = parseFloat(statsRecords[0].totalCost.toString())
const isExceeded = currentUsage >= costLimit
return {
isExceeded,
currentUsage,
limit: costLimit,
message: isExceeded
? `Usage limit exceeded: ${currentUsage.toFixed(2)}$ used of ${costLimit}$ limit. Please upgrade your plan to continue.`
: undefined
}
} catch (error) {
logger.error('Error in server-side usage limit check', { error, userId })
// Be conservative in case of error - allow execution but log the issue
return {
isExceeded: false,
currentUsage: 0,
limit: 0,
message: `Error checking usage limits: ${error instanceof Error ? error.message : String(error)}`
}
}
}
+11 -17
View File
@@ -272,16 +272,6 @@ export function generateApiKey(): string {
return `sim_${nanoid(32)}`
}
/**
* Determines if the application is running on the hosted/production version
* @returns boolean indicating if the app is running on the hosted version
*/
export function isHostedVersion(): boolean {
return (
typeof window !== 'undefined' && process.env.NEXT_PUBLIC_APP_URL === 'https://www.simstudio.ai'
)
}
/**
* Rotates through available API keys for a provider
* @param provider - The provider to get a key for (e.g., 'openai')
@@ -289,21 +279,25 @@ export function isHostedVersion(): boolean {
* @throws Error if no API keys are configured for rotation
*/
export function getRotatingApiKey(provider: string): string {
if (provider !== 'openai') {
if (provider !== 'openai' && provider !== 'anthropic') {
throw new Error(`No rotation implemented for provider: ${provider}`)
}
// Get all OpenAI keys from environment
const keys = []
// Add keys if they exist in environment variables
if (process.env.OPENAI_API_KEY_1) keys.push(process.env.OPENAI_API_KEY_1)
if (process.env.OPENAI_API_KEY_2) keys.push(process.env.OPENAI_API_KEY_2)
if (process.env.OPENAI_API_KEY_3) keys.push(process.env.OPENAI_API_KEY_3)
if (provider === 'openai') {
if (process.env.OPENAI_API_KEY_1) keys.push(process.env.OPENAI_API_KEY_1)
if (process.env.OPENAI_API_KEY_2) keys.push(process.env.OPENAI_API_KEY_2)
if (process.env.OPENAI_API_KEY_3) keys.push(process.env.OPENAI_API_KEY_3)
} else if (provider === 'anthropic') {
if (process.env.ANTHROPIC_API_KEY_1) keys.push(process.env.ANTHROPIC_API_KEY_1)
if (process.env.ANTHROPIC_API_KEY_2) keys.push(process.env.ANTHROPIC_API_KEY_2)
if (process.env.ANTHROPIC_API_KEY_3) keys.push(process.env.ANTHROPIC_API_KEY_3)
}
if (keys.length === 0) {
throw new Error(
'No API keys configured for rotation. Please configure OPENAI_API_KEY_1, OPENAI_API_KEY_2, or OPENAI_API_KEY_3.'
`No API keys configured for rotation. Please configure ${provider.toUpperCase()}_API_KEY_1, ${provider.toUpperCase()}_API_KEY_2, or ${provider.toUpperCase()}_API_KEY_3.`
)
}
+2 -4
View File
@@ -1,14 +1,12 @@
import { NextRequest } from 'next/server'
import { getRedisClient } from '../redis'
import { isProd } from '@/lib/environment'
// Configuration
const RATE_LIMIT_WINDOW = 60 // 1 minute window (in seconds)
const WAITLIST_MAX_REQUESTS = 5 // 5 requests per minute per IP
const WAITLIST_BLOCK_DURATION = 15 * 60 // 15 minutes block (in seconds)
// Environment detection
const isProduction = process.env.NODE_ENV === 'production'
// Fallback in-memory store for development or if Redis fails
const inMemoryStore = new Map<
string,
@@ -16,7 +14,7 @@ const inMemoryStore = new Map<
>()
// Clean up in-memory store periodically (only used in development)
if (!isProduction && typeof setInterval !== 'undefined') {
if (!isProd && typeof setInterval !== 'undefined') {
setInterval(
() => {
const now = Math.floor(Date.now() / 1000)
+6491 -1632
View File
File diff suppressed because it is too large Load Diff
+2
View File
@@ -30,6 +30,7 @@
"@anthropic-ai/sdk": "^0.39.0",
"@aws-sdk/client-s3": "^3.779.0",
"@aws-sdk/s3-request-presigner": "^3.779.0",
"@better-auth/stripe": "^1.2.7",
"@browserbasehq/stagehand": "^2.0.0",
"@cerebras/cerebras_cloud_sdk": "^1.23.0",
"@hookform/resolvers": "^4.1.3",
@@ -87,6 +88,7 @@
"react-simple-code-editor": "^0.14.1",
"reactflow": "^11.11.4",
"resend": "^4.1.2",
"stripe": "^17.7.0",
"tailwind-merge": "^2.6.0",
"tailwindcss-animate": "^1.0.7",
"uuid": "^11.1.0",
+61 -19
View File
@@ -1,37 +1,79 @@
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
import * as environmentModule from '@/lib/environment'
import { getApiKey } from './utils'
// Skip the tests that need proper module mocking
const isHostedSpy = vi.spyOn(environmentModule, 'isHosted', 'get')
const mockGetRotatingApiKey = vi.fn().mockReturnValue('rotating-server-key')
const originalRequire = module.require
describe('getApiKey', () => {
// Save original env and reset between tests
const originalEnv = { ...process.env }
beforeEach(() => {
vi.clearAllMocks()
isHostedSpy.mockReturnValue(false)
module.require = vi.fn(() => ({
getRotatingApiKey: mockGetRotatingApiKey
}))
})
afterEach(() => {
// Reset env vars after each test
process.env = { ...originalEnv }
module.require = originalRequire
})
it('should return user-provided key for non-gpt-4o models', () => {
const key = getApiKey('openai', 'o1', 'user-key')
it('should return user-provided key when not in hosted environment', () => {
isHostedSpy.mockReturnValue(false)
// For OpenAI
const key1 = getApiKey('openai', 'gpt-4', 'user-key-openai')
expect(key1).toBe('user-key-openai')
// For Anthropic
const key2 = getApiKey('anthropic', 'claude-3', 'user-key-anthropic')
expect(key2).toBe('user-key-anthropic')
})
it('should throw error if no key provided in non-hosted environment', () => {
isHostedSpy.mockReturnValue(false)
expect(() => getApiKey('openai', 'gpt-4')).toThrow('API key is required for openai gpt-4')
expect(() => getApiKey('anthropic', 'claude-3')).toThrow('API key is required for anthropic claude-3')
})
it('should fall back to user key in hosted environment if rotation fails', () => {
isHostedSpy.mockReturnValue(true)
module.require = vi.fn(() => {
throw new Error('Rotation failed')
})
const key = getApiKey('openai', 'gpt-4', 'user-fallback-key')
expect(key).toBe('user-fallback-key')
})
it('should throw error in hosted environment if rotation fails and no user key', () => {
isHostedSpy.mockReturnValue(true)
module.require = vi.fn(() => {
throw new Error('Rotation failed')
})
expect(() => getApiKey('openai', 'gpt-4')).toThrow('No API key available for openai gpt-4')
})
it('should require user key for non-OpenAI/Anthropic providers even in hosted environment', () => {
isHostedSpy.mockReturnValue(true)
const key = getApiKey('other-provider', 'some-model', 'user-key')
expect(key).toBe('user-key')
})
it('should throw error if no key provided for non-gpt-4o models', () => {
expect(() => getApiKey('openai', 'o1')).toThrow('API key is required for openai o1')
})
it('should require user key for gpt-4o on non-hosted environments', () => {
process.env.NEXT_PUBLIC_APP_URL = 'http://localhost:3000'
// Should work with user key
const key = getApiKey('openai', 'gpt-4o', 'user-key')
expect(key).toBe('user-key')
// Should throw without user key
expect(() => getApiKey('openai', 'gpt-4o')).toThrow('API key is required for openai gpt-4o')
expect(() => getApiKey('other-provider', 'some-model')).toThrow(
'API key is required for other-provider some-model'
)
})
})
+12 -8
View File
@@ -1,5 +1,6 @@
import { createLogger } from '@/lib/logs/console-logger'
import { useCustomToolsStore } from '@/stores/custom-tools/store'
import { isProd, getCostMultiplier } from '@/lib/environment'
import { anthropicProvider } from './anthropic'
import { cerebrasProvider } from './cerebras'
import { deepseekProvider } from './deepseek'
@@ -10,6 +11,7 @@ import { openaiProvider } from './openai'
import { getModelPricing } from './pricing'
import { ProviderConfig, ProviderId, ProviderToolConfig } from './types'
import { xAIProvider } from './xai'
import { isHosted } from '@/lib/environment'
const logger = createLogger('ProviderUtils')
@@ -426,11 +428,13 @@ export function calculateCost(
const outputCost = completionTokens * (pricing.output / 1_000_000)
const totalCost = inputCost + outputCost
const costMultiplier = getCostMultiplier()
return {
input: parseFloat(inputCost.toFixed(6)),
output: parseFloat(outputCost.toFixed(6)),
total: parseFloat(totalCost.toFixed(6)),
input: parseFloat((inputCost * costMultiplier).toFixed(6)),
output: parseFloat((outputCost * costMultiplier).toFixed(6)),
total: parseFloat((totalCost * costMultiplier).toFixed(6)),
pricing,
}
}
@@ -471,15 +475,15 @@ export function getApiKey(provider: string, model: string, userProvidedKey?: str
// If user provided a key, use it as a fallback
const hasUserKey = !!userProvidedKey
// Only use server key rotation for OpenAI's gpt-4o model on the hosted platform
const isHostedVersion = process.env.NEXT_PUBLIC_APP_URL === 'https://www.simstudio.ai'
const isGPT4o = model === 'gpt-4o' && provider === 'openai'
// Use server key rotation for all OpenAI models and Anthropic's Claude models on the hosted platform
const isOpenAIModel = provider === 'openai'
const isClaudeModel = provider === 'anthropic'
if (isHostedVersion && isGPT4o) {
if (isHosted && (isOpenAIModel || isClaudeModel)) {
try {
// Import the key rotation function
const { getRotatingApiKey } = require('@/lib/utils')
const serverKey = getRotatingApiKey('openai')
const serverKey = getRotatingApiKey(provider)
return serverKey
} catch (error) {
// If server key fails and we have a user key, fallback to that