From 490ba4eb87c1f0b4925a9085f023740bff418bb6 Mon Sep 17 00:00:00 2001 From: Waleed Latif Date: Mon, 10 Mar 2025 16:38:31 -0700 Subject: [PATCH] improvement: actually follow json-schema format for response format and make structured output enforced for providers --- .../connection-blocks/connection-blocks.tsx | 37 ++++++++- app/w/[id]/hooks/use-block-connections.ts | 40 ++++++++- blocks/blocks/agent.ts | 82 ++++++++----------- components/ui/tag-dropdown.tsx | 63 +++++++++++--- providers/anthropic/index.ts | 64 +++++++++++++-- providers/cerebras/index.ts | 5 +- providers/google/index.ts | 5 +- providers/groq/index.ts | 5 +- providers/openai/index.ts | 10 ++- providers/types.ts | 8 +- providers/xai/index.ts | 35 +------- 11 files changed, 238 insertions(+), 116 deletions(-) diff --git a/app/w/[id]/components/workflow-block/components/connection-blocks/connection-blocks.tsx b/app/w/[id]/components/workflow-block/components/connection-blocks/connection-blocks.tsx index 899f9c1f55..90e19f3d35 100644 --- a/app/w/[id]/components/workflow-block/components/connection-blocks/connection-blocks.tsx +++ b/app/w/[id]/components/workflow-block/components/connection-blocks/connection-blocks.tsx @@ -44,18 +44,45 @@ export function ConnectionBlocks({ blockId, setIsConnecting }: ConnectionBlocksP setIsConnecting(false) } + // Helper function to extract fields from JSON Schema + const extractFieldsFromSchema = (connection: ConnectedBlock): ResponseField[] => { + // Handle legacy format with fields array + if (connection.responseFormat?.fields) { + return connection.responseFormat.fields + } + + // Handle new JSON Schema format + const schema = connection.responseFormat?.schema || connection.responseFormat + // Safely check if schema and properties exist + if ( + !schema || + typeof schema !== 'object' || + !('properties' in schema) || + typeof schema.properties !== 'object' + ) { + return [] + } + return Object.entries(schema.properties).map(([name, prop]: [string, any]) => ({ + name, + type: Array.isArray(prop) ? 'array' : prop.type || 'string', + description: prop.description, + })) + } + return (
{incomingConnections.map((connection) => (
{Array.isArray(connection.outputType) ? ( + // Handle array of field names connection.outputType.map((fieldName) => { - const field = connection.responseFormat?.fields.find( - (f: ResponseField) => f.name === fieldName - ) || { + // Try to find field in response format + const fields = extractFieldsFromSchema(connection) + const field = fields.find((f) => f.name === fieldName) || { name: fieldName, type: 'string', } + return ( {connection.name.replace(/\s+/g, '').toLowerCase()} - .{connection.outputType} + + {typeof connection.outputType === 'string' ? `.${connection.outputType}` : ''} +
)} diff --git a/app/w/[id]/hooks/use-block-connections.ts b/app/w/[id]/hooks/use-block-connections.ts index ac8c5df77b..3a7b993f62 100644 --- a/app/w/[id]/hooks/use-block-connections.ts +++ b/app/w/[id]/hooks/use-block-connections.ts @@ -14,10 +14,42 @@ export interface ConnectedBlock { outputType: string | string[] name: string responseFormat?: { - fields: Field[] + // Support both formats + fields?: Field[] + name?: string + schema?: { + type: string + properties: Record + required?: string[] + } } } +// Helper function to extract fields from JSON Schema +function extractFieldsFromSchema(schema: any): Field[] { + if (!schema || typeof schema !== 'object') { + return [] + } + + // Handle legacy format with fields array + if (Array.isArray(schema.fields)) { + return schema.fields + } + + // Handle new JSON Schema format + const schemaObj = schema.schema || schema + if (!schemaObj || !schemaObj.properties || typeof schemaObj.properties !== 'object') { + return [] + } + + // Extract fields from schema properties + return Object.entries(schemaObj.properties).map(([name, prop]: [string, any]) => ({ + name, + type: prop.type || 'string', + description: prop.description, + })) +} + export function useBlockConnections(blockId: string) { const { edges, blocks } = useWorkflowStore( (state) => ({ @@ -43,7 +75,7 @@ export function useBlockConnections(blockId: string) { responseFormat = typeof responseFormatValue === 'string' && responseFormatValue ? JSON.parse(responseFormatValue) - : undefined + : responseFormatValue // Handle case where it's already an object } catch (e) { console.error('Failed to parse response format:', e) responseFormat = undefined @@ -55,8 +87,8 @@ export function useBlockConnections(blockId: string) { type: 'string', })) - // If we have a valid response format, use its fields as the output types - const outputFields = responseFormat?.fields || defaultOutputs + // Extract fields from the response format using our helper function + const outputFields = responseFormat ? extractFieldsFromSchema(responseFormat) : defaultOutputs return { id: sourceBlock.id, diff --git a/blocks/blocks/agent.ts b/blocks/blocks/agent.ts index 78b4d37a74..b32f301199 100644 --- a/blocks/blocks/agent.ts +++ b/blocks/blocks/agent.ts @@ -81,6 +81,7 @@ export const AgentBlock: BlockConfig = { title: 'Response Format', type: 'code', layout: 'full', + placeholder: `Enter JSON schema...`, }, ], tools: { @@ -115,63 +116,46 @@ export const AgentBlock: BlockConfig = { type: 'json', required: false, description: - 'Define the expected response format. If not provided, returns plain text content.', + 'Define the expected response format using JSON Schema. If not provided, returns plain text content.', schema: { type: 'object', properties: { - fields: { - type: 'array', - items: { - type: 'object', + name: { + type: 'string', + description: 'A name for your schema (optional)', + }, + schema: { + type: 'object', + description: 'The JSON Schema definition', + properties: { + type: { + type: 'string', + enum: ['object'], + description: 'Must be "object" for a valid JSON Schema', + }, properties: { - name: { - type: 'string', - description: 'Name of the field', - }, - type: { - type: 'string', - enum: ['string', 'number', 'boolean', 'array', 'object'], - description: 'Type of the field', - }, - isArray: { - type: 'boolean', - description: 'Whether this field contains multiple values', - }, - items: { - type: 'object', - description: 'Schema for array items (required when isArray is true)', - properties: { - type: { - type: 'string', - enum: ['string', 'number', 'boolean', 'object'], - }, - properties: { - type: 'array', - items: { - type: 'object', - properties: { - name: { type: 'string' }, - type: { - type: 'string', - enum: ['string', 'number', 'boolean'], - }, - }, - required: ['name', 'type'], - }, - }, - }, - }, - description: { - type: 'string', - description: 'Description of what this field represents', - }, + type: 'object', + description: 'Object containing property definitions', + }, + required: { + type: 'array', + items: { type: 'string' }, + description: 'Array of required property names', + }, + additionalProperties: { + type: 'boolean', + description: 'Whether additional properties are allowed', }, - required: ['name', 'type'], - additionalProperties: false, }, + required: ['type', 'properties'], + }, + strict: { + type: 'boolean', + description: 'Whether to enforce strict schema validation', + default: true, }, }, - required: ['fields'], + required: ['schema'], }, }, temperature: { type: 'number', required: false }, diff --git a/components/ui/tag-dropdown.tsx b/components/ui/tag-dropdown.tsx index 937118ad3a..9ae4f66c0b 100644 --- a/components/ui/tag-dropdown.tsx +++ b/components/ui/tag-dropdown.tsx @@ -30,6 +30,33 @@ interface TagDropdownProps { style?: React.CSSProperties } +// Add a helper function to extract fields from JSON Schema +const extractFieldsFromSchema = (responseFormat: any): Field[] => { + if (!responseFormat) return [] + + // Handle legacy format with fields array + if (Array.isArray(responseFormat.fields)) { + return responseFormat.fields + } + + // Handle new JSON Schema format + const schema = responseFormat.schema || responseFormat + if ( + !schema || + typeof schema !== 'object' || + !('properties' in schema) || + typeof schema.properties !== 'object' + ) { + return [] + } + + return Object.entries(schema.properties).map(([name, prop]: [string, any]) => ({ + name, + type: Array.isArray(prop) ? 'array' : prop.type || 'string', + description: prop.description, + })) +} + export const TagDropdown: React.FC = ({ visible, onSelect, @@ -109,13 +136,18 @@ export const TagDropdown: React.FC = ({ const responseFormatValue = useSubBlockStore .getState() .getValue(activeSourceBlockId, 'responseFormat') - if (typeof responseFormatValue === 'string' && responseFormatValue) { - const responseFormat = JSON.parse(responseFormatValue) - if (responseFormat?.fields) { - return { - tags: responseFormat.fields.map( - (field: Field) => `${normalizedBlockName}.${field.name}` - ), + if (responseFormatValue) { + const responseFormat = + typeof responseFormatValue === 'string' + ? JSON.parse(responseFormatValue) + : responseFormatValue + + if (responseFormat) { + const fields = extractFieldsFromSchema(responseFormat) + if (fields.length > 0) { + return { + tags: fields.map((field: Field) => `${normalizedBlockName}.${field.name}`), + } } } } @@ -144,12 +176,17 @@ export const TagDropdown: React.FC = ({ const responseFormatValue = useSubBlockStore .getState() .getValue(edge.source, 'responseFormat') - if (typeof responseFormatValue === 'string' && responseFormatValue) { - const responseFormat = JSON.parse(responseFormatValue) - if (responseFormat?.fields) { - return responseFormat.fields.map( - (field: Field) => `${normalizedBlockName}.${field.name}` - ) + if (responseFormatValue) { + const responseFormat = + typeof responseFormatValue === 'string' + ? JSON.parse(responseFormatValue) + : responseFormatValue + + if (responseFormat) { + const fields = extractFieldsFromSchema(responseFormat) + if (fields.length > 0) { + return fields.map((field: Field) => `${normalizedBlockName}.${field.name}`) + } } } } catch (e) { diff --git a/providers/anthropic/index.ts b/providers/anthropic/index.ts index 2375988c09..fa55a0d094 100644 --- a/providers/anthropic/index.ts +++ b/providers/anthropic/index.ts @@ -100,19 +100,73 @@ export const anthropicProvider: ProviderConfig = { // If response format is specified, add strict formatting instructions if (request.responseFormat) { - systemPrompt = `${systemPrompt}\n\nIMPORTANT RESPONSE FORMAT INSTRUCTIONS: + // Get the schema from the response format + const schema = request.responseFormat.schema || request.responseFormat + + // Build a system prompt for structured output based on the JSON schema + let schemaInstructions = '' + + if (schema && schema.properties) { + // Create a template of the expected JSON structure + const jsonTemplate = Object.entries(schema.properties).reduce( + (acc: Record, [key, prop]: [string, any]) => { + let exampleValue + const propType = prop.type || 'string' + + // Generate appropriate example values based on type + switch (propType) { + case 'string': + exampleValue = '"value"' + break + case 'number': + exampleValue = '0' + break + case 'boolean': + exampleValue = 'true' + break + case 'array': + exampleValue = '[]' + break + case 'object': + exampleValue = '{}' + break + default: + exampleValue = '"value"' + } + + acc[key] = exampleValue + return acc + }, + {} + ) + + // Generate field descriptions + const fieldDescriptions = Object.entries(schema.properties) + .map(([key, prop]: [string, any]) => { + const type = prop.type || 'string' + const description = prop.description ? `: ${prop.description}` : '' + return `${key} (${type})${description}` + }) + .join('\n') + + // Format the JSON template as a string + const jsonTemplateStr = JSON.stringify(jsonTemplate, null, 2) + + schemaInstructions = ` +IMPORTANT RESPONSE FORMAT INSTRUCTIONS: 1. Your response must be EXACTLY in this format, with no additional fields: -{ -${request.responseFormat.fields.map((field) => ` "${field.name}": ${field.type === 'string' ? '"value"' : field.type === 'array' ? '[]' : field.type === 'object' ? '{}' : field.type === 'number' ? '0' : 'true/false'}`).join(',\n')} -} +${jsonTemplateStr} Field descriptions: -${request.responseFormat.fields.map((field) => `${field.name} (${field.type})${field.description ? `: ${field.description}` : ''}`).join('\n')} +${fieldDescriptions} 2. DO NOT include any explanatory text before or after the JSON 3. DO NOT wrap the response in an array 4. DO NOT add any fields not specified in the schema 5. Your response MUST be valid JSON and include all the specified fields with their correct types` + } + + systemPrompt = `${systemPrompt}${schemaInstructions}` } // Build the request payload diff --git a/providers/cerebras/index.ts b/providers/cerebras/index.ts index 0b3cedb11d..60422451c1 100644 --- a/providers/cerebras/index.ts +++ b/providers/cerebras/index.ts @@ -67,7 +67,10 @@ export const cerebrasProvider: ProviderConfig = { // Add response format for structured output if specified if (request.responseFormat) { - payload.response_format = { type: 'json_object' } + payload.response_format = { + type: 'json_schema', + schema: request.responseFormat.schema || request.responseFormat, + } } // Add tools if provided diff --git a/providers/google/index.ts b/providers/google/index.ts index 0936fdd6c2..e43ca6350f 100644 --- a/providers/google/index.ts +++ b/providers/google/index.ts @@ -69,7 +69,10 @@ export const googleProvider: ProviderConfig = { // Add response format for structured output if specified if (request.responseFormat) { - payload.response_format = { type: 'json_object' } + payload.response_format = { + type: 'json_schema', + schema: request.responseFormat.schema || request.responseFormat, + } } // Add tools if provided diff --git a/providers/groq/index.ts b/providers/groq/index.ts index dab4092d70..e5b1d1072f 100644 --- a/providers/groq/index.ts +++ b/providers/groq/index.ts @@ -68,7 +68,10 @@ export const groqProvider: ProviderConfig = { // Add response format for structured output if specified if (request.responseFormat) { - payload.response_format = { type: 'json_object' } + payload.response_format = { + type: 'json_schema', + schema: request.responseFormat.schema || request.responseFormat, + } } // Add tools if provided diff --git a/providers/openai/index.ts b/providers/openai/index.ts index ea6012fb9a..0370a4729d 100644 --- a/providers/openai/index.ts +++ b/providers/openai/index.ts @@ -68,7 +68,15 @@ export const openaiProvider: ProviderConfig = { // Add response format for structured output if specified if (request.responseFormat) { - payload.response_format = { type: 'json_object' } + // Use OpenAI's JSON schema format + payload.response_format = { + type: 'json_schema', + json_schema: { + name: request.responseFormat.name || 'response_schema', + schema: request.responseFormat.schema || request.responseFormat, + strict: request.responseFormat.strict !== false, + }, + } } // Add tools if provided diff --git a/providers/types.ts b/providers/types.ts index 4c68011772..626320c8ae 100644 --- a/providers/types.ts +++ b/providers/types.ts @@ -86,11 +86,9 @@ export interface ProviderRequest { apiKey: string messages?: Message[] responseFormat?: { - fields: Array<{ - name: string - type: 'string' | 'number' | 'boolean' | 'array' | 'object' - description?: string - }> + name: string + schema: any + strict?: boolean } local_execution?: boolean } diff --git a/providers/xai/index.ts b/providers/xai/index.ts index aaefe1a5a6..d9790f800d 100644 --- a/providers/xai/index.ts +++ b/providers/xai/index.ts @@ -64,38 +64,9 @@ export const xAIProvider: ProviderConfig = { payload.response_format = { type: 'json_schema', json_schema: { - name: 'structured_response', - schema: { - type: 'object', - properties: request.responseFormat.fields.reduce( - (acc, field) => ({ - ...acc, - [field.name]: { - type: - field.type === 'array' - ? 'array' - : field.type === 'object' - ? 'object' - : field.type === 'number' - ? 'number' - : field.type === 'boolean' - ? 'boolean' - : 'string', - description: field.description || '', - ...(field.type === 'array' && { - items: { type: 'string' }, - }), - ...(field.type === 'object' && { - additionalProperties: true, - }), - }, - }), - {} - ), - required: request.responseFormat.fields.map((f) => f.name), - additionalProperties: false, - }, - strict: true, + name: request.responseFormat.name || 'structured_response', + schema: request.responseFormat.schema || request.responseFormat, + strict: request.responseFormat.strict !== false, }, }