diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts index eddcd9cb71..9404f81878 100644 --- a/app/api/chat/route.ts +++ b/app/api/chat/route.ts @@ -1,102 +1,104 @@ -import { OpenAI } from 'openai' import { NextResponse } from 'next/server' -import { z } from 'zod' +import { OpenAI } from 'openai' import { ChatCompletionMessageParam } from 'openai/resources/chat/completions' +import { z } from 'zod' // Validation schemas const MessageSchema = z.object({ role: z.enum(['user', 'assistant', 'system']), - content: z.string() + content: z.string(), }) const RequestSchema = z.object({ messages: z.array(MessageSchema), workflowState: z.object({ blocks: z.record(z.any()), - edges: z.array(z.any()) - }) + edges: z.array(z.any()), + }), }) // Define function schemas with strict typing const workflowActions = { addBlock: { - description: "Add one new block to the workflow", + description: 'Add one new block to the workflow', parameters: { - type: "object", - required: ["type"], + type: 'object', + required: ['type'], properties: { type: { - type: "string", - enum: ["agent", "api", "condition", "function", "router"], - description: "The type of block to add" + type: 'string', + enum: ['agent', 'api', 'condition', 'function', 'router'], + description: 'The type of block to add', }, name: { - type: "string", - description: "Optional custom name for the block. Do not provide a name unless the user has specified it." + type: 'string', + description: + 'Optional custom name for the block. Do not provide a name unless the user has specified it.', }, position: { - type: "object", - description: "Optional position for the block. Do not provide a position unless the user has specified it.", + type: 'object', + description: + 'Optional position for the block. Do not provide a position unless the user has specified it.', properties: { - x: { type: "number" }, - y: { type: "number" } - } - } - } - } + x: { type: 'number' }, + y: { type: 'number' }, + }, + }, + }, + }, }, addEdge: { - description: "Create a connection (edge) between two blocks", + description: 'Create a connection (edge) between two blocks', parameters: { - type: "object", - required: ["sourceId", "targetId"], + type: 'object', + required: ['sourceId', 'targetId'], properties: { sourceId: { - type: "string", - description: "ID of the source block" + type: 'string', + description: 'ID of the source block', }, targetId: { - type: "string", - description: "ID of the target block" + type: 'string', + description: 'ID of the target block', }, - sourceHandle: { - type: "string", - description: "Optional handle identifier for the source connection point" + sourceHandle: { + type: 'string', + description: 'Optional handle identifier for the source connection point', }, - targetHandle: { - type: "string", - description: "Optional handle identifier for the target connection point" - } - } - } + targetHandle: { + type: 'string', + description: 'Optional handle identifier for the target connection point', + }, + }, + }, }, removeBlock: { - description: "Remove a block from the workflow", + description: 'Remove a block from the workflow', parameters: { - type: "object", - required: ["id"], + type: 'object', + required: ['id'], properties: { - id: { type: "string", description: "ID of the block to remove" } - } - } + id: { type: 'string', description: 'ID of the block to remove' }, + }, + }, }, removeEdge: { - description: "Remove a connection (edge) between blocks", + description: 'Remove a connection (edge) between blocks', parameters: { - type: "object", - required: ["id"], + type: 'object', + required: ['id'], properties: { - id: { type: "string", description: "ID of the edge to remove" } - } - } - } + id: { type: 'string', description: 'ID of the edge to remove' }, + }, + }, + }, } // System prompt that references workflow state const getSystemPrompt = (workflowState: any) => { const blockCount = Object.keys(workflowState.blocks).length const edgeCount = workflowState.edges.length - + // Create a summary of existing blocks const blockSummary = Object.values(workflowState.blocks) .map((block: any) => `- ${block.type} block named "${block.name}" with id ${block.id}`) @@ -110,10 +112,14 @@ const getSystemPrompt = (workflowState: any) => { return `You are a workflow assistant that helps users modify their workflow by adding/removing blocks and connections. Current Workflow State: -${blockCount === 0 ? 'The workflow is empty.' : `${blockSummary} +${ + blockCount === 0 + ? 'The workflow is empty.' + : `${blockSummary} Connections: -${edgeCount === 0 ? 'No connections between blocks.' : edgeSummary}`} +${edgeCount === 0 ? 'No connections between blocks.' : edgeSummary}` +} When users request changes: - Consider existing blocks when suggesting connections @@ -133,10 +139,7 @@ export async function POST(request: Request) { // Validate API key const apiKey = request.headers.get('X-OpenAI-Key') if (!apiKey) { - return NextResponse.json( - { error: 'OpenAI API key is required' }, - { status: 401 } - ) + return NextResponse.json({ error: 'OpenAI API key is required' }, { status: 401 }) } // Parse and validate request body @@ -146,52 +149,53 @@ export async function POST(request: Request) { // Initialize OpenAI client const openai = new OpenAI({ apiKey }) - + // Create message history with workflow context const messageHistory = [ { role: 'system', content: getSystemPrompt(workflowState) }, - ...messages + ...messages, ] // Make OpenAI API call with workflow context const completion = await openai.chat.completions.create({ - model: "gpt-4o", + model: 'gpt-4o', messages: messageHistory as ChatCompletionMessageParam[], tools: Object.entries(workflowActions).map(([name, config]) => ({ type: 'function', function: { name, description: config.description, - parameters: config.parameters - } + parameters: config.parameters, + }, })), - tool_choice: "auto" + tool_choice: 'auto', }) const message = completion.choices[0].message // Process tool calls if present if (message.tool_calls) { - console.log(message.tool_calls) - const actions = message.tool_calls.map(call => ({ + console.log(message.tool_calls) + const actions = message.tool_calls.map((call) => ({ name: call.function.name, - parameters: JSON.parse(call.function.arguments) + parameters: JSON.parse(call.function.arguments), })) return NextResponse.json({ message: message.content || "I've updated the workflow based on your request.", - actions + actions, }) } // Return response with no actions - return NextResponse.json({ - message: message.content || "I'm not sure what changes to make to the workflow. Can you please provide more specific instructions?" + return NextResponse.json({ + message: + message.content || + "I'm not sure what changes to make to the workflow. Can you please provide more specific instructions?", }) - } catch (error) { console.error('Chat API error:', error) - + // Handle specific error types if (error instanceof z.ZodError) { return NextResponse.json( @@ -200,9 +204,6 @@ export async function POST(request: Request) { ) } - return NextResponse.json( - { error: 'Failed to process chat message' }, - { status: 500 } - ) + return NextResponse.json({ error: 'Failed to process chat message' }, { status: 500 }) } -} \ No newline at end of file +} diff --git a/components/ui/command.tsx b/components/ui/command.tsx index b8adf23b03..45e653fd39 100644 --- a/components/ui/command.tsx +++ b/components/ui/command.tsx @@ -12,6 +12,12 @@ import { cn } from '@/lib/utils' // This file is not typed correctly from shadcn, so we're disabling the type checker // @ts-nocheck +// This file is not typed correctly from shadcn, so we're disabling the type checker +// @ts-nocheck + +// This file is not typed correctly from shadcn, so we're disabling the type checker +// @ts-nocheck + const Command = React.forwardRef< React.ElementRef, React.ComponentPropsWithoutRef & { diff --git a/executor/index.ts b/executor/index.ts index 0ad49f6a59..006f4a1abf 100644 --- a/executor/index.ts +++ b/executor/index.ts @@ -14,13 +14,16 @@ import { resolveBlockReferences, resolveEnvVariables } from './utils' * Handles parallel execution, state management, and special block types. */ export class Executor { + private loopIterations: Map + constructor( private workflow: SerializedWorkflow, - // Initial block states can be passed in (e.g., for resuming workflows or pre-populating data) private initialBlockStates: Record = {}, private environmentVariables: Record = {}, private processedConditionBlocks: Set = new Set() - ) {} + ) { + this.loopIterations = new Map() + } /** * Main entry point for workflow execution. @@ -116,10 +119,9 @@ export class Executor { private async executeInParallel(context: ExecutionContext): Promise { const { blocks, connections } = this.workflow - // Track iterations per loop - const loopIterations = new Map() + this.loopIterations.clear() for (const [loopId, loop] of Object.entries(this.workflow.loops || {})) { - loopIterations.set(loopId, 0) + this.loopIterations.set(loopId, 0) } // Build dependency graphs: inDegree (number of incoming edges) and adjacency (outgoing connections) @@ -135,7 +137,6 @@ export class Executor { // Helper functions for identifying entry points and feedback edges const isEntryBlock = (blockId: string): boolean => { - const block = blocks.find((b) => b.id === blockId) const starterBlock = blocks.find((b) => b.metadata?.id === 'starter') // Entry blocks are those that are directly connected to the starter block @@ -334,11 +335,10 @@ export class Executor { const executedLoopBlocks = layerResults.filter((blockId) => loopBlocks.has(blockId)) if (executedLoopBlocks.length > 0) { - const iterations = loopIterations.get(loopId) || 0 + const iterations = this.loopIterations.get(loopId) || 0 // Only process if we haven't hit max iterations - if (iterations < loop.maxIterations - 1) { - // Removed the -1 + if (iterations < loop.maxIterations) { // Check if any block in the loop has outgoing connections to other blocks in the loop const hasLoopConnection = executedLoopBlocks.some((blockId) => { const outgoingConns = connections.filter((conn) => conn.source === blockId) @@ -351,19 +351,8 @@ export class Executor { return block?.metadata?.id === 'condition' }) - if (hasLoopConnection) { - // Reset the loop blocks' inDegrees - resetLoopBlocksDegrees(loopId) - for (const blockId of loop.nodes) { - if (inDegree.get(blockId) === 0) { - queue.push(blockId) - } - } - - // Only increment counter when we complete a loop cycle - if (isLoopComplete) { - loopIterations.set(loopId, iterations + 1) - } + if (hasLoopConnection && isLoopComplete) { + this.loopIterations.set(loopId, iterations + 1) } } } @@ -373,6 +362,48 @@ export class Executor { return lastOutput } + private resetLoopBlocksDegrees( + loopId: string, + inDegree: Map, + isFeedbackEdge: (conn: (typeof this.workflow.connections)[number]) => boolean + ): void { + const loop = this.workflow.loops?.[loopId] + if (!loop) return + + for (const blockId of loop.nodes) { + // For each block in the loop, recalculate its initial inDegree + let degree = 0 + for (const conn of this.workflow.connections) { + if (conn.target === blockId && loop.nodes.includes(conn.source)) { + // Count non-feedback edges within the loop + if (!isFeedbackEdge(conn)) { + degree++ + } + } + } + inDegree.set(blockId, degree) + } + } + + private shouldResetLoop(loopId: string, blockId: string, chosenPath: string): boolean { + const loop = this.workflow.loops?.[loopId] + if (!loop) return false + + const iterations = this.loopIterations.get(loopId) || 0 + const block = this.workflow.blocks.find((b) => b.id === blockId) + const isConditionBlock = block?.metadata?.id === 'condition' + + // Get execution order within the loop + const loopBlocks = loop.nodes + const sourceIndex = loopBlocks.indexOf(blockId) + const targetIndex = loopBlocks.indexOf(chosenPath) + + // Check if this is a feedback path (points to an earlier block in the loop) + const isFeedbackPath = targetIndex < sourceIndex + + return isConditionBlock && isFeedbackPath && iterations < loop.maxIterations + } + /** * Executes a single block with appropriate tool or provider. * Handles different block types (router, evaluator, condition, agent). @@ -1023,13 +1054,50 @@ export class Executor { const conditionId = conn.sourceHandle.replace('condition-', '') const activeCondition = activeConditionalPaths?.get(sourceBlockId) - // Only decrement if this connection matches the active condition path + // Only process if this is the active condition path if (activeCondition === conditionId) { - // Only decrement once per condition block, regardless of number of connections if (!this.processedConditionBlocks.has(`${sourceBlockId}-${conn.target}`)) { + // Check if this is a loop-back connection first + const loopId = Object.keys(this.workflow.loops || {}).find( + (id) => + this.workflow.loops?.[id].nodes.includes(conn.target) && + this.workflow.loops?.[id].nodes.includes(sourceBlockId) + ) + + if (loopId) { + const loop = this.workflow.loops?.[loopId] + if (loop) { + const sourceIndex = loop.nodes.indexOf(sourceBlockId) + const targetIndex = loop.nodes.indexOf(conn.target) + const isFeedbackPath = targetIndex < sourceIndex + + if (isFeedbackPath) { + const iterations = this.loopIterations.get(loopId) || 0 + if (iterations < loop.maxIterations) { + // Reset all blocks in the loop + this.resetLoopBlocksDegrees(loopId, inDegree, (conn) => { + if (!conn.sourceHandle?.startsWith('condition-')) return false + const loopBlocks = loop.nodes + const srcIndex = loopBlocks.indexOf(conn.source) + const tgtIndex = loopBlocks.indexOf(conn.target) + return tgtIndex < srcIndex + }) + + // Add loop entry block to queue + const entryBlock = loop.nodes[0] + if (inDegree.get(entryBlock) === 0) { + queue.push(entryBlock) + } + + this.loopIterations.set(loopId, iterations + 1) + } + } + } + } + const newDegree = (inDegree.get(conn.target) || 0) - 1 inDegree.set(conn.target, newDegree) - if (newDegree === 0) { + if (newDegree === 0 && !queue.includes(conn.target)) { queue.push(conn.target) } this.processedConditionBlocks.add(`${sourceBlockId}-${conn.target}`) diff --git a/stores/chat/store.ts b/stores/chat/store.ts index 4e2551ed82..24c02b3505 100644 --- a/stores/chat/store.ts +++ b/stores/chat/store.ts @@ -1,9 +1,9 @@ import { create } from 'zustand' import { devtools } from 'zustand/middleware' -import { useWorkflowStore } from '../workflow/store' import { useEnvironmentStore } from '../environment/store' -import { ChatStore, ChatMessage } from './types' -import { getNextBlockNumber, calculateBlockPosition } from './utils' +import { useWorkflowStore } from '../workflow/store' +import { ChatMessage, ChatStore } from './types' +import { calculateBlockPosition, getNextBlockNumber } from './utils' export const useChatStore = create()( devtools( @@ -15,12 +15,14 @@ export const useChatStore = create()( sendMessage: async (content: string) => { try { set({ isProcessing: true, error: null }) - + const workflowStore = useWorkflowStore.getState() const apiKey = useEnvironmentStore.getState().getVariable('OPENAI_API_KEY') - + if (!apiKey) { - throw new Error('OpenAI API key not found. Please add it to your environment variables.') + throw new Error( + 'OpenAI API key not found. Please add it to your environment variables.' + ) } // User message @@ -33,34 +35,34 @@ export const useChatStore = create()( // Format messages for OpenAI API const formattedMessages = [ - ...get().messages.map(msg => ({ + ...get().messages.map((msg) => ({ role: msg.role, - content: msg.content + content: msg.content, })), { role: newMessage.role, - content: newMessage.content - } + content: newMessage.content, + }, ] // Add message to local state first - set(state => ({ - messages: [...state.messages, newMessage] + set((state) => ({ + messages: [...state.messages, newMessage], })) const response = await fetch('/api/chat', { method: 'POST', - headers: { + headers: { 'Content-Type': 'application/json', - 'X-OpenAI-Key': apiKey + 'X-OpenAI-Key': apiKey, }, body: JSON.stringify({ messages: formattedMessages, workflowState: { blocks: workflowStore.blocks, - edges: workflowStore.edges - } - }) + edges: workflowStore.edges, + }, + }), }) if (!response.ok) { @@ -70,35 +72,28 @@ export const useChatStore = create()( const data = await response.json() console.log('OPENAI RESPONSE', data) - + // Handle any actions returned from the API if (data.actions) { // Process all block additions first to properly calculate positions - const blockActions = data.actions.filter( - (action: any) => action.name === 'addBlock' - ) - + const blockActions = data.actions.filter((action: any) => action.name === 'addBlock') + blockActions.forEach((action: any, index: number) => { const { type, name } = action.parameters const id = crypto.randomUUID() - + // Calculate position based on current blocks and action index - const position = calculateBlockPosition( - workflowStore.blocks, - index - ) - + const position = calculateBlockPosition(workflowStore.blocks, index) + // Generate name if not provided const blockName = name || `${type} ${getNextBlockNumber(workflowStore.blocks, type)}` - + workflowStore.addBlock(id, type, blockName, position) }) // Handle other actions (edges, removals, etc.) - const otherActions = data.actions.filter( - (action: any) => action.name !== 'addBlock' - ) - + const otherActions = data.actions.filter((action: any) => action.name !== 'addBlock') + otherActions.forEach((action: any) => { switch (action.name) { case 'addEdge': { @@ -109,7 +104,7 @@ export const useChatStore = create()( target: targetId, sourceHandle, targetHandle, - type: 'custom' + type: 'custom', }) break } @@ -127,16 +122,18 @@ export const useChatStore = create()( // Add assistant's response to chat if (data.message) { - set(state => ({ - messages: [...state.messages, { - id: crypto.randomUUID(), - role: 'assistant', - content: data.message, - timestamp: Date.now() - }] + set((state) => ({ + messages: [ + ...state.messages, + { + id: crypto.randomUUID(), + role: 'assistant', + content: data.message, + timestamp: Date.now(), + }, + ], })) } - } catch (error) { console.error('Chat error:', error) set({ error: error instanceof Error ? error.message : 'Unknown error' }) @@ -146,8 +143,8 @@ export const useChatStore = create()( }, clearChat: () => set({ messages: [], error: null }), - setError: (error) => set({ error }) + setError: (error) => set({ error }), }), { name: 'chat-store' } ) -) \ No newline at end of file +) diff --git a/stores/chat/types.ts b/stores/chat/types.ts index 62e17a7813..4fe5bed20d 100644 --- a/stores/chat/types.ts +++ b/stores/chat/types.ts @@ -17,4 +17,4 @@ export interface ChatActions { setError: (error: string | null) => void } -export type ChatStore = ChatState & ChatActions \ No newline at end of file +export type ChatStore = ChatState & ChatActions diff --git a/stores/chat/utils.ts b/stores/chat/utils.ts index 3b0b107ac4..a09566fb4a 100644 --- a/stores/chat/utils.ts +++ b/stores/chat/utils.ts @@ -21,13 +21,13 @@ export const calculateBlockPosition = ( ySpacing = 150 ) => { const blocksCount = Object.keys(existingBlocks).length - + // Calculate position based on existing blocks and current action index const row = Math.floor((blocksCount + index) / 5) // 5 blocks per row const col = (blocksCount + index) % 5 - + return { - x: startX + (col * xSpacing), - y: startY + (row * ySpacing) + x: startX + col * xSpacing, + y: startY + row * ySpacing, } -} \ No newline at end of file +} diff --git a/stores/workflow/utils.ts b/stores/workflow/utils.ts index 685c27ca7e..ca94872197 100644 --- a/stores/workflow/utils.ts +++ b/stores/workflow/utils.ts @@ -6,7 +6,10 @@ import { Edge } from 'reactflow' * @param startNode - Starting node for cycle detection * @returns Array of all unique cycles found in the graph */ -export function detectCycle(edges: Edge[], startNode: string): { hasCycle: boolean; paths: string[][] } { +export function detectCycle( + edges: Edge[], + startNode: string +): { hasCycle: boolean; paths: string[][] } { const visited = new Set() const recursionStack = new Set() const allCycles: string[][] = [] @@ -18,9 +21,7 @@ export function detectCycle(edges: Edge[], startNode: string): { hasCycle: boole currentPath.push(node) // Get all neighbors of current node - const neighbors = edges - .filter(edge => edge.source === node) - .map(edge => edge.target) + const neighbors = edges.filter((edge) => edge.source === node).map((edge) => edge.target) for (const neighbor of neighbors) { if (!recursionStack.has(neighbor)) { @@ -48,6 +49,6 @@ export function detectCycle(edges: Edge[], startNode: string): { hasCycle: boole return { hasCycle: allCycles.length > 0, - paths: allCycles + paths: allCycles, } -} \ No newline at end of file +}