From 5cf1d7a1ec4b9ec002b880d1ba9b366eda10d7c7 Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Wed, 29 Jul 2026 17:14:35 +0300 Subject: [PATCH 1/3] Refactoring --- crates/nagare/src/chat/mod.rs | 23 ++++++++--------------- crates/uzu/examples/tool_calls.rs | 4 ++-- crates/uzu/tests/tool/calls.rs | 5 +---- 3 files changed, 11 insertions(+), 21 deletions(-) diff --git a/crates/nagare/src/chat/mod.rs b/crates/nagare/src/chat/mod.rs index 23ddd4917..0e0575161 100644 --- a/crates/nagare/src/chat/mod.rs +++ b/crates/nagare/src/chat/mod.rs @@ -162,25 +162,18 @@ impl ChatSession { self.instance.lock().await.peak_memory_usage() } - pub async fn add_tool_function( + pub async fn add_tool( &mut self, - function: impl Into, + descriptor: impl Into, ) -> Result<(), ChatSessionError> { - self.add_tool_function_definitions(vec![function.into()]).await + self.add_tools(vec![descriptor.into()]).await } - pub async fn add_tool_descriptor( + pub async fn add_tools( &mut self, - definition: ToolDescriptor, + descriptors: Vec, ) -> Result<(), ChatSessionError> { - self.add_tool_function_definitions(vec![definition]).await - } - - pub async fn add_tool_function_definitions( - &mut self, - definitions: Vec, - ) -> Result<(), ChatSessionError> { - if definitions.is_empty() { + if descriptors.is_empty() { return Ok(()); } let Some(registry) = self.tool_registry.as_ref() else { @@ -193,8 +186,8 @@ impl ChatSession { // an incoming name must displace a stale caller-provided declaration too let owned_namespaces = { let mut registry_guard = registry.lock().await; - for def in definitions { - registry_guard.add_function(def) + for desc in descriptors { + registry_guard.add_function(desc) } registry_guard.get_namespaces() }; diff --git a/crates/uzu/examples/tool_calls.rs b/crates/uzu/examples/tool_calls.rs index b56a38631..367ab609e 100644 --- a/crates/uzu/examples/tool_calls.rs +++ b/crates/uzu/examples/tool_calls.rs @@ -35,9 +35,9 @@ async fn main() -> Result<(), Box> { } let mut session = engine.chat(model, ChatConfig::default()).await?; - session.add_tool_function(get_current_location).await?; + session.add_tool(get_current_location).await?; session - .add_tool_descriptor(uzu_tool_closure! { + .add_tool(uzu_tool_closure! { /// Returns temperature in provided location get_current_temperature: | /// Latitude in decimal degrees. diff --git a/crates/uzu/tests/tool/calls.rs b/crates/uzu/tests/tool/calls.rs index f81a07289..b33a94481 100644 --- a/crates/uzu/tests/tool/calls.rs +++ b/crates/uzu/tests/tool/calls.rs @@ -117,10 +117,7 @@ async fn run_tool_calls_test( println!("Prompt: {}", case.prompt); let mut session = engine.chat(model.clone(), ChatConfig::default()).await.unwrap(); - session - .add_tool_function_definitions(case.tools.iter().map(|definition| definition()).collect()) - .await - .unwrap(); + session.add_tools(case.tools.iter().map(|definition| definition()).collect()).await.unwrap(); let mut messages = Vec::new(); if with_system_message { From e43acc2bc6c6fdcddd69bf56f9b1cee54b820eac Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Thu, 30 Jul 2026 16:45:51 +0300 Subject: [PATCH 2/3] Add tool calls for TypeScript --- bindings/typescript/examples/toolCalls.ts | 82 +++++++++++++ bindings/typescript/package.json | 2 +- bindings/typescript/pnpm-lock.yaml | 12 +- bindings/typescript/src/index.ts | 1 + bindings/typescript/src/napi/index.d.mts | 6 + bindings/typescript/src/napi/index.d.ts | 6 + bindings/typescript/src/napi/index.js | 1 + bindings/typescript/src/napi/index.mjs | 1 + bindings/typescript/src/tool.ts | 120 +++++++++++++++++++ bindings/typescript/tests/tool.test.ts | 78 +++++++++++++ crates/nagare/src/chat/bindings_napi.rs | 133 ++++++++++++++++++++++ crates/nagare/src/chat/mod.rs | 2 + 12 files changed, 437 insertions(+), 7 deletions(-) create mode 100644 bindings/typescript/examples/toolCalls.ts create mode 100644 bindings/typescript/src/tool.ts create mode 100644 bindings/typescript/tests/tool.test.ts create mode 100644 crates/nagare/src/chat/bindings_napi.rs diff --git a/bindings/typescript/examples/toolCalls.ts b/bindings/typescript/examples/toolCalls.ts new file mode 100644 index 000000000..d209c4d7c --- /dev/null +++ b/bindings/typescript/examples/toolCalls.ts @@ -0,0 +1,82 @@ +import { + ChatConfig, + ChatMessage, + ChatReplyConfig, + Engine, + EngineConfig, + SamplingMethodGreedy, + SamplingPolicyCustom, + uzuToolFunction, +} from '@trymirai/uzu'; +import * as z from 'zod'; + + +const Coordinate = z.object({ + latitude: z.number().describe('Latitude in decimal degrees.'), + longitude: z.number().describe('Longitude in decimal degrees.'), +}); + +type Coordinate = z.infer; + + +const getCurrentLocation = uzuToolFunction({ + name: 'get_location', + description: 'Return the current location in coordinates', + parameters: z.object({}), + returns: Coordinate, + handler: (): Coordinate => ({ + latitude: 51.5074, + longitude: -0.1278, + }), +}); + + +async function calculateCurrentTemperature({latitude, longitude}: Coordinate): Promise { + if (!Number.isFinite(Math.hypot(latitude, longitude))) { + throw new RangeError('Coordinates must be finite'); + } + return 25; +} + +const getCurrentTemperature = uzuToolFunction({ + name: 'get_current_temperature', + description: 'Return the temperature at the provided coordinates', + parameters: Coordinate, + returns: z.number(), + handler: calculateCurrentTemperature, +}); + + +async function main() { + const engine = await Engine.create(EngineConfig.create()); + const model = await engine.model('mlx-community/Qwen3.5-9B-MLX-8bit'); + if (!model) { + throw new Error('Model not found'); + } + + for await (const update of await engine.download(model)) { + console.log('Download progress:', update.progress); + } + + const session = await engine.chat(model, ChatConfig.create()); + await session.addTool(getCurrentLocation); + await session.addTool(getCurrentTemperature); + + const messages = [ + ChatMessage.system().withText('You are a helpful assistant'), + ChatMessage.user().withText('What temperature is it now at my location?'), + ]; + const config = ChatReplyConfig.create().withSamplingPolicy( + new SamplingPolicyCustom(new SamplingMethodGreedy()), + ); + const replies = await session.reply(messages, config); + const message = replies[replies.length - 1]?.message; + if (message) { + console.log('Reasoning:', message.reasoning ?? ''); + console.log('Text:', message.text ?? ''); + } +} + +main().catch((error: unknown) => { + console.error(error); +}); diff --git a/bindings/typescript/package.json b/bindings/typescript/package.json index 2decc50d0..19f1cab2a 100644 --- a/bindings/typescript/package.json +++ b/bindings/typescript/package.json @@ -54,7 +54,7 @@ "publint": "^0.2.12", "ts-jest": "^29.1.0", "ts-node": "^10.5.0", - "tsc-multi": "https://github.com/stainless-api/tsc-multi/releases/download/v1.1.9/tsc-multi.tgz", + "tsc-multi": "https://github.com/stainless-api/tsc-multi/releases/download/v1.1.11/tsc-multi.tgz", "tsconfig-paths": "^4.0.0", "tslib": "^2.8.1", "typescript": "5.8.3", diff --git a/bindings/typescript/pnpm-lock.yaml b/bindings/typescript/pnpm-lock.yaml index 0925b8977..efc4a6eb2 100644 --- a/bindings/typescript/pnpm-lock.yaml +++ b/bindings/typescript/pnpm-lock.yaml @@ -63,8 +63,8 @@ importers: specifier: ^10.5.0 version: 10.9.2(@swc/core@1.13.5)(@types/node@20.19.19)(typescript@5.8.3) tsc-multi: - specifier: https://github.com/stainless-api/tsc-multi/releases/download/v1.1.9/tsc-multi.tgz - version: https://github.com/stainless-api/tsc-multi/releases/download/v1.1.9/tsc-multi.tgz(typescript@5.8.3) + specifier: https://github.com/stainless-api/tsc-multi/releases/download/v1.1.11/tsc-multi.tgz + version: https://github.com/stainless-api/tsc-multi/releases/download/v1.1.11/tsc-multi.tgz(typescript@5.8.3) tsconfig-paths: specifier: ^4.0.0 version: 4.2.0 @@ -4098,10 +4098,10 @@ packages: '@swc/wasm': optional: true - tsc-multi@https://github.com/stainless-api/tsc-multi/releases/download/v1.1.9/tsc-multi.tgz: + tsc-multi@https://github.com/stainless-api/tsc-multi/releases/download/v1.1.11/tsc-multi.tgz: resolution: - { tarball: https://github.com/stainless-api/tsc-multi/releases/download/v1.1.9/tsc-multi.tgz } - version: 1.1.9 + { tarball: https://github.com/stainless-api/tsc-multi/releases/download/v1.1.11/tsc-multi.tgz } + version: 1.1.11 engines: { node: '>=14' } hasBin: true peerDependencies: @@ -7017,7 +7017,7 @@ snapshots: optionalDependencies: '@swc/core': 1.13.5 - tsc-multi@https://github.com/stainless-api/tsc-multi/releases/download/v1.1.9/tsc-multi.tgz(typescript@5.8.3): + tsc-multi@https://github.com/stainless-api/tsc-multi/releases/download/v1.1.11/tsc-multi.tgz(typescript@5.8.3): dependencies: debug: 4.4.3 fast-glob: 3.3.3 diff --git a/bindings/typescript/src/index.ts b/bindings/typescript/src/index.ts index aa9a87e51..1f0a258c1 100644 --- a/bindings/typescript/src/index.ts +++ b/bindings/typescript/src/index.ts @@ -1 +1,2 @@ export * from './napi/index'; +export * from './tool'; diff --git a/bindings/typescript/src/napi/index.d.mts b/bindings/typescript/src/napi/index.d.mts index f8dad3205..db5e78aff 100644 --- a/bindings/typescript/src/napi/index.d.mts +++ b/bindings/typescript/src/napi/index.d.mts @@ -1,6 +1,8 @@ /* auto-generated by NAPI-RS */ /* eslint-disable */ export declare class ChatSession { + addTool(tool: NativeTool): Promise + addTools(tools: Array): Promise get state(): Promise get messages(): Promise> reset(): Promise @@ -35,6 +37,10 @@ export declare class ClassificationSession { classify(input: Array): Promise } +export declare class NativeTool { + constructor(definition: ToolFunction, invokeJson: (arg0: string, arg1: string) => Promise, cancel: (arg: string) => void) +} + export declare class TextToSpeechSession { get state(): Promise synthesize(input: string): Promise diff --git a/bindings/typescript/src/napi/index.d.ts b/bindings/typescript/src/napi/index.d.ts index f8dad3205..db5e78aff 100644 --- a/bindings/typescript/src/napi/index.d.ts +++ b/bindings/typescript/src/napi/index.d.ts @@ -1,6 +1,8 @@ /* auto-generated by NAPI-RS */ /* eslint-disable */ export declare class ChatSession { + addTool(tool: NativeTool): Promise + addTools(tools: Array): Promise get state(): Promise get messages(): Promise> reset(): Promise @@ -35,6 +37,10 @@ export declare class ClassificationSession { classify(input: Array): Promise } +export declare class NativeTool { + constructor(definition: ToolFunction, invokeJson: (arg0: string, arg1: string) => Promise, cancel: (arg: string) => void) +} + export declare class TextToSpeechSession { get state(): Promise synthesize(input: string): Promise diff --git a/bindings/typescript/src/napi/index.js b/bindings/typescript/src/napi/index.js index d125d56d4..d646fefb3 100644 --- a/bindings/typescript/src/napi/index.js +++ b/bindings/typescript/src/napi/index.js @@ -513,6 +513,7 @@ module.exports.ChatSessionStream = nativeBinding.ChatSessionStream module.exports.ChatSessionStreamChunkError = nativeBinding.ChatSessionStreamChunkError module.exports.ChatSessionStreamChunkReplies = nativeBinding.ChatSessionStreamChunkReplies module.exports.ClassificationSession = nativeBinding.ClassificationSession +module.exports.NativeTool = nativeBinding.NativeTool module.exports.TextToSpeechSession = nativeBinding.TextToSpeechSession module.exports.TextToSpeechSessionStream = nativeBinding.TextToSpeechSessionStream module.exports.TextToSpeechSessionStreamChunkError = nativeBinding.TextToSpeechSessionStreamChunkError diff --git a/bindings/typescript/src/napi/index.mjs b/bindings/typescript/src/napi/index.mjs index ffff55107..6f020d6ba 100644 --- a/bindings/typescript/src/napi/index.mjs +++ b/bindings/typescript/src/napi/index.mjs @@ -10,6 +10,7 @@ export const ChatSessionStream = cjs.ChatSessionStream; export const ChatSessionStreamChunkError = cjs.ChatSessionStreamChunkError; export const ChatSessionStreamChunkReplies = cjs.ChatSessionStreamChunkReplies; export const ClassificationSession = cjs.ClassificationSession; +export const NativeTool = cjs.NativeTool; export const TextToSpeechSession = cjs.TextToSpeechSession; export const TextToSpeechSessionStream = cjs.TextToSpeechSessionStream; export const TextToSpeechSessionStreamChunkError = cjs.TextToSpeechSessionStreamChunkError; diff --git a/bindings/typescript/src/tool.ts b/bindings/typescript/src/tool.ts new file mode 100644 index 000000000..b24d4e3ba --- /dev/null +++ b/bindings/typescript/src/tool.ts @@ -0,0 +1,120 @@ +import * as z from 'zod'; + +import { NativeTool, ToolFunction, Value } from './napi/index'; + +export interface UzuToolContext { + readonly signal: AbortSignal; +} + +export interface UzuToolInvokeOptions { + readonly signal?: AbortSignal; +} + +export interface UzuToolFunctionOptions { + readonly name: string; + readonly description?: string; + readonly parameters: Parameters; + readonly returns: ResultSchema; + readonly handler: ( + parameters: z.output, + context: UzuToolContext, + ) => z.input | Promise>; +} + +export class UzuToolFunction< + Parameters extends z.ZodObject, + ResultSchema extends z.ZodType, +> extends NativeTool { + readonly name: string; + readonly description: string; + readonly parameters: Parameters; + readonly returns: ResultSchema; + readonly parametersSchema: Record; + readonly returnSchema: Record; + readonly handler: UzuToolFunctionOptions['handler']; + + constructor(options: UzuToolFunctionOptions) { + const name = options.name.trim(); + if (!name) { + throw new TypeError('tool name must not be empty'); + } + + const description = options.description ?? ''; + const parametersSchema = z.toJSONSchema(options.parameters); + const returnSchema = z.toJSONSchema(options.returns); + const activeInvocations = new Map(); + const cancelledInvocations = new Set(); + + const invokeJson = async (argumentsJson: string, invocationId: string): Promise => { + const controller = new AbortController(); + activeInvocations.set(invocationId, controller); + if (cancelledInvocations.delete(invocationId)) { + controller.abort(); + } + + try { + const rawArguments: unknown = JSON.parse(argumentsJson); + const parameters = await options.parameters.parseAsync(rawArguments); + const rawResult = await options.handler(parameters, { + signal: controller.signal, + }); + const result = await options.returns.parseAsync(rawResult); + return serializeResult(result); + } finally { + activeInvocations.delete(invocationId); + cancelledInvocations.delete(invocationId); + } + }; + const cancel = (invocationId: string): void => { + const controller = activeInvocations.get(invocationId); + if (controller) { + controller.abort(); + } else { + cancelledInvocations.add(invocationId); + } + }; + + super( + new ToolFunction( + name, + description, + new Value(JSON.stringify(parametersSchema)), + new Value(JSON.stringify(returnSchema)), + ), + invokeJson, + cancel, + ); + + this.name = name; + this.description = description; + this.parameters = options.parameters; + this.returns = options.returns; + this.parametersSchema = parametersSchema; + this.returnSchema = returnSchema; + this.handler = options.handler; + } + + async invoke( + input: z.input, + options: UzuToolInvokeOptions = {}, + ): Promise> { + const parameters = await this.parameters.parseAsync(input); + const signal = options.signal ?? new AbortController().signal; + const rawResult = await this.handler(parameters, { signal }); + return this.returns.parseAsync(rawResult); + } +} + +export function uzuToolFunction( + options: UzuToolFunctionOptions, +): UzuToolFunction { + return new UzuToolFunction(options); +} + +function serializeResult(result: unknown): string { + const json = JSON.stringify(result === undefined ? null : result); + if (json === undefined) { + throw new TypeError('tool result must be JSON serializable'); + } + return json; +} diff --git a/bindings/typescript/tests/tool.test.ts b/bindings/typescript/tests/tool.test.ts new file mode 100644 index 000000000..521f8c69a --- /dev/null +++ b/bindings/typescript/tests/tool.test.ts @@ -0,0 +1,78 @@ +import { ChatSession, UzuToolFunction, uzuToolFunction } from '@trymirai/uzu'; +import * as z from 'zod'; + +test('tool factory builds schemas and invokes a typed handler', async () => { + const tool = uzuToolFunction({ + name: 'add', + description: 'Add two numbers', + parameters: z.object({ + left: z.number().describe('Left operand.'), + right: z.number().describe('Right operand.'), + }), + returns: z.number(), + handler: ({ left, right }) => Promise.resolve(left + right), + }); + + expect(tool).toBeInstanceOf(UzuToolFunction); + expect(tool.name).toBe('add'); + expect(tool.description).toBe('Add two numbers'); + expect(tool.parametersSchema).toMatchObject({ + type: 'object', + properties: { + left: { + type: 'number', + description: 'Left operand.', + }, + right: { + type: 'number', + description: 'Right operand.', + }, + }, + required: ['left', 'right'], + additionalProperties: false, + }); + expect(tool.returnSchema).toMatchObject({ + type: 'number', + }); + await expect(tool.invoke({ left: 2, right: 3 })).resolves.toBe(5); +}); + +test('tool invocation validates arguments and results', async () => { + const tool = uzuToolFunction({ + name: 'length', + parameters: z.object({ + value: z.string(), + }), + returns: z.number(), + handler: ({ value }) => value.length, + }); + + await expect(tool.invoke({ value: 'uzu' })).resolves.toBe(3); + await expect(tool.invoke({ value: 3 } as never)).rejects.toBeInstanceOf(z.ZodError); + + const invalidResult = uzuToolFunction({ + name: 'invalid_result', + parameters: z.object({}), + returns: z.number(), + handler: () => 'not a number' as never, + }); + await expect(invalidResult.invoke({})).rejects.toBeInstanceOf(z.ZodError); +}); + +test('tool handler receives the invocation abort signal', async () => { + const controller = new AbortController(); + const tool = uzuToolFunction({ + name: 'is_cancelled', + parameters: z.object({}), + returns: z.boolean(), + handler: (_parameters, context) => context.signal.aborted, + }); + + controller.abort(); + await expect(tool.invoke({}, { signal: controller.signal })).resolves.toBe(true); +}); + +test('chat session exposes tool registration methods', () => { + expect(typeof ChatSession.prototype.addTool).toBe('function'); + expect(typeof ChatSession.prototype.addTools).toBe('function'); +}); diff --git a/crates/nagare/src/chat/bindings_napi.rs b/crates/nagare/src/chat/bindings_napi.rs new file mode 100644 index 000000000..c9fac8db9 --- /dev/null +++ b/crates/nagare/src/chat/bindings_napi.rs @@ -0,0 +1,133 @@ +use std::{ + future::Future, + pin::Pin, + sync::Arc, + task::{Context, Poll}, +}; + +use napi::{ + Status, + bindgen_prelude::{AsyncBlock, AsyncBlockBuilder, ClassInstance, Env, FnArgs, Function, Promise}, + threadsafe_function::{ThreadsafeFunction, ThreadsafeFunctionCallMode}, +}; +use shoji::types::basic::{ToolFunction, Value}; +use uuid::Uuid; + +use super::ChatSession; +use crate::tool::func_def::{ErrorFuture, ToolDescriptor}; + +type InvokeArguments = FnArgs<(String, String)>; +type InvokeFunction = ThreadsafeFunction, InvokeArguments, Status, false, true>; +type CancelFunction = ThreadsafeFunction; +type JavaScriptFuture = Pin> + Send>>; + +struct JavaScriptInvocation { + future: JavaScriptFuture, + cancellation: Option<(Arc, String)>, +} + +impl Future for JavaScriptInvocation { + type Output = napi::Result; + + fn poll( + mut self: Pin<&mut Self>, + context: &mut Context<'_>, + ) -> Poll { + let result = self.future.as_mut().poll(context); + if result.is_ready() { + self.cancellation = None; + } + result + } +} + +impl Drop for JavaScriptInvocation { + fn drop(&mut self) { + let Some((cancel, invocation_id)) = self.cancellation.take() else { + return; + }; + let _ = cancel.call(invocation_id, ThreadsafeFunctionCallMode::NonBlocking); + } +} + +#[napi_derive::napi] +pub struct NativeTool { + descriptor: ToolDescriptor, +} + +#[napi_derive::napi] +impl NativeTool { + #[napi(constructor)] + pub fn new( + definition: ToolFunction, + invoke_json: Function<'_, FnArgs<(String, String)>, Promise>, + cancel: Function<'_, String, ()>, + ) -> napi::Result { + let invoke_json = Arc::new(invoke_json.build_threadsafe_function().weak::().build()?); + let cancel = Arc::new(cancel.build_threadsafe_function().weak::().build()?); + let descriptor = ToolDescriptor::new( + definition.name, + definition.description, + definition.parameters, + definition.return_definition, + Box::new(move |arguments| { + let invoke_json = invoke_json.clone(); + let cancel = cancel.clone(); + Box::new(call_javascript_tool(invoke_json, cancel, arguments)) + }), + ); + Ok(Self { + descriptor, + }) + } +} + +#[napi_derive::napi] +impl ChatSession { + #[napi(js_name = "addTool")] + pub fn add_tool_bindings_napi( + &self, + tool: ClassInstance<'_, NativeTool>, + env: Env, + ) -> napi::Result> { + let descriptor = tool.descriptor.clone(); + let mut session = self.clone(); + AsyncBlockBuilder::new(async move { session.add_tool(descriptor).await.map_err(Into::into) }).build(&env) + } + + #[napi(js_name = "addTools")] + pub fn add_tools_bindings_napi( + &self, + tools: Vec>, + env: Env, + ) -> napi::Result> { + let descriptors = tools.iter().map(|tool| tool.descriptor.clone()).collect(); + let mut session = self.clone(); + AsyncBlockBuilder::new(async move { session.add_tools(descriptors).await.map_err(Into::into) }).build(&env) + } +} + +async fn call_javascript_tool( + invoke_json: Arc, + cancel: Arc, + arguments: Value, +) -> Result { + let invocation_id = Uuid::new_v4().to_string(); + let invocation_id_for_call = invocation_id.clone(); + let future = Box::pin(async move { + let arguments = InvokeArguments::from((arguments.json, invocation_id_for_call)); + let promise = invoke_json.call_async_catch(arguments).await?; + promise.await + }); + let invocation = JavaScriptInvocation { + future, + cancellation: Some((cancel, invocation_id)), + }; + let json = invocation.await.map_err(javascript_error)?; + let value = serde_json::from_str::(&json)?; + Ok(Value::from(value)) +} + +fn javascript_error(error: napi::Error) -> ErrorFuture { + error.to_string().into() +} diff --git a/crates/nagare/src/chat/mod.rs b/crates/nagare/src/chat/mod.rs index 0e0575161..8942a2648 100644 --- a/crates/nagare/src/chat/mod.rs +++ b/crates/nagare/src/chat/mod.rs @@ -1,3 +1,5 @@ +#[cfg(feature = "bindings-napi")] +mod bindings_napi; mod error; pub mod message; pub mod token; From f053cfbfe6fc079b1876d807f48bd5dad0e4a5aa Mon Sep 17 00:00:00 2001 From: Aleksandr Golokoz <5617556+agolokoz@users.noreply.github.com> Date: Thu, 30 Jul 2026 17:10:52 +0300 Subject: [PATCH 3/3] Fix the cancellation ID leak --- bindings/typescript/src/tool.ts | 34 +++++++++++--- .../tests/tool-cancellation.test.ts | 45 +++++++++++++++++++ 2 files changed, 74 insertions(+), 5 deletions(-) create mode 100644 bindings/typescript/tests/tool-cancellation.test.ts diff --git a/bindings/typescript/src/tool.ts b/bindings/typescript/src/tool.ts index b24d4e3ba..4ce0a4131 100644 --- a/bindings/typescript/src/tool.ts +++ b/bindings/typescript/src/tool.ts @@ -2,6 +2,8 @@ import * as z from 'zod'; import { NativeTool, ToolFunction, Value } from './napi/index'; +const CANCELLED_INVOCATION_RETENTION_MS = 30_000; + export interface UzuToolContext { readonly signal: AbortSignal; } @@ -43,12 +45,22 @@ export class UzuToolFunction< const parametersSchema = z.toJSONSchema(options.parameters); const returnSchema = z.toJSONSchema(options.returns); const activeInvocations = new Map(); - const cancelledInvocations = new Set(); + const cancelledInvocations = new Map>(); + + const takeCancellation = (invocationId: string): boolean => { + const expiration = cancelledInvocations.get(invocationId); + if (expiration === undefined) { + return false; + } + clearTimeout(expiration); + cancelledInvocations.delete(invocationId); + return true; + }; const invokeJson = async (argumentsJson: string, invocationId: string): Promise => { const controller = new AbortController(); activeInvocations.set(invocationId, controller); - if (cancelledInvocations.delete(invocationId)) { + if (takeCancellation(invocationId)) { controller.abort(); } @@ -62,16 +74,28 @@ export class UzuToolFunction< return serializeResult(result); } finally { activeInvocations.delete(invocationId); - cancelledInvocations.delete(invocationId); + takeCancellation(invocationId); } }; const cancel = (invocationId: string): void => { const controller = activeInvocations.get(invocationId); if (controller) { controller.abort(); - } else { - cancelledInvocations.add(invocationId); + return; + } + + // Invocation and cancellation use separate nonblocking native callbacks, so + // cancellation can arrive before invocation starts. Retain it briefly for + // that case, but expire IDs from callbacks that arrive after completion. + const expiration = setTimeout(() => { + cancelledInvocations.delete(invocationId); + }, CANCELLED_INVOCATION_RETENTION_MS); + expiration.unref?.(); + const previousExpiration = cancelledInvocations.get(invocationId); + if (previousExpiration !== undefined) { + clearTimeout(previousExpiration); } + cancelledInvocations.set(invocationId, expiration); }; super( diff --git a/bindings/typescript/tests/tool-cancellation.test.ts b/bindings/typescript/tests/tool-cancellation.test.ts new file mode 100644 index 000000000..a93c2fede --- /dev/null +++ b/bindings/typescript/tests/tool-cancellation.test.ts @@ -0,0 +1,45 @@ +import * as z from 'zod'; + +type InvokeJson = (argumentsJson: string, invocationId: string) => Promise; +type Cancel = (invocationId: string) => void; + +jest.mock('../src/napi/index', () => ({ + NativeTool: class { + readonly testInvokeJson: InvokeJson; + readonly testCancel: Cancel; + + constructor(_definition: unknown, invokeJson: InvokeJson, cancel: Cancel) { + this.testInvokeJson = invokeJson; + this.testCancel = cancel; + } + }, + ToolFunction: class {}, + Value: class {}, +})); + +import { uzuToolFunction } from '../src/tool'; + +interface TestNativeTool { + readonly testInvokeJson: InvokeJson; + readonly testCancel: Cancel; +} + +test('retains pre-start cancellations but expires late cancellations', async () => { + jest.useFakeTimers(); + const tool = uzuToolFunction({ + name: 'is_cancelled', + parameters: z.object({}), + returns: z.boolean(), + handler: (_parameters, context) => context.signal.aborted, + }) as unknown as TestNativeTool; + + tool.testCancel('before-start'); + await expect(tool.testInvokeJson('{}', 'before-start')).resolves.toBe('true'); + + await expect(tool.testInvokeJson('{}', 'already-finished')).resolves.toBe('false'); + tool.testCancel('already-finished'); + jest.runOnlyPendingTimers(); + + await expect(tool.testInvokeJson('{}', 'already-finished')).resolves.toBe('false'); + jest.useRealTimers(); +});