Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 2 additions & 4 deletions apps/content/docs/integrations/ai-sdk.md
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,7 @@ const getWeatherTool = implementTool(getWeatherContract, {
location,
temperature: 72 + Math.floor(Math.random() * 21) - 10,
}),
// ...add any additional configuration or overrides here
})
```

Expand Down Expand Up @@ -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.
:::
23 changes: 19 additions & 4 deletions packages/ai-sdk/src/tool.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,7 @@ describe('implementTool', () => {
})

describe('createTool', () => {
const abortSignal = (new AbortController()).signal
const base = os.$meta<AiSdkToolMeta>({})

const inputSchema = z.object({
Expand All @@ -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}!`,
Expand All @@ -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')
})
})
22 changes: 15 additions & 7 deletions packages/ai-sdk/src/tool.ts
Original file line number Diff line number Diff line change
@@ -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'

Expand Down Expand Up @@ -88,8 +88,6 @@ export function implementTool<TOutInput, TInOutput>(
* 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
Expand Down Expand Up @@ -147,9 +145,19 @@ export function createTool<
const options = resolveMaybeOptionalOptions(rest)

return implementTool(procedure, {
execute: (input: InferSchemaOutput<TInputSchema>) => {
return call(procedure, input as InferSchemaInput<TInputSchema>, 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<TInputSchema>,
{ signal: callingOptions.abortSignal, ...options },
) as Promise<InferSchemaInput<TOutputSchema>>
}) satisfies (Tool<InferSchemaOutput<TInputSchema>, InferSchemaInput<TOutputSchema>>['execute']),
Comment thread
dinwwwh marked this conversation as resolved.
...options,
} as any)
}
18 changes: 18 additions & 0 deletions packages/server/src/procedure-client.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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')
})
17 changes: 12 additions & 5 deletions packages/trpc/src/to-orpc-router.test.ts
Original file line number Diff line number Diff line change
@@ -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'
Expand Down Expand Up @@ -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 () => {
Expand Down
28 changes: 6 additions & 22 deletions packages/trpc/src/to-orpc-router.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down Expand Up @@ -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<unknown, unknown> {
function toStandardSchema(schema: undefined | Parser): undefined | ORPC.Schema<unknown, unknown> {
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
}
Loading