diff --git a/apps/content/docs/integrations/ai-sdk.md b/apps/content/docs/integrations/ai-sdk.md index cd86903ba..2f25e12e4 100644 --- a/apps/content/docs/integrations/ai-sdk.md +++ b/apps/content/docs/integrations/ai-sdk.md @@ -160,6 +160,7 @@ const getWeatherTool = implementTool(getWeatherContract, { location, temperature: 72 + Math.floor(Math.random() * 21) - 10, }), + // ...add any additional configuration or overrides here }) ``` @@ -210,13 +211,10 @@ const getWeatherProcedure = base const getWeatherTool = createTool(getWeatherProcedure, { context: {}, // provide initial context if needed + // ...add any additional configuration or overrides here }) ``` ::: warning The `createTool` helper requires a procedure with an `input` schema defined ::: - -::: warning -Validation occurs twice (once for the tool, once for the procedure call). So validation may fail if `inputSchema` or `outputSchema` transform the data into different shapes. -::: diff --git a/packages/ai-sdk/src/tool.test.ts b/packages/ai-sdk/src/tool.test.ts index 33594653d..c8e00020d 100644 --- a/packages/ai-sdk/src/tool.test.ts +++ b/packages/ai-sdk/src/tool.test.ts @@ -71,6 +71,7 @@ describe('implementTool', () => { }) describe('createTool', () => { + const abortSignal = (new AbortController()).signal const base = os.$meta({}) const inputSchema = z.object({ @@ -80,7 +81,7 @@ describe('createTool', () => { greeting: z.string().describe('Greeting message'), }) - it('can create a tool', () => { + it('can create a tool', async () => { const handler = vi.fn(async ({ input }) => { return { greeting: `Hello, ${input.name}!`, @@ -103,13 +104,27 @@ describe('createTool', () => { expect(tool.outputSchema).toBe(outputSchema) expect(tool.description).toBe('Greet a person') - return expect( - (tool as any).execute({ name: 'Alice' }), - ).resolves.toEqual({ greeting: 'Hello, Alice!' }) + await expect((tool as any).execute({ name: 'Alice' }, { abortSignal })).resolves.toEqual({ greeting: 'Hello, Alice!' }) expect(handler).toHaveBeenCalledWith(expect.objectContaining({ + signal: abortSignal, input: { name: 'Alice' }, context: { authToken: 'auth-token' }, })) }) + + it('disable validation at oRPC level to avoid twice times validation', async () => { + const procedure + = base + .route({ + summary: 'Greet a person', + }) + .input(inputSchema) + .output(outputSchema) + .handler(({ input }) => input as any) + + const tool = createTool(procedure) + + await expect(tool.execute?.('invalid' as any, { abortSignal } as any)).resolves.toEqual('invalid') + }) }) diff --git a/packages/ai-sdk/src/tool.ts b/packages/ai-sdk/src/tool.ts index 90748aa15..ebbfc6599 100644 --- a/packages/ai-sdk/src/tool.ts +++ b/packages/ai-sdk/src/tool.ts @@ -1,9 +1,9 @@ import type { ClientOptions } from '@orpc/client' import type { AnySchema, ContractProcedure, ErrorMap, InferSchemaInput, InferSchemaOutput, Meta, Schema } from '@orpc/contract' -import type { Context, CreateProcedureClientOptions, Procedure } from '@orpc/server' +import type { Context, CreateProcedureClientOptions } from '@orpc/server' import type { MaybeOptionalOptions, SetOptional } from '@orpc/shared' import type { Tool } from 'ai' -import { call } from '@orpc/server' +import { call, Procedure } from '@orpc/server' import { resolveMaybeOptionalOptions } from '@orpc/shared' import { tool } from 'ai' @@ -88,8 +88,6 @@ export function implementTool( * by leveraging existing procedure definitions. * * @warning Requires a contract with an `input` schema defined. - * @warning Validation occurs twice (once for the tool, once for the procedure call). - * So validation may fail if inputSchema or outputSchema transform the data into different shapes. * * @example * ```ts @@ -147,9 +145,19 @@ export function createTool< const options = resolveMaybeOptionalOptions(rest) return implementTool(procedure, { - execute: (input: InferSchemaOutput) => { - return call(procedure, input as InferSchemaInput, options) - }, + execute: ((input, callingOptions) => { + const disabledValidation = new Procedure({ + ...procedure['~orpc'], + inputValidationIndex: Number.NaN, // disable input validation + outputValidationIndex: Number.NaN, // disable output validation + }) + + return call( + disabledValidation, + input as InferSchemaInput, + { signal: callingOptions.abortSignal, ...options }, + ) as Promise> + }) satisfies (Tool, InferSchemaInput>['execute']), ...options, } as any) } diff --git a/packages/server/src/procedure-client.test.ts b/packages/server/src/procedure-client.test.ts index ef6acec84..9071509f3 100644 --- a/packages/server/src/procedure-client.test.ts +++ b/packages/server/src/procedure-client.test.ts @@ -652,3 +652,21 @@ it('has helper `output` in meta', async () => { expect(preMid1).toReturnWith(Promise.resolve({ output: { val: '99990' }, context: {} })) }) + +it('support disable input/output validation by setting validation index to NaN', async () => { + const procedure = new Procedure({ + inputSchema: schema, + outputSchema: schema, + errorMap: {}, + route: {}, + meta: {}, + handler: ({ input }) => input, + middlewares: [], + inputValidationIndex: Number.NaN, + outputValidationIndex: Number.NaN, + }) + + const client = createProcedureClient(procedure) + + await expect(client('invalid' as any)).resolves.toEqual('invalid') +}) diff --git a/packages/trpc/src/to-orpc-router.test.ts b/packages/trpc/src/to-orpc-router.test.ts index 659b07d83..97670726c 100644 --- a/packages/trpc/src/to-orpc-router.test.ts +++ b/packages/trpc/src/to-orpc-router.test.ts @@ -1,4 +1,4 @@ -import { call, createRouterClient, getEventMeta, isLazy, isProcedure, ORPCError, unlazy } from '@orpc/server' +import { call, createRouterClient, getEventMeta, isLazy, isProcedure, ORPCError, Procedure, unlazy } from '@orpc/server' import { isAsyncIteratorObject } from '@orpc/shared' import { tracked, TRPCError } from '@trpc/server' import * as z from 'zod' @@ -31,15 +31,22 @@ describe('toORPCRouter', async () => { expect(await unlazy(orpcRouter.lazy.lazy.throw)).toEqual({ default: expect.toSatisfy(isProcedure) }) }) - it('with disabled input/output', async () => { + it('with input/output schema and validation happen inside handler only', async () => { expect((orpcRouter as any).ping['~orpc'].inputSchema['~standard'].vendor).toBe('zod') expect((orpcRouter as any).ping['~orpc'].inputSchema._def).toBe(inputSchema._def) + expect((orpcRouter as any).ping['~orpc'].inputValidationIndex).toBe(Number.NaN) // input validation is disabled + expect((orpcRouter as any).ping['~orpc'].outputSchema['~standard'].vendor).toBe('zod') expect((orpcRouter as any).ping['~orpc'].outputSchema._def).toBe(outputSchema._def) + expect((orpcRouter as any).ping['~orpc'].outputValidationIndex).toBe(Number.NaN) // output validation is disabled + + const withoutHandlerProcedure = new Procedure({ + ...(orpcRouter as any).ping['~orpc'], + handler: async ({ input }) => input, + }) - const invalidValue = 'INVALID' - expect((orpcRouter as any).ping['~orpc'].inputSchema['~standard'].validate(invalidValue)).toEqual({ value: invalidValue }) - expect((orpcRouter as any).ping['~orpc'].outputSchema['~standard'].validate(invalidValue)).toEqual({ value: invalidValue }) + await expect(call(withoutHandlerProcedure, 'invalid')).resolves.toEqual('invalid') // validation not happen at oRPC level + await expect(call((orpcRouter as any).ping, 'invalid')).rejects.toThrow('Invalid input') // validation happen at tRPC level }) it('meta/route', async () => { diff --git a/packages/trpc/src/to-orpc-router.ts b/packages/trpc/src/to-orpc-router.ts index bfbf26142..743761732 100644 --- a/packages/trpc/src/to-orpc-router.ts +++ b/packages/trpc/src/to-orpc-router.ts @@ -89,12 +89,12 @@ function toORPCProcedure(procedure: AnyProcedure) { return new ORPC.Procedure({ errorMap: {}, meta: procedure._def.meta ?? {}, - inputValidationIndex: 0, - outputValidationIndex: 0, route: get(procedure._def.meta, ['route']) ?? {}, middlewares: [], - inputSchema: toDisabledStandardSchema(procedure._def.inputs.at(-1)), - outputSchema: toDisabledStandardSchema((procedure as any)._def.output), + inputSchema: toStandardSchema(procedure._def.inputs.at(-1)), + outputSchema: toStandardSchema((procedure._def as any).output), + inputValidationIndex: Number.NaN, // disable input validation + outputValidationIndex: Number.NaN, // disable output validation handler: async ({ context, signal, path, input, lastEventId }) => { try { const trpcInput = lastEventId !== undefined && (input === undefined || isObject(input)) @@ -151,26 +151,10 @@ function toORPCProcedure(procedure: AnyProcedure) { * Wraps a TRPC schema to disable validation in the ORPC context. * This is necessary because tRPC procedure calling already validates the input/output, */ -function toDisabledStandardSchema(schema: undefined | Parser): undefined | ORPC.Schema { +function toStandardSchema(schema: undefined | Parser): undefined | ORPC.Schema { if (!isTypescriptObject(schema) || !('~standard' in schema) || !isTypescriptObject(schema['~standard'])) { return undefined } - return new Proxy(schema as any, { - get: (target, prop) => { - if (prop === '~standard') { - return new Proxy(target['~standard'], { - get: (target, prop) => { - if (prop === 'validate') { - return (value: any) => ({ value }) - } - - return Reflect.get(target, prop, target) - }, - }) - } - - return Reflect.get(target, prop, target) - }, - }) + return schema as any }