fix(blocks): fixed router and evaluator block, added autofill & tests (#274)

* fixed router block, refined system prompt, added tests

* add autofill for evaluator

* fixed evaluator block to execute server-side, added tests
This commit is contained in:
Waleed Latif
2025-04-17 12:23:25 -07:00
committed by GitHub
parent de581f89dd
commit 382046973a
6 changed files with 556 additions and 326 deletions
@@ -6,9 +6,10 @@ import { useWorkflowStore } from '@/stores/workflows/workflow/store'
import { getProviderFromModel } from '@/providers/utils'
/**
* Helper to handle API key auto-fill for agent blocks
* Helper to handle API key auto-fill for provider-based blocks
* Used for agent, router, evaluator, and any other blocks that use LLM providers
*/
function handleAgentBlockApiKey(
function handleProviderBasedApiKey(
blockId: string,
subBlockId: string,
modelValue: string | null | undefined,
@@ -119,8 +120,8 @@ function storeApiKeyValue(
subBlockStore.unmarkParamAsCleared(blockId, 'apiKey')
}
// For agent blocks, store the API key under the provider name
if (blockType === 'agent' && modelValue) {
// For provider-based blocks, store the API key under the provider name
if ((blockType === 'agent' || blockType === 'router' || blockType === 'evaluator') && modelValue) {
const provider = getProviderFromModel(modelValue)
if (provider && provider !== 'ollama') {
subBlockStore.setToolParam(provider, 'apiKey', String(newValue))
@@ -179,68 +180,13 @@ export function useSubBlockValue<T = any>(
blockId ? state.getValue(blockId, 'model') : null
)
// Compute the modelValue after the hook call
const modelValue = blockType === 'agent' ? (modelSubBlockValue as string) : null
// Determine if this is a provider-based block type
const isProviderBasedBlock = blockType === 'agent' || blockType === 'router' || blockType === 'evaluator'
// When model changes for an agent block's API key, immediately check if we need to clear it
useEffect(() => {
// Only run for agent blocks with API key fields when model changes
if (blockType === 'agent' && isApiKey && modelValue !== prevModelRef.current) {
// Update the previous model reference
prevModelRef.current = modelValue
// Compute the modelValue based on block type
const modelValue = isProviderBasedBlock ? (modelSubBlockValue as string) : null
// For agent blocks, always clear the field if needed
// But only fill with saved values if auto-fill is enabled
if (modelValue) {
const provider = getProviderFromModel(modelValue)
// Skip if we couldn't determine a provider
if (!provider || provider === 'ollama') return
const subBlockStore = useSubBlockStore.getState()
// Check if there's a saved value for this provider
const savedValue = subBlockStore.resolveToolParamValue(provider, 'apiKey', blockId)
if (savedValue && savedValue !== '' && isAutoFillEnvVarsEnabled) {
// Only auto-fill if the feature is enabled
subBlockStore.setValue(blockId, subBlockId, savedValue)
} else {
// Always clear immediately when switching to a model with no saved key
// or when auto-fill is disabled
subBlockStore.setValue(blockId, subBlockId, '')
}
}
}
}, [blockId, subBlockId, blockType, isApiKey, modelValue, isAutoFillEnvVarsEnabled, storeValue])
// When component mounts, check for existing API key in toolParamsStore
useEffect(() => {
// Skip autofill if the feature is disabled in settings
if (!isAutoFillEnvVarsEnabled) return
// Only process API key fields
if (!isApiKey) return
// Handle agent blocks differently, they need to use the model to determine provider
if (blockType === 'agent') {
handleAgentBlockApiKey(blockId, subBlockId, modelValue, storeValue)
} else {
// Normal handling for non-agent blocks
handleStandardBlockApiKey(blockId, subBlockId, blockType, storeValue)
}
}, [blockId, subBlockId, blockType, storeValue, isApiKey, isAutoFillEnvVarsEnabled, modelValue])
// Update the ref if the store value changes
// This ensures we're always working with the latest value
useEffect(() => {
// Use deep comparison for objects to prevent unnecessary updates
if (!isEqual(valueRef.current, storeValue)) {
valueRef.current = storeValue !== undefined ? storeValue : initialValue
}
}, [storeValue, initialValue])
// Set value function that handles deep equality for complex objects
// Hook to set a value in the subblock store
const setValue = useCallback(
(newValue: T) => {
// Use deep comparison to avoid unnecessary updates for complex objects
@@ -272,6 +218,71 @@ export function useSubBlockValue<T = any>(
[blockId, subBlockId, blockType, isApiKey, storeValue, triggerWorkflowUpdate, modelValue]
)
// Return the current value and setter
// Initialize valueRef on first render
useEffect(() => {
valueRef.current = storeValue !== undefined ? storeValue : initialValue
}, [])
// When component mounts, check for existing API key in toolParamsStore
useEffect(() => {
// Skip autofill if the feature is disabled in settings
if (!isAutoFillEnvVarsEnabled) return
// Only process API key fields
if (!isApiKey) return
// Handle different block types
if (isProviderBasedBlock) {
handleProviderBasedApiKey(blockId, subBlockId, modelValue, storeValue)
} else {
// Normal handling for non-provider blocks
handleStandardBlockApiKey(blockId, subBlockId, blockType, storeValue)
}
}, [blockId, subBlockId, blockType, storeValue, isApiKey, isAutoFillEnvVarsEnabled, modelValue, isProviderBasedBlock])
// Monitor for model changes in provider-based blocks
useEffect(() => {
// Only process API key fields in model-based blocks
if (!isApiKey || !isProviderBasedBlock) return
// Check if the model has changed
if (modelValue !== prevModelRef.current) {
// Update the previous model reference
prevModelRef.current = modelValue
// For provider-based blocks, always clear the field if needed
// But only fill with saved values if auto-fill is enabled
if (modelValue) {
const provider = getProviderFromModel(modelValue)
// Skip if we couldn't determine a provider
if (!provider || provider === 'ollama') return
const subBlockStore = useSubBlockStore.getState()
// Check if there's a saved value for this provider
const savedValue = subBlockStore.resolveToolParamValue(provider, 'apiKey', blockId)
if (savedValue && savedValue !== '' && isAutoFillEnvVarsEnabled) {
// Only auto-fill if the feature is enabled
subBlockStore.setValue(blockId, subBlockId, savedValue)
} else {
// Always clear immediately when switching to a model with no saved key
// or when auto-fill is disabled
subBlockStore.setValue(blockId, subBlockId, '')
}
}
}
}, [blockId, subBlockId, blockType, isApiKey, modelValue, isAutoFillEnvVarsEnabled, storeValue, isProviderBasedBlock])
// Update the ref if the store value changes
// This ensures we're always working with the latest value
useEffect(() => {
// Use deep comparison for objects to prevent unnecessary updates
if (!isEqual(valueRef.current, storeValue)) {
valueRef.current = storeValue !== undefined ? storeValue : initialValue
}
}, [storeValue, initialValue])
return [valueRef.current as T | null, setValue] as const
}
+3 -2
View File
@@ -55,7 +55,7 @@ ID: ${block.id}
Type: ${block.type}
Title: ${block.title}
Description: ${block.description}
Category: ${block.category}
System Prompt: ${JSON.stringify(block.subBlocks?.systemPrompt || '')}
Configuration: ${JSON.stringify(block.subBlocks, null, 2)}
${block.currentState ? `Current State: ${JSON.stringify(block.currentState, null, 2)}` : ''}
---`
@@ -64,7 +64,8 @@ ${block.currentState ? `Current State: ${JSON.stringify(block.currentState, null
Routing Instructions:
1. Analyze the input request carefully against each block's:
- Primary purpose (from description)
- Primary purpose (from title, description, and system prompt)
- Look for keywords in the system prompt that match the user's request
- Configuration settings
- Current state (if available)
- Processing capabilities
@@ -1,14 +1,13 @@
import '../../__test-utils__/mock-dependencies'
import { beforeEach, describe, expect, it, Mock, vi } from 'vitest'
import { BlockOutput } from '@/blocks/types'
import { executeProviderRequest } from '@/providers'
import { getProviderFromModel } from '@/providers/utils'
import { SerializedBlock } from '@/serializer/types'
import { ExecutionContext } from '../../types'
import { EvaluatorBlockHandler } from './evaluator-handler'
const mockGetProviderFromModel = getProviderFromModel as Mock
const mockExecuteProviderRequest = executeProviderRequest as Mock
const mockFetch = global.fetch as Mock
describe('EvaluatorBlockHandler', () => {
let handler: EvaluatorBlockHandler
@@ -37,11 +36,12 @@ describe('EvaluatorBlockHandler', () => {
workflowId: 'test-workflow-id',
blockStates: new Map(),
blockLogs: [],
metadata: {},
metadata: { duration: 0 },
environmentVariables: {},
decisions: { router: new Map(), condition: new Map() },
loopIterations: new Map(),
loopItems: new Map(),
completedLoops: new Set(),
executedBlocks: new Set(),
activeExecutionPath: new Set(),
}
@@ -51,12 +51,20 @@ describe('EvaluatorBlockHandler', () => {
// Default mock implementations
mockGetProviderFromModel.mockReturnValue('openai')
mockExecuteProviderRequest.mockResolvedValue({
content: JSON.stringify({ score1: 5, score2: 8 }),
model: 'mock-model',
tokens: { prompt: 50, completion: 10, total: 60 },
cost: 0.002,
timing: { total: 200 },
// Set up fetch mock to return a successful response
mockFetch.mockImplementation(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: JSON.stringify({ score1: 5, score2: 8 }),
model: 'mock-model',
tokens: { prompt: 50, completion: 10, total: 60 },
cost: 0.002,
timing: { total: 200 },
}),
})
})
})
@@ -77,27 +85,6 @@ describe('EvaluatorBlockHandler', () => {
temperature: 0.1,
}
const expectedProviderRequest = {
model: 'gpt-4o',
systemPrompt: expect.stringContaining(inputs.content),
responseFormat: {
name: 'evaluation_response',
schema: {
type: 'object',
properties: {
score1: { type: 'number' },
score2: { type: 'number' },
},
required: ['score1', 'score2'],
additionalProperties: false,
},
strict: true,
},
messages: [{ role: 'user', content: expect.stringContaining('Please evaluate the content') }],
temperature: 0.1,
apiKey: undefined,
}
const expectedOutput: BlockOutput = {
response: {
content: 'This is the content to evaluate.',
@@ -111,7 +98,36 @@ describe('EvaluatorBlockHandler', () => {
const result = await handler.execute(mockBlock, inputs, mockContext)
expect(mockGetProviderFromModel).toHaveBeenCalledWith('gpt-4o')
expect(mockExecuteProviderRequest).toHaveBeenCalledWith('openai', expectedProviderRequest)
expect(mockFetch).toHaveBeenCalledWith(
expect.any(String),
expect.objectContaining({
method: 'POST',
headers: expect.any(Object),
body: expect.any(String),
})
)
// Verify the request body contains the expected data
const fetchCallArgs = mockFetch.mock.calls[0]
const requestBody = JSON.parse(fetchCallArgs[1].body)
expect(requestBody).toMatchObject({
provider: 'openai',
model: 'gpt-4o',
systemPrompt: expect.stringContaining(inputs.content),
responseFormat: expect.objectContaining({
schema: {
type: 'object',
properties: {
score1: { type: 'number' },
score2: { type: 'number' },
},
required: ['score1', 'score2'],
additionalProperties: false,
}
}),
temperature: 0.1,
})
expect(result).toEqual(expectedOutput)
})
@@ -121,22 +137,28 @@ describe('EvaluatorBlockHandler', () => {
content: JSON.stringify(contentObj),
metrics: [{ name: 'clarity', description: 'Clarity score', range: { min: 1, max: 5 } }],
}
mockExecuteProviderRequest.mockResolvedValueOnce({
content: JSON.stringify({ clarity: 4 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: JSON.stringify({ clarity: 4 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
}),
})
})
await handler.execute(mockBlock, inputs, mockContext)
expect(mockExecuteProviderRequest).toHaveBeenCalledWith(
expect.any(String),
expect.objectContaining({
systemPrompt: expect.stringContaining(JSON.stringify(contentObj, null, 2)),
})
)
const fetchCallArgs = mockFetch.mock.calls[0]
const requestBody = JSON.parse(fetchCallArgs[1].body)
expect(requestBody).toMatchObject({
systemPrompt: expect.stringContaining(JSON.stringify(contentObj, null, 2)),
})
})
it('should process object content correctly', async () => {
@@ -147,22 +169,28 @@ describe('EvaluatorBlockHandler', () => {
{ name: 'completeness', description: 'Data completeness', range: { min: 0, max: 1 } },
],
}
mockExecuteProviderRequest.mockResolvedValueOnce({
content: JSON.stringify({ completeness: 1 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: JSON.stringify({ completeness: 1 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
}),
})
})
await handler.execute(mockBlock, inputs, mockContext)
expect(mockExecuteProviderRequest).toHaveBeenCalledWith(
expect.any(String),
expect.objectContaining({
systemPrompt: expect.stringContaining(JSON.stringify(contentObj, null, 2)),
})
)
const fetchCallArgs = mockFetch.mock.calls[0]
const requestBody = JSON.parse(fetchCallArgs[1].body)
expect(requestBody).toMatchObject({
systemPrompt: expect.stringContaining(JSON.stringify(contentObj, null, 2)),
})
})
it('should parse valid JSON response correctly', async () => {
@@ -170,12 +198,19 @@ describe('EvaluatorBlockHandler', () => {
content: 'Test content',
metrics: [{ name: 'quality', description: 'Quality score', range: { min: 1, max: 10 } }],
}
mockExecuteProviderRequest.mockResolvedValueOnce({
content: '```json\n{ "quality": 9 }\n```',
model: 'm',
tokens: {},
cost: 0,
timing: {},
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: '```json\n{ "quality": 9 }\n```',
model: 'm',
tokens: {},
cost: 0,
timing: {},
}),
})
})
const result = await handler.execute(mockBlock, inputs, mockContext)
@@ -188,12 +223,19 @@ describe('EvaluatorBlockHandler', () => {
content: 'Test content',
metrics: [{ name: 'score', description: 'Score', range: { min: 0, max: 5 } }],
}
mockExecuteProviderRequest.mockResolvedValueOnce({
content: 'Sorry, I cannot provide a score.',
model: 'm',
tokens: {},
cost: 0,
timing: {},
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: 'Sorry, I cannot provide a score.',
model: 'm',
tokens: {},
cost: 0,
timing: {},
}),
})
})
const result = await handler.execute(mockBlock, inputs, mockContext)
@@ -209,12 +251,19 @@ describe('EvaluatorBlockHandler', () => {
{ name: 'fluency', description: 'Flu', range: { min: 0, max: 1 } },
],
}
mockExecuteProviderRequest.mockResolvedValueOnce({
content: '{ "accuracy": 1, "fluency": invalid }',
model: 'm',
tokens: {},
cost: 0,
timing: {},
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: '{ "accuracy": 1, "fluency": invalid }',
model: 'm',
tokens: {},
cost: 0,
timing: {},
}),
})
})
const result = await handler.execute(mockBlock, inputs, mockContext)
@@ -227,12 +276,19 @@ describe('EvaluatorBlockHandler', () => {
content: 'Test',
metrics: [{ name: 'CamelCaseScore', description: 'Desc', range: { min: 0, max: 10 } }],
}
mockExecuteProviderRequest.mockResolvedValueOnce({
content: JSON.stringify({ camelcasescore: 7 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: JSON.stringify({ camelcasescore: 7 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
}),
})
})
const result = await handler.execute(mockBlock, inputs, mockContext)
@@ -248,12 +304,19 @@ describe('EvaluatorBlockHandler', () => {
{ name: 'missingScore', description: 'Desc2', range: { min: 0, max: 5 } },
],
}
mockExecuteProviderRequest.mockResolvedValueOnce({
content: JSON.stringify({ presentScore: 4 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: JSON.stringify({ presentScore: 4 }),
model: 'm',
tokens: {},
cost: 0,
timing: {},
}),
})
})
const result = await handler.execute(mockBlock, inputs, mockContext)
@@ -261,4 +324,19 @@ describe('EvaluatorBlockHandler', () => {
expect((result as any).response.presentscore).toBe(4)
expect((result as any).response.missingscore).toBe(0)
})
it('should handle server error responses', async () => {
const inputs = { content: 'Test error handling.' }
// Override fetch mock to return an error
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: false,
status: 500,
json: () => Promise.resolve({ error: 'Server error' }),
})
})
await expect(handler.execute(mockBlock, inputs, mockContext)).rejects.toThrow('Server error')
})
})
@@ -1,6 +1,5 @@
import { createLogger } from '@/lib/logs/console-logger'
import { BlockOutput } from '@/blocks/types'
import { executeProviderRequest } from '@/providers'
import { getProviderFromModel } from '@/providers/utils'
import { SerializedBlock } from '@/serializer/types'
import { BlockHandler, ExecutionContext } from '../../types'
@@ -102,129 +101,162 @@ export class EvaluatorBlockHandler implements BlockHandler {
'Evaluate the content and provide scores for each metric as JSON.'
}
// Make sure we force JSON output in the request
const response = await executeProviderRequest(providerId, {
model: model,
systemPrompt: systemPromptObj.systemPrompt,
responseFormat: systemPromptObj.responseFormat,
messages: [
{
role: 'user',
content:
'Please evaluate the content provided in the system prompt. Return ONLY a valid JSON with metric scores.',
},
],
temperature: inputs.temperature || 0,
apiKey: inputs.apiKey,
})
// Parse response content with robust error handling
let parsedContent: Record<string, any> = {}
try {
const contentStr = response.content.trim()
let jsonStr = ''
const baseUrl = process.env.NEXT_PUBLIC_APP_URL || ''
const url = new URL('/api/providers', baseUrl)
// Make sure we force JSON output in the request
const providerRequest = {
provider: providerId,
model: model,
systemPrompt: systemPromptObj.systemPrompt,
responseFormat: systemPromptObj.responseFormat,
context: JSON.stringify([
{
role: 'user',
content: 'Please evaluate the content provided in the system prompt. Return ONLY a valid JSON with metric scores.',
},
]),
temperature: inputs.temperature || 0,
apiKey: inputs.apiKey,
workflowId: context.workflowId,
}
const response = await fetch(url.toString(), {
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
body: JSON.stringify(providerRequest),
})
// Method 1: Extract content between first { and last }
const fullMatch = contentStr.match(/(\{[\s\S]*\})/) // Regex to find JSON structure
if (fullMatch) {
jsonStr = fullMatch[0]
}
// Method 2: Try to find and extract just the JSON part
else if (contentStr.includes('{') && contentStr.includes('}')) {
const startIdx = contentStr.indexOf('{')
const endIdx = contentStr.lastIndexOf('}') + 1
jsonStr = contentStr.substring(startIdx, endIdx)
}
// Method 3: Just use the raw content as a last resort
else {
jsonStr = contentStr
if (!response.ok) {
// Try to extract a helpful error message
let errorMessage = `Provider API request failed with status ${response.status}`
try {
const errorData = await response.json()
if (errorData.error) {
errorMessage = errorData.error
}
} catch (e) {
// If JSON parsing fails, use the original error message
}
throw new Error(errorMessage)
}
// Try to parse the extracted JSON
const result = await response.json()
// Parse response content with robust error handling
let parsedContent: Record<string, any> = {}
try {
parsedContent = JSON.parse(jsonStr)
} catch (parseError) {
logger.error('Failed to parse extracted JSON:', parseError)
throw new Error('Invalid JSON in response')
const contentStr = result.content.trim()
let jsonStr = ''
// Method 1: Extract content between first { and last }
const fullMatch = contentStr.match(/(\{[\s\S]*\})/) // Regex to find JSON structure
if (fullMatch) {
jsonStr = fullMatch[0]
}
// Method 2: Try to find and extract just the JSON part
else if (contentStr.includes('{') && contentStr.includes('}')) {
const startIdx = contentStr.indexOf('{')
const endIdx = contentStr.lastIndexOf('}') + 1
jsonStr = contentStr.substring(startIdx, endIdx)
}
// Method 3: Just use the raw content as a last resort
else {
jsonStr = contentStr
}
// Try to parse the extracted JSON
try {
parsedContent = JSON.parse(jsonStr)
} catch (parseError) {
logger.error('Failed to parse extracted JSON:', parseError)
throw new Error('Invalid JSON in response')
}
} catch (error) {
logger.error('Error parsing evaluator response:', error)
logger.error('Raw response content:', result.content)
// Fallback to empty object
parsedContent = {}
}
} catch (error) {
logger.error('Error parsing evaluator response:', error)
logger.error('Raw response content:', response.content)
// Fallback to empty object
parsedContent = {}
}
// Extract and process metric scores with proper validation
const metricScores: Record<string, any> = {}
// Extract and process metric scores with proper validation
const metricScores: Record<string, any> = {}
try {
// Ensure metrics is an array before processing
const validMetrics = Array.isArray(inputs.metrics) ? inputs.metrics : []
try {
// Ensure metrics is an array before processing
const validMetrics = Array.isArray(inputs.metrics) ? inputs.metrics : []
// If we have a successful parse, extract the metrics
if (Object.keys(parsedContent).length > 0) {
validMetrics.forEach((metric: any) => {
// Check if metric and name are valid before proceeding
if (!metric || !metric.name) {
logger.warn('Skipping invalid metric entry during score extraction:', metric)
return // Skip this iteration
}
const metricName = metric.name
const lowerCaseMetricName = metricName.toLowerCase()
// Try multiple possible ways the metric might be represented
if (parsedContent[metricName] !== undefined) {
metricScores[lowerCaseMetricName] = Number(parsedContent[metricName])
} else if (parsedContent[metricName.toLowerCase()] !== undefined) {
metricScores[lowerCaseMetricName] = Number(parsedContent[metricName.toLowerCase()])
} else if (parsedContent[metricName.toUpperCase()] !== undefined) {
metricScores[lowerCaseMetricName] = Number(parsedContent[metricName.toUpperCase()])
} else {
// Last resort - try to find any key that might contain this metric name
const matchingKey = Object.keys(parsedContent).find((key) => {
// Add check for key validity before calling toLowerCase()
return typeof key === 'string' && key.toLowerCase().includes(lowerCaseMetricName)
})
if (matchingKey) {
metricScores[lowerCaseMetricName] = Number(parsedContent[matchingKey])
} else {
logger.warn(`Metric "${metricName}" not found in LLM response`)
metricScores[lowerCaseMetricName] = 0
// If we have a successful parse, extract the metrics
if (Object.keys(parsedContent).length > 0) {
validMetrics.forEach((metric: any) => {
// Check if metric and name are valid before proceeding
if (!metric || !metric.name) {
logger.warn('Skipping invalid metric entry during score extraction:', metric)
return // Skip this iteration
}
}
})
} else {
// If we couldn't parse any content, set all metrics to 0
validMetrics.forEach((metric: any) => {
// Ensure metric and name are valid before setting default score
if (metric && metric.name) {
metricScores[metric.name.toLowerCase()] = 0
} else {
logger.warn('Skipping invalid metric entry when setting default scores:', metric)
}
})
const metricName = metric.name
const lowerCaseMetricName = metricName.toLowerCase()
// Try multiple possible ways the metric might be represented
if (parsedContent[metricName] !== undefined) {
metricScores[lowerCaseMetricName] = Number(parsedContent[metricName])
} else if (parsedContent[metricName.toLowerCase()] !== undefined) {
metricScores[lowerCaseMetricName] = Number(parsedContent[metricName.toLowerCase()])
} else if (parsedContent[metricName.toUpperCase()] !== undefined) {
metricScores[lowerCaseMetricName] = Number(parsedContent[metricName.toUpperCase()])
} else {
// Last resort - try to find any key that might contain this metric name
const matchingKey = Object.keys(parsedContent).find((key) => {
// Add check for key validity before calling toLowerCase()
return typeof key === 'string' && key.toLowerCase().includes(lowerCaseMetricName)
})
if (matchingKey) {
metricScores[lowerCaseMetricName] = Number(parsedContent[matchingKey])
} else {
logger.warn(`Metric "${metricName}" not found in LLM response`)
metricScores[lowerCaseMetricName] = 0
}
}
})
} else {
// If we couldn't parse any content, set all metrics to 0
validMetrics.forEach((metric: any) => {
// Ensure metric and name are valid before setting default score
if (metric && metric.name) {
metricScores[metric.name.toLowerCase()] = 0
} else {
logger.warn('Skipping invalid metric entry when setting default scores:', metric)
}
})
}
} catch (e) {
logger.error('Error extracting metric scores:', e)
}
} catch (e) {
logger.error('Error extracting metric scores:', e)
}
// Create result with metrics as direct fields for easy access
const result = {
response: {
content: inputs.content,
model: response.model,
tokens: {
prompt: response.tokens?.prompt || 0,
completion: response.tokens?.completion || 0,
total: response.tokens?.total || 0,
// Create result with metrics as direct fields for easy access
const outputResult = {
response: {
content: inputs.content,
model: result.model,
tokens: {
prompt: result.tokens?.prompt || 0,
completion: result.tokens?.completion || 0,
total: result.tokens?.total || 0,
},
...metricScores,
},
...metricScores,
},
}
}
return result
return outputResult
} catch (error) {
logger.error('Evaluator execution failed:', error)
throw error
}
}
}
@@ -2,7 +2,6 @@ import '../../__test-utils__/mock-dependencies'
import { beforeEach, describe, expect, it, Mock, Mocked, MockedClass, vi } from 'vitest'
import { generateRouterPrompt } from '@/blocks/blocks/router'
import { BlockOutput } from '@/blocks/types'
import { executeProviderRequest } from '@/providers'
import { getProviderFromModel } from '@/providers/utils'
import { SerializedBlock, SerializedWorkflow } from '@/serializer/types'
import { PathTracker } from '../../path'
@@ -11,8 +10,8 @@ import { RouterBlockHandler } from './router-handler'
const mockGenerateRouterPrompt = generateRouterPrompt as Mock
const mockGetProviderFromModel = getProviderFromModel as Mock
const mockExecuteProviderRequest = executeProviderRequest as Mock
const MockPathTracker = PathTracker as MockedClass<typeof PathTracker>
const mockFetch = global.fetch as Mock
describe('RouterBlockHandler', () => {
let handler: RouterBlockHandler
@@ -66,11 +65,12 @@ describe('RouterBlockHandler', () => {
workflowId: 'test-workflow-id',
blockStates: new Map(),
blockLogs: [],
metadata: {},
metadata: { duration: 0 },
environmentVariables: {},
decisions: { router: new Map(), condition: new Map() },
loopIterations: new Map(),
loopItems: new Map(),
completedLoops: new Set(),
executedBlocks: new Set(),
activeExecutionPath: new Set(),
workflow: mockWorkflow as SerializedWorkflow,
@@ -82,12 +82,20 @@ describe('RouterBlockHandler', () => {
// Default mock implementations
mockGetProviderFromModel.mockReturnValue('openai')
mockGenerateRouterPrompt.mockReturnValue('Generated System Prompt')
mockExecuteProviderRequest.mockResolvedValue({
content: 'target-block-1',
model: 'mock-model',
tokens: { prompt: 100, completion: 5, total: 105 },
cost: 0.003,
timing: { total: 300 },
// Set up fetch mock to return a successful response
mockFetch.mockImplementation(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: 'target-block-1',
model: 'mock-model',
tokens: { prompt: 100, completion: 5, total: 105 },
cost: 0.003,
timing: { total: 300 },
}),
})
})
})
@@ -110,7 +118,10 @@ describe('RouterBlockHandler', () => {
type: 'target',
title: 'Option A',
description: 'Choose A',
subBlocks: { p: 'a' },
subBlocks: {
p: 'a',
systemPrompt: ''
},
currentState: undefined,
},
{
@@ -118,19 +129,14 @@ describe('RouterBlockHandler', () => {
type: 'target',
title: 'Option B',
description: 'Choose B',
subBlocks: { p: 'b' },
subBlocks: {
p: 'b',
systemPrompt: ''
},
currentState: undefined,
},
]
const expectedProviderRequest = {
model: 'gpt-4o',
systemPrompt: 'Generated System Prompt',
messages: [{ role: 'user', content: 'Choose the best option.' }],
temperature: 0.5,
apiKey: undefined,
}
const expectedOutput: BlockOutput = {
response: {
content: 'Choose the best option.',
@@ -148,7 +154,26 @@ describe('RouterBlockHandler', () => {
expect(mockGenerateRouterPrompt).toHaveBeenCalledWith(inputs.prompt, expectedTargetBlocks)
expect(mockGetProviderFromModel).toHaveBeenCalledWith('gpt-4o')
expect(mockExecuteProviderRequest).toHaveBeenCalledWith('openai', expectedProviderRequest)
expect(mockFetch).toHaveBeenCalledWith(
expect.any(String),
expect.objectContaining({
method: 'POST',
headers: expect.any(Object),
body: expect.any(String),
})
)
// Verify the request body contains the expected data
const fetchCallArgs = mockFetch.mock.calls[0]
const requestBody = JSON.parse(fetchCallArgs[1].body)
expect(requestBody).toMatchObject({
provider: 'openai',
model: 'gpt-4o',
systemPrompt: 'Generated System Prompt',
context: JSON.stringify([{ role: 'user', content: 'Choose the best option.' }]),
temperature: 0.5,
})
expect(result).toEqual(expectedOutput)
})
@@ -160,17 +185,25 @@ describe('RouterBlockHandler', () => {
await expect(handler.execute(mockBlock, inputs, mockContext)).rejects.toThrow(
'Target block target-block-1 not found'
)
expect(mockExecuteProviderRequest).not.toHaveBeenCalled()
expect(mockFetch).not.toHaveBeenCalled()
})
it('should throw error if LLM response is not a valid target block ID', async () => {
const inputs = { prompt: 'Test' }
mockExecuteProviderRequest.mockResolvedValueOnce({
content: 'invalid-block-id',
model: 'm',
tokens: {},
cost: 0,
timing: {},
// Override fetch mock to return an invalid block ID
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: true,
json: () =>
Promise.resolve({
content: 'invalid-block-id',
model: 'mock-model',
tokens: {},
cost: 0,
timing: {},
}),
})
})
await expect(handler.execute(mockBlock, inputs, mockContext)).rejects.toThrow(
@@ -184,9 +217,27 @@ describe('RouterBlockHandler', () => {
await handler.execute(mockBlock, inputs, mockContext)
expect(mockGetProviderFromModel).toHaveBeenCalledWith('gpt-4o')
expect(mockExecuteProviderRequest).toHaveBeenCalledWith(
expect.any(String),
expect.objectContaining({ model: 'gpt-4o', temperature: 0 })
)
const fetchCallArgs = mockFetch.mock.calls[0]
const requestBody = JSON.parse(fetchCallArgs[1].body)
expect(requestBody).toMatchObject({
model: 'gpt-4o',
temperature: 0
})
})
it('should handle server error responses', async () => {
const inputs = { prompt: 'Test error handling.' }
// Override fetch mock to return an error
mockFetch.mockImplementationOnce(() => {
return Promise.resolve({
ok: false,
status: 500,
json: () => Promise.resolve({ error: 'Server error' }),
})
})
await expect(handler.execute(mockBlock, inputs, mockContext)).rejects.toThrow('Server error')
})
})
+89 -32
View File
@@ -1,7 +1,6 @@
import { createLogger } from '@/lib/logs/console-logger'
import { generateRouterPrompt } from '@/blocks/blocks/router'
import { BlockOutput } from '@/blocks/types'
import { executeProviderRequest } from '@/providers'
import { getProviderFromModel } from '@/providers/utils'
import { SerializedBlock } from '@/serializer/types'
import { PathTracker } from '../../path'
@@ -38,38 +37,78 @@ export class RouterBlockHandler implements BlockHandler {
const providerId = getProviderFromModel(routerConfig.model)
const response = await executeProviderRequest(providerId, {
model: routerConfig.model,
systemPrompt: generateRouterPrompt(routerConfig.prompt, targetBlocks),
messages: [{ role: 'user', content: routerConfig.prompt }],
temperature: routerConfig.temperature,
apiKey: routerConfig.apiKey,
})
const chosenBlockId = response.content.trim().toLowerCase()
const chosenBlock = targetBlocks?.find((b) => b.id === chosenBlockId)
if (!chosenBlock) {
throw new Error(`Invalid routing decision: ${chosenBlockId}`)
}
const tokens = response.tokens || { prompt: 0, completion: 0, total: 0 }
return {
response: {
content: inputs.prompt,
model: response.model,
tokens: {
prompt: tokens.prompt || 0,
completion: tokens.completion || 0,
total: tokens.total || 0,
try {
const baseUrl = process.env.NEXT_PUBLIC_APP_URL || ''
const url = new URL('/api/providers', baseUrl)
// Create the provider request with proper message formatting
const messages = [{ role: 'user', content: routerConfig.prompt }]
const systemPrompt = generateRouterPrompt(routerConfig.prompt, targetBlocks)
const providerRequest = {
provider: providerId,
model: routerConfig.model,
systemPrompt: systemPrompt,
context: JSON.stringify(messages),
temperature: routerConfig.temperature,
apiKey: routerConfig.apiKey,
workflowId: context.workflowId,
}
const response = await fetch(url.toString(), {
method: 'POST',
headers: {
'Content-Type': 'application/json',
},
selectedPath: {
blockId: chosenBlock.id,
blockType: chosenBlock.type || 'unknown',
blockTitle: chosenBlock.title || 'Untitled Block',
body: JSON.stringify(providerRequest),
})
if (!response.ok) {
// Try to extract a helpful error message
let errorMessage = `Provider API request failed with status ${response.status}`
try {
const errorData = await response.json()
if (errorData.error) {
errorMessage = errorData.error
}
} catch (e) {
// If JSON parsing fails, use the original error message
}
throw new Error(errorMessage)
}
const result = await response.json()
const chosenBlockId = result.content.trim().toLowerCase()
const chosenBlock = targetBlocks?.find((b) => b.id === chosenBlockId)
if (!chosenBlock) {
logger.error(`Invalid routing decision. Response content: "${result.content}", available blocks:`,
targetBlocks?.map(b => ({ id: b.id, title: b.title })) || []
)
throw new Error(`Invalid routing decision: ${chosenBlockId}`)
}
const tokens = result.tokens || { prompt: 0, completion: 0, total: 0 }
return {
response: {
content: inputs.prompt,
model: result.model,
tokens: {
prompt: tokens.prompt || 0,
completion: tokens.completion || 0,
total: tokens.total || 0,
},
selectedPath: {
blockId: chosenBlock.id,
blockType: chosenBlock.type || 'unknown',
blockTitle: chosenBlock.title || 'Untitled Block',
},
},
},
}
} catch (error) {
logger.error('Router execution failed:', error)
throw error
}
}
@@ -89,12 +128,30 @@ export class RouterBlockHandler implements BlockHandler {
if (!targetBlock) {
throw new Error(`Target block ${conn.target} not found`)
}
// Extract system prompt for agent blocks
let systemPrompt = ''
if (targetBlock.metadata?.id === 'agent') {
// Try to get system prompt from different possible locations
systemPrompt = targetBlock.config?.params?.systemPrompt ||
targetBlock.inputs?.systemPrompt ||
''
// If system prompt is still not found, check if we can extract it from inputs
if (!systemPrompt && targetBlock.inputs) {
systemPrompt = targetBlock.inputs.systemPrompt || ''
}
}
return {
id: targetBlock.id,
type: targetBlock.metadata?.id,
title: targetBlock.metadata?.name,
description: targetBlock.metadata?.description,
subBlocks: targetBlock.config.params,
subBlocks: {
...targetBlock.config.params,
systemPrompt: systemPrompt
},
currentState: context.blockStates.get(targetBlock.id)?.output,
}
})