From af959be605e6e2b52ab9b4952910cc769e63fa4e Mon Sep 17 00:00:00 2001 From: lforst <8118419+lforst@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:43:22 +0000 Subject: [PATCH] refactor(instrumentation)!: Separate API wrapping from tracing --- .agents/skills/instrumentation/SKILL.md | 23 +- AGENTS.md | 8 + .../template/app/api/test/route.ts | 30 +- js/src/auto-instrumentations/README.md | 65 +- .../orchestrion-js/transformer.ts | 18 +- .../orchestrion-js/transforms.ts | 31 +- js/src/cli/auto-instrumentation.test.ts | 39 +- js/src/debug-logger.test.ts | 39 +- js/src/global-instrumentation-hooks.test.ts | 675 ++------ js/src/global-instrumentation-hooks.ts | 798 +--------- js/src/instrumentation/README.md | 210 +-- .../core/channel-definitions.test.ts | 96 +- .../core/channel-definitions.ts | 382 +---- .../core/channel-tracing.test.ts | 295 ++-- .../instrumentation/core/channel-tracing.ts | 1008 +++++------- js/src/instrumentation/core/index.ts | 13 +- .../core/observe-result.test.ts | 58 + js/src/instrumentation/core/observe-result.ts | 44 + js/src/instrumentation/core/plugin.ts | 475 +----- js/src/instrumentation/core/tracing-types.ts | 39 + js/src/instrumentation/core/types.ts | 73 +- js/src/instrumentation/index.ts | 33 +- .../plugins/ai-sdk-channels.ts | 433 +++--- .../plugins/ai-sdk-plugin.streaming.test.ts | 96 +- .../plugins/ai-sdk-plugin.test.ts | 34 +- .../instrumentation/plugins/ai-sdk-plugin.ts | 1365 ++++++++++------- .../plugins/anthropic-channels.ts | 93 +- .../plugins/anthropic-plugin.test.ts | 5 +- .../plugins/anthropic-plugin.ts | 283 ++-- .../plugins/anthropic-sessions-plugin.test.ts | 34 +- .../plugins/bedrock-runtime-channels.ts | 37 +- .../plugins/bedrock-runtime-plugin.test.ts | 65 +- .../plugins/bedrock-runtime-plugin.ts | 76 +- .../plugins/claude-agent-sdk-channels.ts | 8 +- .../plugins/claude-agent-sdk-plugin.test.ts | 23 +- .../plugins/cloudflare-agents-channels.ts | 27 +- .../plugins/cloudflare-agents-plugin.test.ts | 45 +- .../plugins/cloudflare-agents-plugin.ts | 175 ++- .../plugins/cloudflare-ai-chat-channels.ts | 11 +- .../cloudflare-ai-chat-instrumentation.ts | 22 +- .../plugins/cloudflare-ai-chat-plugin.test.ts | 110 +- .../plugins/cloudflare-ai-chat-plugin.ts | 220 +-- .../plugins/cloudflare-think-channels.ts | 27 +- .../plugins/cloudflare-think-plugin.ts | 248 +-- .../plugins/cohere-channels.ts | 50 +- .../instrumentation/plugins/cohere-plugin.ts | 151 +- .../plugins/cursor-sdk-channels.ts | 74 +- .../plugins/cursor-sdk-plugin.test.ts | 49 +- .../plugins/cursor-sdk-plugin.ts | 392 ++--- .../plugins/elevenlabs-channels.ts | 25 +- .../plugins/elevenlabs-plugin.ts | 436 +++--- .../instrumentation/plugins/flue-channels.ts | 19 +- .../plugins/flue-plugin.test.ts | 61 +- js/src/instrumentation/plugins/flue-plugin.ts | 65 +- .../plugins/genkit-channels.ts | 89 +- .../plugins/genkit-plugin.test.ts | 19 +- .../instrumentation/plugins/genkit-plugin.ts | 432 +++--- .../plugins/github-copilot-channels.ts | 48 +- .../plugins/github-copilot-plugin.ts | 214 +-- .../plugins/google-adk-channels.ts | 63 +- .../plugins/google-adk-plugin.test.ts | 197 ++- .../plugins/google-adk-plugin.ts | 544 +++---- .../plugins/google-genai-channels.ts | 122 +- .../plugins/google-genai-plugin.test.ts | 123 +- .../plugins/google-genai-plugin.ts | 715 ++++----- .../plugins/google-generative-ai-channels.ts | 33 +- .../plugins/google-generative-ai-plugin.ts | 426 ++--- .../instrumentation/plugins/groq-channels.ts | 81 +- js/src/instrumentation/plugins/groq-plugin.ts | 234 +-- .../plugins/huggingface-channels.ts | 85 +- .../plugins/huggingface-plugin.ts | 197 ++- .../huggingface-transformers-channels.ts | 16 +- .../huggingface-transformers-plugin.ts | 141 +- .../plugins/instrumentation-names.test.ts | 13 +- .../plugins/langchain-channels.ts | 38 +- .../plugins/langchain-plugin.test.ts | 41 +- .../plugins/langchain-plugin.ts | 50 +- .../plugins/langgraph-sdk-channels.ts | 10 +- .../plugins/langsmith-channels.ts | 51 +- .../plugins/langsmith-plugin.test.ts | 225 +-- .../plugins/langsmith-plugin.ts | 98 +- .../plugins/mistral-channels.ts | 167 +- .../instrumentation/plugins/mistral-plugin.ts | 337 ++-- .../plugins/ollama-channels.ts | 52 +- .../instrumentation/plugins/ollama-plugin.ts | 91 +- .../plugins/openai-agents-channels.ts | 19 +- .../plugins/openai-agents-plugin.test.ts | 26 +- .../plugins/openai-agents-plugin.ts | 112 +- .../plugins/openai-channels.ts | 373 +++-- .../plugins/openai-codex-channels.ts | 44 +- .../plugins/openai-codex-plugin.test.ts | 29 +- .../plugins/openai-codex-plugin.ts | 173 ++- .../instrumentation/plugins/openai-media.ts | 359 ++--- .../instrumentation/plugins/openai-plugin.ts | 615 +++++--- .../plugins/openrouter-agent-channels.ts | 69 +- .../plugins/openrouter-agent-plugin.test.ts | 18 +- .../plugins/openrouter-agent-plugin.ts | 298 ++-- .../plugins/openrouter-channels.ts | 131 +- .../plugins/openrouter-plugin.test.ts | 14 +- .../plugins/openrouter-plugin.ts | 604 ++++---- .../plugins/pi-coding-agent-channels.ts | 10 +- .../plugins/pi-coding-agent-plugin.test.ts | 8 +- .../plugins/pi-coding-agent-plugin.ts | 21 +- .../plugins/strands-agent-sdk-channels.ts | 10 +- .../plugins/strands-agent-sdk-plugin.test.ts | 15 +- .../plugins/typesafe-channels.ts | 24 +- .../plugins/voyageai-channels.ts | 64 +- .../plugins/voyageai-plugin.ts | 165 +- js/src/instrumentation/registry.test.ts | 23 +- .../instrumentation/test-utils/invocation.ts | 61 + js/src/isomorph.ts | 24 +- js/src/openai-promise-utils.test.ts | 58 +- js/src/wrappers/ai-sdk/ai-sdk.ts | 44 +- .../wrappers/ai-sdk/harness-agent-context.ts | 29 +- js/src/wrappers/anthropic.ts | 30 +- js/src/wrappers/bedrock-runtime.ts | 11 +- .../claude-agent-sdk/claude-agent-sdk.ts | 2 +- js/src/wrappers/cloudflare-agent.test.ts | 39 +- js/src/wrappers/cloudflare-agent.ts | 8 +- js/src/wrappers/cloudflare-ai-chat.test.ts | 34 +- js/src/wrappers/cloudflare-think.test.ts | 34 +- js/src/wrappers/cloudflare-think.ts | 11 +- js/src/wrappers/cohere.ts | 21 +- js/src/wrappers/cursor-sdk.test.ts | 47 +- js/src/wrappers/cursor-sdk.ts | 39 +- js/src/wrappers/genkit.test.ts | 42 +- js/src/wrappers/genkit.ts | 56 +- js/src/wrappers/github-copilot.ts | 16 +- js/src/wrappers/google-adk.test.ts | 31 +- js/src/wrappers/google-adk.ts | 31 +- js/src/wrappers/google-genai.test.ts | 38 +- js/src/wrappers/google-genai.ts | 53 +- js/src/wrappers/groq.ts | 41 +- js/src/wrappers/huggingface-transformers.ts | 21 +- js/src/wrappers/huggingface.ts | 75 +- js/src/wrappers/langsmith.test.ts | 35 +- js/src/wrappers/langsmith.ts | 46 +- js/src/wrappers/mistral.ts | 91 +- js/src/wrappers/oai.ts | 55 +- js/src/wrappers/oai_responses.ts | 23 +- js/src/wrappers/ollama.test.ts | 30 +- js/src/wrappers/ollama.ts | 14 +- js/src/wrappers/openai-codex.ts | 24 +- js/src/wrappers/openai-promise-utils.ts | 67 +- js/src/wrappers/openrouter-agent.test.ts | 14 +- js/src/wrappers/openrouter-agent.ts | 12 +- js/src/wrappers/openrouter.ts | 42 +- js/src/wrappers/pi-coding-agent.test.ts | 20 +- js/src/wrappers/strands-agent-sdk.test.ts | 16 +- js/src/wrappers/strands-agent-sdk.ts | 4 +- .../error-handling.test.ts | 32 +- .../event-content.test.ts | 112 +- .../configurable-global-hook-registry.cjs | 15 +- .../fixtures/global-hook-listener.cjs | 4 +- .../incompatible-global-hook-registry.cjs | 28 +- .../fixtures/listener-cjs.cjs | 40 +- .../fixtures/listener-esm.mjs | 39 +- .../orchestrion-js/arguments_mutation/test.js | 29 +- .../orchestrion-js/ast_query_cjs/test.js | 6 +- .../orchestrion-js/callback_cjs/test.js | 6 +- .../class_expression_cjs/test.js | 6 +- .../orchestrion-js/class_method_cjs/test.js | 6 +- .../orchestrion-js/common/preamble.js | 40 +- .../const_class_export_alias_mjs/test.mjs | 6 +- .../fixtures/orchestrion-js/decl_cjs/test.js | 6 +- .../fixtures/orchestrion-js/decl_mjs/test.mjs | 6 +- .../decl_mjs_mismatched_type/test.mjs | 6 +- .../export_alias_class_mjs/test.mjs | 6 +- .../orchestrion-js/export_alias_mjs/test.mjs | 6 +- .../orchestrion-js/iife_nested_class/test.js | 4 +- .../fixtures/orchestrion-js/index_cjs/test.js | 12 +- .../instance_method_subclass_cjs/test.js | 6 +- .../let_class_export_alias_mjs/test.mjs | 6 +- .../multiple_class_method_cjs/test.js | 12 +- .../orchestrion-js/multiple_load_cjs/test.js | 6 +- .../orchestrion-js/nested_functions/test.js | 4 +- .../orchestrion-js/object_method_cjs/test.js | 6 +- .../object_property_named_cjs/test.js | 6 +- .../object_property_this_cjs/test.js | 6 +- .../orchestrion-js/private_method_cjs/test.js | 6 +- .../orchestrion-js/promise_subclass/test.js | 6 +- .../var_class_export_alias_mjs/test.mjs | 6 +- .../var_named_class_export_alias_mjs/test.mjs | 6 +- .../orchestrion-js/windows_path/test.js | 6 +- .../wrap_promise_non_promise/test.js | 6 +- .../auto-instrumentations/loader-hook.test.ts | 2 +- .../multiple-instrumentations.test.ts | 48 +- .../streaming-and-responses.test.ts | 30 +- .../auto-instrumentations/test-helpers.ts | 138 +- .../transformation.test.ts | 82 +- 190 files changed, 9231 insertions(+), 10348 deletions(-) create mode 100644 js/src/instrumentation/core/observe-result.test.ts create mode 100644 js/src/instrumentation/core/observe-result.ts create mode 100644 js/src/instrumentation/core/tracing-types.ts create mode 100644 js/src/instrumentation/test-utils/invocation.ts diff --git a/.agents/skills/instrumentation/SKILL.md b/.agents/skills/instrumentation/SKILL.md index 74e344508..f3adf1525 100644 --- a/.agents/skills/instrumentation/SKILL.md +++ b/.agents/skills/instrumentation/SKILL.md @@ -1,14 +1,14 @@ --- name: instrumentation -description: Add or update Braintrust SDK instrumentation. Use when working on instrumentation of any kind - like wrappers, auto-instrumentation configs, tracing channels, provider plugins, vendored SDK typings, or instrumentation-specific tests. +description: Add or update Braintrust SDK instrumentation. Use when working on instrumentation of any kind - like wrappers, auto-instrumentation configs, invocation hooks, provider plugins, vendored SDK typings, or instrumentation-specific tests. --- # Instrumentation Rules Read first based on the task: -- `js/src/instrumentation/README.md` for plugin and tracing-channel architecture -- Closest file in `js/src/instrumentation/core/` when changing shared channel semantics +- `js/src/instrumentation/README.md` for wrapping and tracing architecture +- Closest file in `js/src/instrumentation/core/` when changing shared invocation semantics - Closest file in `js/src/instrumentation/plugins/` when changing provider-specific extraction or span mapping - Closest file in `js/src/wrappers/` when manual wrappers and auto-instrumentation need to stay aligned - Closest test in `js/tests/auto-instrumentations/` when changing hook, loader, bundler, or transform behavior @@ -16,8 +16,8 @@ Read first based on the task: Map the change before editing: -- `js/src/instrumentation/core/` - tracing-channel helpers, stream patching, shared types -- `js/src/instrumentation/plugins/` - provider-specific channel subscriptions and event-to-span conversion +- `js/src/instrumentation/core/` - invocation definitions, independent tracing helpers, stream patching, shared types +- `js/src/instrumentation/plugins/` - explicit interceptor registration and provider-specific tracing functions - `js/src/wrappers/` - manual instrumentation entrypoints that should mirror the same logical contracts - `js/src/auto-instrumentations/` - loader and bundler instrumentation config - `js/tests/auto-instrumentations/` - functional coverage for transformed code @@ -27,10 +27,17 @@ Map the change before editing: - Inputs are untrusted: treat args, results, events, headers, and metadata as hostile. Prototype pollution is a concrete risk here. Avoid unsafe property access patterns, prototype-sensitive operations, and unnecessary mutation of third-party objects. - Support both auto-instrumentation and manual instrumentation. Auto-instrumentation does not cover every environment, loader, or framework. - For orchestrion auto-instrumentation, prefer targeting public API functions. Instrumenting internal helpers is more likely to break across library versions. -- Auto and manual paths should share logic through the same typed channel. For new and migrated instrumentation, prefer `invoke` in manual wrappers and `intercept` in provider plugins so the target can be scoped with `AsyncLocalStorage.run()` and its arguments, receiver, or output can be patched. Keep tracing-style hooks only as the compatibility path for instrumentation that has not migrated yet. Manual wrappers should not directly emit observability data. +- Keep wrapping and tracing separate, including internal APIs. + Define typed invocation hooks with `defineInterceptor`; `intercept` and `invoke` only compose and execute wrappers. + Definitions describe call arguments, return values, and opaque additional data, without span provenance or tracing methods. + The wrapping runtime must work without SDK initialization and have no dependency on spans, logging, or tracing lifecycle events. + Put span creation, context propagation, and finalization in separate tracing functions. + Plugins explicitly register those functions through `intercept`; shared tracing helpers accept callables and tracing configuration, never hooks or registration responsibilities. + Do not add combined helpers such as `traceInvocation` or `interceptAndTrace`, or retain a tracing-event compatibility path. + Manual wrappers and generated wrappers use the same hook through `invoke` and do not directly emit observability data. - Reuse shared repo utilities before introducing local helpers. Check `js/util/index.ts`, neighboring instrumentation files, and existing plugins/wrappers for utilities like `isObject`, merge helpers, and sanitizers before adding ad hoc replacements. - If a public instrumentation surface changes, check whether the export surface also needs updates in `js/src/instrumentation/index.ts` or `js/src/exports.ts`. -- Preserve async context propagation. Changes around tracing channels, stream patching, or loader hooks must keep the current span context across awaits and stream consumption. +- Preserve async context propagation. Changes around invocation hooks, stream patching, or loader hooks must keep the current span context across awaits and stream consumption. - Maintain isomorphic behavior. Node and browser/bundled paths must use compatible channel implementations and avoid channel-registry mismatches. - Setup, teardown, and patching must be idempotent. Enabling twice, disabling twice, or applying a patch twice should remain safe. - Promise/stream behavior must be preserved. Patches need to keep subclass/helper semantics intact. @@ -44,7 +51,7 @@ Map the change before editing: Preserve useful text, metadata, metrics, and remote references when attachment capture is disabled, and omit inline bytes from logged payloads. - Do not modify package READMEs during instrumentation work unless the user specifically requests it or the change corrects outdated information. - We want to limit our instrumentation to operations that are relevant for AI generations and operations (LLMs, embeddings, media generation, ...). Things like creating entities on platforms (CRUD for Workflows of Agent entities) is irrelevant to us. -- When building instrumentation, we should always have a vendored type/interface for what we are wrapping. The type or interface should not be larger than what is relevant to the instrumentation. The type or interface should be used for typing tracing channels and also should be used to assert the type on whatever is passed into wrappers as soon as the wrapper has verified that the passed in value is plausibly what should be wrapped. +- When building instrumentation, we should always have a vendored type/interface for what we are wrapping. The type or interface should not be larger than what is relevant to the instrumentation. The type or interface should be used for typing invocation hooks and also should be used to assert the type on whatever is passed into wrappers as soon as the wrapper has verified that the passed in value is plausibly what should be wrapped. ## Process diff --git a/AGENTS.md b/AGENTS.md index f474c4bcb..22fc16f15 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -33,6 +33,14 @@ Keep public exports minimal. Generally, export only the requested runtime APIs a Use the normal Orchestrion config plus plugin/channel path by default. Special-case source patches should be rare exceptions only when the target SDK cannot be instrumented through the standard transformer path, and the reason should be documented next to the patch. +API wrapping and span instrumentation are separate concepts. +Use `defineInterceptor` to define invocation hooks and `intercept`/`invoke` exclusively as generic wrapping machinery. +Wrapping definitions and runtime code must not depend on spans, tracing lifecycle events, provenance, or SDK initialization. +Keep span creation, context propagation, and finalization in separate tracing functions. +Provider plugins explicitly register those functions through interceptors; tracing helpers accept callables and tracing configuration, never hooks or registration responsibilities. +Do not introduce combined APIs such as `traceInvocation` or `interceptAndTrace`, including internal convenience APIs. +Manual wrappers and generated wrappers only invoke hooks. + Instrumentation patches generally do not need to be removed during teardown. Prefer leaving behavior-preserving patches installed when they are idempotent; do not add unpatching machinery by default. Span names should generally remain stable across calls and versions. Do not include dynamic values such as model names in span names; record those values in metadata instead. diff --git a/e2e/scenarios/nextjs-auto-instrumentation/template/app/api/test/route.ts b/e2e/scenarios/nextjs-auto-instrumentation/template/app/api/test/route.ts index 282e699dd..e499491cb 100644 --- a/e2e/scenarios/nextjs-auto-instrumentation/template/app/api/test/route.ts +++ b/e2e/scenarios/nextjs-auto-instrumentation/template/app/api/test/route.ts @@ -7,12 +7,13 @@ import type { AddressInfo } from "node:net"; export const dynamic = "force-dynamic"; type InstrumentationHook = { - subscribe(handlers: InstrumentationHookHandlers): void; - unsubscribe(handlers: InstrumentationHookHandlers): boolean; -}; - -type InstrumentationHookHandlers = { - start(): void; + intercept( + interceptor: ( + target: (...args: unknown[]) => unknown, + receiver: unknown, + args: unknown[], + ) => unknown, + ): () => void; }; export async function GET() { @@ -43,18 +44,15 @@ export async function GET() { const hooks = ( globalThis as typeof globalThis & { - __braintrust_instrumentation_hooks?: Map; + __braintrust_invocation_hooks_v2?: Map; } - ).__braintrust_instrumentation_hooks; + ).__braintrust_invocation_hooks_v2; const hook = hooks?.get("orchestrion:openai:chat.completions.create"); let hookFired = false; - const subscriber = { - start: () => { - hookFired = true; - }, - }; - - hook?.subscribe(subscriber); + const remove = hook?.intercept((target, receiver, args) => { + hookFired = true; + return Reflect.apply(target, receiver, args); + }); try { const client = new OpenAI({ @@ -67,7 +65,7 @@ export async function GET() { messages: [{ role: "user", content: "hi" }], }); } finally { - hook?.unsubscribe(subscriber); + remove?.(); mockServer.close(); } diff --git a/js/src/auto-instrumentations/README.md b/js/src/auto-instrumentations/README.md index d9c784fea..cfa20d2cc 100644 --- a/js/src/auto-instrumentations/README.md +++ b/js/src/auto-instrumentations/README.md @@ -35,58 +35,26 @@ same identifier. ## Generated Runtime Contract -For every configured channel, transformed modules lazily look up: +Generated wrappers lazily look up invocation hooks in `globalThis.__braintrust_invocation_hooks_v2`. +They pass the original target, receiver, complete arguments, and module version to `invoke`. +They do not emit tracing events, create spans, or select tracing operators. +Legacy `functionQuery.kind` and `callbackIndex` fields remain accepted for source compatibility but do not select runtime tracing behavior. -```js -globalThis.__braintrust_instrumentation_hooks?.get( - "orchestrion:openai:chat.completions.create", -); -``` - -The lookup is retried until a hook exists, then cached. This has two important -properties: - -- Loading an instrumented provider before Braintrust is safe; calls run normally. -- Registering Braintrust later enables tracing without retransformation. - -The generated code only applies `traceInvocation`, passing the configured -operator, original target, receiver, complete arguments, and `moduleVersion`. -The normal hook runtime handles interceptor composition, the no-listener fast -path, tracing context construction, and the legacy `tracePromise`, `traceSync`, -or `traceCallback` dispatch around the effective intercepted call. - -The hook lifecycle mirrors tracing channels: - -1. `start` before the target call -2. `end` after its synchronous portion -3. `asyncStart` and `asyncEnd` when an asynchronous result settles -4. `error` for synchronous throws, promise rejections, or callback errors - -The same context object is passed through every phase. Subscribers may mutate -arguments or returned streams before user code continues. - -Invocation interceptors compose as nested middleware and may replace arguments, -the receiver, the returned value, or the entire implementation. Tracing remains -the outer compatibility layer, so tracing subscribers observe the interceptor's -effective result. +Calls run normally before a hook is registered. +Lookup retries until registration succeeds, then caches the hook. +Interceptors compose in registration order and may replace arguments, receivers, results, or the complete implementation. ## Global Registry -The SDK installs `globalThis.__braintrust_instrumentation_hooks` with a -non-enumerable, non-writable property descriptor. Its value is a mutable -`Map` shared by all Braintrust SDK copies in the realm. - -The implementation lives in `src/global-instrumentation-hooks.ts`. It supports: - -- composable `invoke` / `intercept` wrappers -- all five lifecycle phases -- multiple subscribers and complete unsubscription -- `bindStore` / `unbindStore` for async-context propagation -- sync, promise, and callback tracing operators -- preservation of Promise subclasses, thenables, and non-Promise return values +The invocation registry is a non-enumerable, non-writable global property containing a shared map. +It is independent of SDK initialization, async-context storage, and tracing. +Manual wrappers use the same hooks through `defineInterceptor` and `invoke`. +Provider plugins explicitly register separate tracing functions through `intercept`. +Do not combine wrapping and tracing in a runtime or convenience API. -Manual wrappers use the same registry through typed channel definitions, so -manual and auto-instrumented paths share lifecycle and span behavior. +Protocol version 2 uses a separately keyed registry so it does not mutate an older SDK's registry. +Previously transformed bundles must be rebuilt with the updated SDK to retain instrumentation. +There is no legacy tracing-event compatibility layer. ## Loaders and Bundlers @@ -110,8 +78,7 @@ bundles. 1. Add the narrowest supported package/version/file/function config under `configs/`. 2. Define a typed channel with the same package and operation identifier. -3. Add or update a plugin that intercepts the typed channel; use the tracing - helpers only for existing instrumentation awaiting migration. +3. Write a separate tracing function and explicitly register it through `intercept` in the provider plugin. 4. Keep manual wrappers on that same typed channel through `invoke`. 5. Add transformation/runtime coverage and a provider e2e scenario when the user-visible trace contract changes. diff --git a/js/src/auto-instrumentations/orchestrion-js/transformer.ts b/js/src/auto-instrumentations/orchestrion-js/transformer.ts index 41de794f5..a1465b45d 100644 --- a/js/src/auto-instrumentations/orchestrion-js/transformer.ts +++ b/js/src/auto-instrumentations/orchestrion-js/transformer.ts @@ -3,13 +3,12 @@ * licensed under Apache-2.0. Modified by Braintrust. */ -import esquery from "esquery"; import { generate } from "astring"; +import esquery from "esquery"; import { parse } from "meriyah"; import { SourceMapGenerator } from "source-map"; import { transforms, type TransformState } from "./transforms"; import type { - FunctionKind, FunctionQuery, InstrumentationConfig, ModuleType, @@ -18,7 +17,6 @@ import type { type AnyNode = any; type ExportAliases = Record; -type TraceOperator = "traceCallback" | "tracePromise" | "traceSync"; /** * Applies instrumentation configs to JavaScript source by parsing it into an @@ -90,7 +88,6 @@ export class Transformer { ...config, moduleVersion: this.version, functionQuery: resolvedFunctionQuery, - operator: this.getOperator(resolvedFunctionQuery.kind), }; esquery.traverse(ast, esquery.parse(query), (...args: any[]) => { @@ -141,7 +138,7 @@ export class Transformer { free(): void {} private visit(state: TransformState, ...args: any[]): void { - const transform = transforms[state.operator]; + const transform = transforms.invoke; const { index = 0 } = state.functionQuery as any; const [node] = args; const type = node.init?.type || node.type; @@ -164,17 +161,6 @@ export class Transformer { (transform as (...args: any[]) => void)(state, ...args); } - private getOperator(kind: FunctionKind): TraceOperator { - switch (kind) { - case "Async": - return "tracePromise"; - case "Callback": - return "traceCallback"; - case "Sync": - return "traceSync"; - } - } - private collectExportAliases(ast: AnyNode): ExportAliases { const aliases: ExportAliases = {}; for (const node of ast.body) { diff --git a/js/src/auto-instrumentations/orchestrion-js/transforms.ts b/js/src/auto-instrumentations/orchestrion-js/transforms.ts index 72f6e03ca..c59b93cfa 100644 --- a/js/src/auto-instrumentations/orchestrion-js/transforms.ts +++ b/js/src/auto-instrumentations/orchestrion-js/transforms.ts @@ -6,10 +6,10 @@ import esquery from "esquery"; import { parse } from "meriyah"; import { - GLOBAL_INSTRUMENTATION_HOOK_BRAND, GLOBAL_INSTRUMENTATION_HOOKS_KEY, GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, GLOBAL_INSTRUMENTATION_HOOKS_REGISTRY_BRAND, + GLOBAL_INVOCATION_HOOK_BRAND, } from "../../global-instrumentation-hooks"; import type { FunctionQuery, InstrumentationConfig } from "./types"; @@ -20,12 +20,10 @@ type TransformFn = ( parent: AnyNode, ancestry: AnyNode[], ) => void; -type TraceOperator = "traceCallback" | "tracePromise" | "traceSync"; export interface TransformState extends InstrumentationConfig { moduleVersion: string; functionQuery: FunctionQuery; - operator: TraceOperator; functionIndex?: number; } @@ -41,7 +39,7 @@ function formatChannelGetter(channelName: string): string { } export const transforms: Record = { - tracingHookDeclaration(state, node) { + invocationHookDeclaration(state, node) { const { channelName, module: { name }, @@ -84,9 +82,9 @@ export const transforms: Record = { (typeof __bt$hook !== "object" && typeof __bt$hook !== "function")) || __bt$hook[Symbol.for(${JSON.stringify( - GLOBAL_INSTRUMENTATION_HOOK_BRAND, + GLOBAL_INVOCATION_HOOK_BRAND, )})] !== ${GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION} || - typeof __bt$hook.traceInvocation !== "function" + typeof __bt$hook.invoke !== "function" ) return undefined; return __bt$hook; } catch { @@ -116,9 +114,7 @@ export const transforms: Record = { node.body.splice(index + 1, 0, ...parse(code).body); }, - traceCallback: traceAny, - tracePromise: traceAny, - traceSync: traceAny, + invoke: traceAny, }; function traceAny( @@ -141,7 +137,7 @@ function traceFunction( node: AnyNode, program: AnyNode, ): void { - transforms.tracingHookDeclaration(state, program, null, []); + transforms.invocationHookDeclaration(state, program, null, []); const isArrowFunction = node.type === "ArrowFunctionExpression"; @@ -195,7 +191,7 @@ function traceInstanceMethod( node: AnyNode, program: AnyNode, ): void { - const { functionQuery, operator } = state; + const { functionQuery } = state; const { methodName } = functionQuery as any; if (!methodName) { @@ -210,7 +206,7 @@ function traceInstanceMethod( let ctor = classBody.body.find(({ kind }: AnyNode) => kind === "constructor"); - transforms.tracingHookDeclaration(state, program, null, []); + transforms.invocationHookDeclaration(state, program, null, []); if (!ctor) { ctor = ( @@ -233,7 +229,7 @@ function traceInstanceMethod( const fn = ctorBody[1].expression.right; - fn.async = operator === "tracePromise"; + fn.async = false; fn.body = wrap( state, { @@ -331,21 +327,18 @@ function wrapInvocation( state: TransformState, argsExpression: string, ): AnyNode { - const { channelName, moduleVersion, operator, functionQuery } = state; + const { channelName, moduleVersion } = state; const channelGetter = formatChannelGetter(channelName); - const callbackIndex = functionQuery.callbackIndex ?? -1; return parse(` function wrapper () { const __bt$hook = ${channelGetter}(); if (!__bt$hook) return __bt$target.apply(this, ${argsExpression}); - return __bt$hook.traceInvocation( - ${JSON.stringify(operator)}, + return __bt$hook.invoke( __bt$target, this, ${argsExpression}, - { moduleVersion: ${JSON.stringify(moduleVersion)} }, - ${callbackIndex} + { moduleVersion: ${JSON.stringify(moduleVersion)} } ); } `); diff --git a/js/src/cli/auto-instrumentation.test.ts b/js/src/cli/auto-instrumentation.test.ts index 2aa10990d..602b53878 100644 --- a/js/src/cli/auto-instrumentation.test.ts +++ b/js/src/cli/auto-instrumentation.test.ts @@ -5,7 +5,7 @@ import * as path from "node:path"; import { execFile } from "node:child_process"; import { promisify } from "node:util"; import { fileURLToPath } from "node:url"; -import { newGlobalTracingChannel } from "../global-instrumentation-hooks"; +import { newGlobalInvocationHook } from "../global-instrumentation-hooks"; import { initializeHandles } from "./index"; import type { FileHandle } from "./types"; @@ -89,14 +89,13 @@ describe("eval auto-instrumentation", () => { expect(output).toContain(googleGenAIChannel); const lifecycle: string[] = []; - const hook = newGlobalTracingChannel(googleGenAIChannel); - const handlers = { - asyncEnd: () => lifecycle.push("asyncEnd"), - asyncStart: () => lifecycle.push("asyncStart"), - end: () => lifecycle.push("end"), - start: () => lifecycle.push("start"), - }; - hook.subscribe(handlers); + const hook = newGlobalInvocationHook(googleGenAIChannel); + const remove = hook.intercept((target, receiver, args) => { + lifecycle.push("called"); + const result = Reflect.apply(target, receiver, args); + result.then(() => lifecycle.push("resolved")); + return result; + }); try { const loadedModule = { exports: {} as Record }; @@ -109,9 +108,9 @@ describe("eval auto-instrumentation", () => { await expect(invoke(loadedModule.exports)).resolves.toEqual({ text: "payload", }); - expect(lifecycle).toEqual(["start", "end", "asyncStart", "asyncEnd"]); + expect(lifecycle).toEqual(["called", "resolved"]); } finally { - hook.unsubscribe(handlers); + remove(); } }, ); @@ -175,7 +174,7 @@ describe("eval auto-instrumentation", () => { }); it.each([ - ["instruments", undefined, ["start", "end", "asyncStart", "asyncEnd"]], + ["instruments", undefined, ["called", "resolved"]], ["respects opt-out for", "anthropic", []], ])("%s external eval dependencies", async (_label, disabled, expected) => { const runnerPath = path.join(fixtureDir, "external-eval-runner.cjs"); @@ -186,19 +185,21 @@ describe("eval auto-instrumentation", () => { `require("tsx/cjs"); require("node:module").register = () => {}; const { loadModule } = require(${JSON.stringify(loadModulePath)}); -const { newGlobalTracingChannel } = require(${JSON.stringify(globalHooksPath)}); +const { newGlobalInvocationHook } = require(${JSON.stringify(globalHooksPath)}); const lifecycle = []; -const channel = newGlobalTracingChannel(${JSON.stringify(anthropicChannel)}); -const handlers = Object.fromEntries( - ["start", "end", "asyncStart", "asyncEnd"].map((name) => [name, () => lifecycle.push(name)]), -); -channel.subscribe(handlers); +const channel = newGlobalInvocationHook(${JSON.stringify(anthropicChannel)}); +const remove = channel.intercept((target, receiver, args) => { + lifecycle.push("called"); + const result = Reflect.apply(target, receiver, args); + result.then(() => lifecycle.push("resolved")); + return result; +}); loadModule({ inFile: ${JSON.stringify(evalFile)}, moduleText: 'const { Messages } = require("@anthropic-ai/sdk/resources/messages/messages.js"); globalThis.__externalEvalResult = new Messages().create("payload");', }); Promise.resolve(globalThis.__externalEvalResult).then((result) => { - channel.unsubscribe(handlers); + remove(); process.stdout.write(JSON.stringify({ lifecycle, result })); });`, ); diff --git a/js/src/debug-logger.test.ts b/js/src/debug-logger.test.ts index 73e34ff87..d5a369fe5 100644 --- a/js/src/debug-logger.test.ts +++ b/js/src/debug-logger.test.ts @@ -1,17 +1,20 @@ import { afterEach, beforeEach, describe, expect, test, vi } from "vitest"; +import { + debugLogger, + getEnvDebugLogLevel, + resetDebugLoggerForTests, +} from "./debug-logger"; +import { + newGlobalInvocationHook, + GLOBAL_INSTRUMENTATION_HOOKS_KEY, +} from "./global-instrumentation-hooks"; import { BraintrustState, _exportsForTestingOnly, initLogger, login, } from "./logger"; -import { - debugLogger, - getEnvDebugLogLevel, - resetDebugLoggerForTests, -} from "./debug-logger"; -import { newGlobalTracingChannel } from "./global-instrumentation-hooks"; import { configureNode } from "./node/config"; configureNode(); @@ -72,24 +75,22 @@ describe("debug logger", () => { expect(debugSpy).toHaveBeenCalledWith("[braintrust]", "debug"); }); - test("global hook failures are logged without escaping the provider call", () => { + test("invalid global hooks are reported without changing the provider call", () => { process.env.BRAINTRUST_DEBUG_LOG_LEVEL = "error"; const errorSpy = vi.spyOn(console, "error").mockImplementation(() => {}); - const subscriberError = new Error("subscriber failed"); - const channel = newGlobalTracingChannel>( - `test:debug-logger:${Math.random()}`, - ); - channel.subscribe({ - start() { - throw subscriberError; - }, - }); - - expect(channel.traceSync(() => "result", {})).toBe("result"); + const name = `test:debug-logger:${Math.random()}`; + newGlobalInvocationHook(name); + const registry = ( + globalThis as unknown as Record> + )[GLOBAL_INSTRUMENTATION_HOOKS_KEY]; + registry.set(name, {}); + expect( + newGlobalInvocationHook(name).invoke(() => "result", undefined, [], {}), + ).toBe("result"); expect(errorSpy).toHaveBeenCalledWith( "[braintrust]", "Global instrumentation hook error:", - subscriberError, + expect.any(Error), ); }); diff --git a/js/src/global-instrumentation-hooks.test.ts b/js/src/global-instrumentation-hooks.test.ts index 687695850..d0e93d577 100644 --- a/js/src/global-instrumentation-hooks.test.ts +++ b/js/src/global-instrumentation-hooks.test.ts @@ -1,615 +1,170 @@ import { AsyncLocalStorage } from "node:async_hooks"; import { randomUUID } from "node:crypto"; -import { tracingChannel } from "node:diagnostics_channel"; -import { afterEach, describe, expect, it, vi } from "vitest"; +import { describe, expect, it, vi } from "vitest"; import { - GLOBAL_INSTRUMENTATION_HOOK_BRAND, GLOBAL_INSTRUMENTATION_HOOKS_KEY, GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, GLOBAL_INSTRUMENTATION_HOOKS_REGISTRY_BRAND, - GLOBAL_INVOCATION_HOOK_BRAND, - newGlobalTracingChannel, - setGlobalHookErrorReporter, + newGlobalInvocationHook, } from "./global-instrumentation-hooks"; -function uniqueChannelName(label: string): string { - return `test:${label}:${randomUUID()}`; -} +const hook = () => newGlobalInvocationHook(`test:${randomUUID()}`); -describe("global instrumentation hooks", () => { - let restoreErrorReporter: (() => void) | undefined; - - afterEach(() => { - restoreErrorReporter?.(); - restoreErrorReporter = undefined; - }); - - it("installs a non-enumerable, immutable global registry", () => { - const name = uniqueChannelName("descriptor"); - const channel = newGlobalTracingChannel(name); +describe("invocation hooks", () => { + it("shares hooks in an immutable, non-enumerable versioned registry", () => { + const name = randomUUID(); + expect(newGlobalInvocationHook(name)).toBe(newGlobalInvocationHook(name)); const descriptor = Object.getOwnPropertyDescriptor( globalThis, GLOBAL_INSTRUMENTATION_HOOKS_KEY, - ); - + )!; expect(descriptor).toMatchObject({ configurable: false, enumerable: false, writable: false, }); - expect(descriptor?.value).toBeInstanceOf(Map); + expect(descriptor.value).toBeInstanceOf(Map); expect( - descriptor?.value[ - Symbol.for(GLOBAL_INSTRUMENTATION_HOOKS_REGISTRY_BRAND) - ], + descriptor.value[Symbol.for(GLOBAL_INSTRUMENTATION_HOOKS_REGISTRY_BRAND)], ).toBe(GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION); - expect( - (channel as any)[Symbol.for(GLOBAL_INSTRUMENTATION_HOOK_BRAND)], - ).toBe(GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION); - expect((channel as any)[Symbol.for(GLOBAL_INVOCATION_HOOK_BRAND)]).toBe( - GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, - ); - expect(newGlobalTracingChannel(name)).toBe(channel); - expect(Object.keys(globalThis)).not.toContain( - GLOBAL_INSTRUMENTATION_HOOKS_KEY, - ); + expect(GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION).toBe(2); }); - it("shares subscriptions and supports complete unsubscription", () => { - const name = uniqueChannelName("subscriptions"); - const first = newGlobalTracingChannel>(name); - const second = newGlobalTracingChannel>(name); - const events: string[] = []; - const handlers = { - start: () => events.push("start"), - end: () => events.push("end"), - }; - - expect(first).toBe(second); - first.subscribe(handlers); - expect(second.hasSubscribers).toBe(true); - expect(second.traceSync(() => 42, {})).toBe(42); - expect(events).toEqual(["start", "end"]); - expect(second.unsubscribe(handlers)).toBe(true); - expect(first.hasSubscribers).toBe(false); - expect(second.unsubscribe(handlers)).toBe(false); + it("exposes only wrapping without requiring SDK initialization", () => { + const invocation = hook(); + for (const method of [ + "subscribe", + "traceInvocation", + "tracePromise", + "start", + "end", + ]) { + expect(method in invocation).toBe(false); + } + const receiver = { value: 4 }; + const target = vi.fn(function (this: typeof receiver, n: number) { + return this.value + n; + }); + expect(invocation.invoke(target, receiver, [3], {})).toBe(7); + expect(target).toHaveBeenCalledOnce(); }); - it("composes invocation interceptors in registration order", () => { - const channel = newGlobalTracingChannel(uniqueChannelName("interceptors")); + it("composes in registration order and permits replacing receiver, arguments and output", () => { + const invocation = hook(); const order: string[] = []; - const firstReceiver = { prefix: "first" }; - const secondReceiver = { prefix: "second" }; - - const removeFirst = channel.intercept( - (target, _thisArg, args, additional: { suffix: string }) => { - order.push(`first:${additional.suffix}`); - return `${target.apply(secondReceiver, [args[0] + 1])}:first`; - }, - ); - const removeSecond = channel.intercept((target, thisArg, args) => { - order.push("second"); - return `${target.apply(thisArg, [args[0] * 2])}:second`; + const removeFirst = invocation.intercept((next, _, args, extra) => { + order.push(extra.label); + return `${next.apply({ offset: 5 }, [args[0] + 1])}:outer`; }); - const target = vi.fn(function (this: { prefix: string }, value: number) { - order.push("target"); - return `${this.prefix}:${value}`; + const removeSecond = invocation.intercept((next, receiver, args) => { + order.push("inner"); + return next.apply(receiver, [args[0] * 2]); }); - + const target = function (this: { offset: number }, n: number) { + order.push("target"); + return this.offset + n; + }; expect( - channel.invoke(target, firstReceiver, [2], { suffix: "extra" }), - ).toBe("second:6:second:first"); - expect(order).toEqual(["first:extra", "second", "target"]); - expect(target).toHaveBeenCalledOnce(); - + invocation.invoke(target, { offset: 0 }, [2], { label: "outer" }), + ).toBe("11:outer"); + expect(order).toEqual(["outer", "inner", "target"]); removeFirst(); removeFirst(); removeSecond(); - expect(channel.hasInterceptors).toBe(false); - expect(channel.invoke(target, firstReceiver, [3], {})).toBe("first:3"); + expect(invocation.hasInterceptors).toBe(false); + expect(invocation.invoke(target, { offset: 0 }, [2], {})).toBe(2); }); - it("allows interceptors to replace calls and propagates their errors", () => { - const channel = newGlobalTracingChannel(uniqueChannelName("replacement")); - const target = vi.fn(() => "original"); - const removeReplacement = channel.intercept(() => "replacement"); - - expect(channel.invoke(target, undefined, [], {})).toBe("replacement"); + it("allows skipping and repeating the target and propagates interceptor failures", () => { + const invocation = hook(); + const target = vi.fn(() => 2); + const remove = invocation.intercept(() => 9); + expect(invocation.invoke(target, undefined, [], {})).toBe(9); expect(target).not.toHaveBeenCalled(); - - removeReplacement(); - const error = new Error("interceptor failed"); - channel.intercept(() => { + remove(); + const removeRepeat = invocation.intercept((next) => next() + next()); + expect(invocation.invoke(target, undefined, [], {})).toBe(4); + expect(target).toHaveBeenCalledTimes(2); + removeRepeat(); + const error = new Error("interceptor"); + invocation.intercept(() => { throw error; }); - expect(() => channel.invoke(target, undefined, [], {})).toThrow(error); - expect(target).not.toHaveBeenCalled(); - }); - - it("upgrades compatible tracing-only hooks with invocation support", () => { - const name = uniqueChannelName("legacy-upgrade"); - const donor = newGlobalTracingChannel(uniqueChannelName("legacy-donor")); - const legacyHook = { - asyncEnd: donor.asyncEnd, - asyncStart: donor.asyncStart, - end: donor.end, - error: donor.error, - hasSubscribers: false, - start: donor.start, - subscribe: vi.fn(), - traceCallback: vi.fn((target: (...args: unknown[]) => unknown) => - target(), - ), - tracePromise: vi.fn((target: (...args: unknown[]) => unknown) => - target(), - ), - traceSync: vi.fn((target: (...args: unknown[]) => unknown) => target()), - unsubscribe: vi.fn(() => true), - }; - Object.defineProperty( - legacyHook, - Symbol.for(GLOBAL_INSTRUMENTATION_HOOK_BRAND), - { value: GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION }, - ); - const registry = Object.getOwnPropertyDescriptor( - globalThis, - GLOBAL_INSTRUMENTATION_HOOKS_KEY, - )?.value as Map; - registry.set(name, legacyHook); - - const upgraded = newGlobalTracingChannel(name); - expect(upgraded).toBe(legacyHook); - const removeInterceptor = upgraded.intercept(() => "upgraded"); - expect(upgraded.invoke(() => "original", undefined, [], {})).toBe( - "upgraded", - ); - expect( - upgraded.traceInvocation( - "traceSync", - () => "original", - undefined, - [], - {}, - ), - ).toBe("upgraded"); - expect(legacyHook.traceSync).toHaveBeenCalledOnce(); - removeInterceptor(); + expect(() => invocation.invoke(target, undefined, [], {})).toThrow(error); + expect(target).toHaveBeenCalledTimes(2); }); - it("supports AsyncLocalStorage wrapping across asynchronous targets", async () => { - const channel = newGlobalTracingChannel(uniqueChannelName("als")); - const storage = new AsyncLocalStorage(); - channel.intercept((target, thisArg, args) => - storage.run("intercepted", () => target.apply(thisArg, args)), + it("preserves exact values, Promise subclasses, streams, callbacks and thrown errors", async () => { + class SpecialPromise extends Promise { + helper() { + return 42; + } + } + const invocation = hook(); + invocation.intercept((next, receiver, args) => + Reflect.apply(next, receiver, args), + ); + const promise = new SpecialPromise((resolve) => resolve(3)); + expect(invocation.invoke(() => promise, undefined, [], {})).toBe(promise); + expect(promise.helper()).toBe(42); + await expect(promise).resolves.toBe(3); + const stream = (async function* () { + yield 1; + })(); + expect(invocation.invoke(() => stream, undefined, [], {})).toBe(stream); + const callback = vi.fn(); + invocation.invoke( + (cb: typeof callback) => cb(null, 7), + undefined, + [callback], + {}, ); - - await expect( - channel.invoke( - async () => { - await Promise.resolve(); - return storage.getStore(); + expect(callback).toHaveBeenCalledWith(null, 7); + const error = new Error("provider"); + expect(() => + invocation.invoke( + () => { + throw error; }, undefined, [], {}, ), - ).resolves.toBe("intercepted"); - expect(storage.getStore()).toBeUndefined(); - }); - - it("keeps tracing outside the invocation interceptor", () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("tracing-order"), - ); - const lifecycle: string[] = []; - const context: Record = {}; - channel.subscribe({ - start: () => lifecycle.push("start"), - end: (event) => lifecycle.push(`end:${event.result}`), - }); - channel.intercept(() => { - lifecycle.push("interceptor"); - return "replacement"; - }); - - expect( - channel.traceSync( - () => channel.invoke(() => "original", undefined, [], {}), - context, - ), - ).toBe("replacement"); - expect(context.result).toBe("replacement"); - expect(lifecycle).toEqual(["start", "interceptor", "end:replacement"]); - }); - - it("does not publish legacy diagnostics-channel events", () => { - const name = uniqueChannelName("hard-cutover"); - const diagnostics = tracingChannel(name); - const diagnosticsEvents: unknown[] = []; - const diagnosticsHandler = (message: unknown) => - diagnosticsEvents.push(message); - diagnostics.start.subscribe(diagnosticsHandler); - - const hook = newGlobalTracingChannel>(name); - const hookEvents: unknown[] = []; - hook.subscribe({ start: (message) => hookEvents.push(message) }); - hook.traceSync(() => "result", {}); - - expect(hookEvents).toHaveLength(1); - expect(diagnosticsEvents).toHaveLength(0); - diagnostics.start.unsubscribe(diagnosticsHandler); - }); - - it("mirrors sync lifecycle ordering and rethrows errors", () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("sync"), - ); - const lifecycle: string[] = []; - const contexts: Record[] = []; - channel.subscribe({ - start: (context) => { - lifecycle.push("start"); - contexts.push(context); - }, - end: (context) => { - lifecycle.push("end"); - contexts.push(context); - }, - error: (context) => { - lifecycle.push("error"); - contexts.push(context); - }, - }); - - const successContext: Record = {}; - expect(channel.traceSync(() => "result", successContext)).toBe("result"); - expect(successContext.result).toBe("result"); - expect(lifecycle).toEqual(["start", "end"]); - expect(contexts).toEqual([successContext, successContext]); - - lifecycle.length = 0; - contexts.length = 0; - const error = new Error("boom"); - const errorContext: Record = {}; - expect(() => - channel.traceSync(() => { - throw error; - }, errorContext), ).toThrow(error); - expect(errorContext.error).toBe(error); - expect(lifecycle).toEqual(["start", "error", "end"]); - expect(contexts).toEqual([errorContext, errorContext, errorContext]); - }); - - it("mirrors promise lifecycle ordering for resolution and rejection", async () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("promise"), - ); - const lifecycle: string[] = []; - channel.subscribe({ - start: () => lifecycle.push("start"), - end: () => lifecycle.push("end"), - asyncStart: () => lifecycle.push("asyncStart"), - asyncEnd: () => lifecycle.push("asyncEnd"), - error: () => lifecycle.push("error"), - }); - - const successContext: Record = {}; - await expect( - channel.tracePromise(async () => "result", successContext), - ).resolves.toBe("result"); - expect(successContext.result).toBe("result"); - expect(lifecycle).toEqual(["start", "end", "asyncStart", "asyncEnd"]); - - lifecycle.length = 0; - const error = new Error("rejected"); - const errorContext: Record = {}; - await expect( - channel.tracePromise(async () => { - throw error; - }, errorContext), - ).rejects.toBe(error); - expect(errorContext.error).toBe(error); - expect(lifecycle).toEqual([ - "start", - "end", - "error", - "asyncStart", - "asyncEnd", - ]); - }); - - it("preserves promise subclasses and non-Promise return values", async () => { - class HelperPromise extends Promise { - withResponse(): Promise<{ data: T }> { - return this.then((data) => ({ data })); - } - } - - const channel = newGlobalTracingChannel>( - uniqueChannelName("promise-subclass"), - ); - const lifecycle: string[] = []; - channel.subscribe({ - end: () => lifecycle.push("end"), - asyncStart: () => lifecycle.push("asyncStart"), - asyncEnd: () => lifecycle.push("asyncEnd"), - }); - const original = new HelperPromise((resolve) => resolve("ok")); - const traced = channel.tracePromise(() => original, {}); - - expect(traced).toBe(original); - await expect(traced.withResponse()).resolves.toEqual({ data: "ok" }); - lifecycle.length = 0; - - const nonPromise = channel.tracePromise( - (() => 42) as unknown as () => PromiseLike, - {}, - ); - expect(nonPromise).toBe(42); - expect(lifecycle).toEqual(["end", "asyncStart", "asyncEnd"]); - }); - - it("returns unusual thenables unchanged when inspecting them fails", () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("hostile-thenable"), - ); - const reportedErrors: unknown[] = []; - const lifecycle: string[] = []; - restoreErrorReporter = setGlobalHookErrorReporter((error) => - reportedErrors.push(error), - ); - channel.subscribe({ - start: () => lifecycle.push("start"), - end: () => lifecycle.push("end"), - asyncStart: () => lifecycle.push("asyncStart"), - asyncEnd: () => lifecycle.push("asyncEnd"), - error: () => lifecycle.push("error"), - }); - - const getterError = new Error("then getter failed"); - const throwingGetter = Object.defineProperty({}, "then", { - get() { - throw getterError; - }, - }) as PromiseLike; - const getterContext: Record = {}; - expect(channel.tracePromise(() => throwingGetter, getterContext)).toBe( - throwingGetter, - ); - expect(getterContext.result).toBe(throwingGetter); - - const constructorError = new Error("constructor getter failed"); - const throwingConstructor = Object.defineProperty( - Object.create(Promise.prototype), - "constructor", - { - get() { - throw constructorError; - }, - }, - ) as PromiseLike; - const constructorContext: Record = {}; - expect( - channel.tracePromise(() => throwingConstructor, constructorContext), - ).toBe(throwingConstructor); - expect(constructorContext.result).toBe(throwingConstructor); - - const invocationError = new Error("then invocation failed"); - const throwingThen = { - then() { - throw invocationError; - }, - } as PromiseLike; - const invocationContext: Record = {}; - expect(channel.tracePromise(() => throwingThen, invocationContext)).toBe( - throwingThen, - ); - expect(invocationContext.result).toBe(throwingThen); - - expect(lifecycle).toEqual([ - "start", - "end", - "asyncStart", - "asyncEnd", - "start", - "end", - "asyncStart", - "asyncEnd", - "start", - "end", - "asyncStart", - "asyncEnd", - ]); - expect(reportedErrors).toEqual([ - getterError, - constructorError, - invocationError, - ]); - }); - - it("contains subscriber and store failures without repeating provider calls", () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("instrumentation-errors"), - ); - const reportedErrors: unknown[] = []; - restoreErrorReporter = setGlobalHookErrorReporter((error) => - reportedErrors.push(error), - ); - - const subscriberError = new Error("subscriber failed"); - channel.subscribe({ - start: () => { - throw subscriberError; - }, - }); - - const storeError = new Error("store failed"); - const brokenStore = { - getStore() { - return undefined; - }, - run(_store: unknown, callback: () => T): T { - callback(); - throw storeError; - }, - }; - channel.start.bindStore(brokenStore); - - const provider = vi.fn(() => "result"); - expect(channel.traceSync(provider, {})).toBe("result"); - expect(provider).toHaveBeenCalledOnce(); - expect(reportedErrors).toContain(subscriberError); - expect(reportedErrors).toContain(storeError); }); - it("does not repeat provider calls when a store invokes its callback late", () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("late-store"), - ); - const reportedErrors: unknown[] = []; - restoreErrorReporter = setGlobalHookErrorReporter((error) => - reportedErrors.push(error), + it("lets callers provide async context without owning a store", async () => { + const invocation = hook(); + const storage = new AsyncLocalStorage(); + invocation.intercept((next, receiver, args) => + storage.run("wrapped", () => Reflect.apply(next, receiver, args)), ); - - let callback: (() => unknown) | undefined; - channel.start.bindStore({ - getStore() { - return undefined; - }, - run(_store: unknown, next: () => T): T { - callback = next; - return undefined as T; + await invocation.invoke( + async () => { + await Promise.resolve(); + expect(storage.getStore()).toBe("wrapped"); }, - }); - - const provider = vi.fn(() => "result"); - expect(channel.traceSync(provider, {})).toBe("result"); - expect(provider).toHaveBeenCalledOnce(); - expect(callback?.()).toBe("result"); - expect(provider).toHaveBeenCalledOnce(); - expect(reportedErrors.map(String)).toEqual([ - "Error: Instrumentation store did not invoke its callback", - "Error: Instrumentation store invoked its callback more than once", - ]); - }); - - it("replaces malformed entries without executing them", () => { - const name = uniqueChannelName("malformed-entry"); - const registry = Object.getOwnPropertyDescriptor( - globalThis, - GLOBAL_INSTRUMENTATION_HOOKS_KEY, - )?.value as Map; - registry.set(name, { hasSubscribers: true }); - - const reportedErrors: unknown[] = []; - restoreErrorReporter = setGlobalHookErrorReporter((error) => - reportedErrors.push(error), + undefined, + [], + {}, ); - const channel = newGlobalTracingChannel(name); - const subscriber = vi.fn(); - channel.subscribe({ start: subscriber }); - const provider = vi.fn(() => "result"); - - expect(channel.traceSync(provider, {})).toBe("result"); - expect(provider).toHaveBeenCalledOnce(); - expect(subscriber).toHaveBeenCalledOnce(); - expect(reportedErrors.map(String)).toEqual([ - `Error: Invalid global instrumentation hook: ${name}`, - ]); - }); - - it("replaces unbranded hook-shaped entries without mutating them", () => { - const name = uniqueChannelName("unbranded-entry"); - const registry = Object.getOwnPropertyDescriptor( - globalThis, - GLOBAL_INSTRUMENTATION_HOOKS_KEY, - )?.value as Map; - const foreignChannel = { - bindStore: vi.fn(), - hasSubscribers: false, - publish: vi.fn(), - runStores: vi.fn(), - subscribe: vi.fn(), - unbindStore: vi.fn(), - unsubscribe: vi.fn(), - }; - const foreignHook = { - asyncEnd: foreignChannel, - asyncStart: foreignChannel, - end: foreignChannel, - error: foreignChannel, - hasSubscribers: false, - start: foreignChannel, - subscribe: vi.fn(), - traceCallback: vi.fn(), - tracePromise: vi.fn(), - traceSync: vi.fn(), - unsubscribe: vi.fn(), - }; - registry.set(name, foreignHook); - - const channel = newGlobalTracingChannel(name); - - expect(channel).not.toBe(foreignHook); - expect(registry.get(name)).toBe(channel); - expect( - (foreignHook as Record)[ - Symbol.for(GLOBAL_INSTRUMENTATION_HOOK_BRAND) - ], - ).toBeUndefined(); + expect(storage.getStore()).toBeUndefined(); }); - it("wraps callbacks without changing arguments or receiver semantics", async () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("callback"), - ); - const lifecycle: string[] = []; - channel.subscribe({ - start: () => lifecycle.push("start"), - end: () => lifecycle.push("end"), - asyncStart: () => lifecycle.push("asyncStart"), - asyncEnd: () => lifecycle.push("asyncEnd"), + it("uses a stable interceptor snapshot during a call", () => { + const invocation = hook(); + const seen: string[] = []; + let removeSecond: () => void; + invocation.intercept((next) => { + removeSecond(); + return next(); }); - - const receiver = { label: "receiver" }; - const result = await new Promise((resolve) => { - channel.traceCallback( - function ( - this: typeof receiver, - value: string, - callback: (error: unknown, value: string) => void, - ) { - expect(this).toBe(receiver); - callback.call(this, null, value); - return "immediate"; - }, - 1, - {}, - receiver, - "done", - function (this: typeof receiver, error: unknown, value: string) { - expect(this).toBe(receiver); - expect(error).toBeNull(); - resolve(value); - }, - ); + removeSecond = invocation.intercept((next) => { + seen.push("second"); + return next(); }); - - expect(result).toBe("done"); - expect(lifecycle).toEqual(["start", "asyncStart", "asyncEnd", "end"]); - }); - - it("runs traced functions inside bound stores", () => { - const channel = newGlobalTracingChannel>( - uniqueChannelName("stores"), - ); - const storage = new AsyncLocalStorage(); - channel.start.bindStore(storage, () => "bound"); - - expect(channel.hasSubscribers).toBe(true); - expect(channel.traceSync(() => storage.getStore(), {})).toBe("bound"); - expect(channel.start.unbindStore(storage)).toBe(true); - expect(channel.hasSubscribers).toBe(false); + invocation.invoke(() => {}, undefined, [], {}); + invocation.invoke(() => {}, undefined, [], {}); + expect(seen).toEqual(["second"]); }); }); diff --git a/js/src/global-instrumentation-hooks.ts b/js/src/global-instrumentation-hooks.ts index 119d1bfd5..c4d8189d7 100644 --- a/js/src/global-instrumentation-hooks.ts +++ b/js/src/global-instrumentation-hooks.ts @@ -1,73 +1,14 @@ -/* - * Adapted from Node.js diagnostics_channel's TracingChannel implementation. - * Copyright Node.js contributors. Licensed under the MIT License. - * See licenses/node-diagnostics-channel/LICENSE. - */ - export const GLOBAL_INSTRUMENTATION_HOOKS_KEY = - "__braintrust_instrumentation_hooks"; -export const GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION = 1; + "__braintrust_invocation_hooks_v2"; +export const GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION = 2; export const GLOBAL_INSTRUMENTATION_HOOKS_REGISTRY_BRAND = "braintrust.global-instrumentation-hooks.registry"; -export const GLOBAL_INSTRUMENTATION_HOOK_BRAND = - "braintrust.global-instrumentation-hooks.hook"; export const GLOBAL_INVOCATION_HOOK_BRAND = "braintrust.global-instrumentation-hooks.invocation-hook"; const registryBrand = Symbol.for(GLOBAL_INSTRUMENTATION_HOOKS_REGISTRY_BRAND); -const hookBrand = Symbol.for(GLOBAL_INSTRUMENTATION_HOOK_BRAND); const invocationHookBrand = Symbol.for(GLOBAL_INVOCATION_HOOK_BRAND); -export interface GlobalHookAsyncLocalStorage { - run(store: T | undefined, callback: () => R): R; - getStore(): T | undefined; -} - -type GlobalHookMessageFunction = ( - message: M, - name: N, -) => void; - -type GlobalHookTransformFunction = (message: M) => S; - -export interface GlobalHookChannel< - M = any, - N extends string | symbol = string, -> { - readonly name: N; - readonly hasSubscribers: boolean; - subscribe(subscription: GlobalHookMessageFunction): void; - unsubscribe(subscription: GlobalHookMessageFunction): boolean; - bindStore( - store: GlobalHookAsyncLocalStorage, - transform?: GlobalHookTransformFunction, - ): void; - unbindStore(store: GlobalHookAsyncLocalStorage): boolean; - publish(message: M): void; - runStores any>( - message: M, - fn: F, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType; -} - -export interface GlobalTracingChannelCollection { - readonly start?: GlobalHookChannel; - readonly end?: GlobalHookChannel; - readonly asyncStart?: GlobalHookChannel; - readonly asyncEnd?: GlobalHookChannel; - readonly error?: GlobalHookChannel; -} - -export interface GlobalHookHandlers { - start?: (context: M, name: string) => void; - end?: (context: M, name: string) => void; - asyncStart?: (context: M, name: string) => void; - asyncEnd?: (context: M, name: string) => void; - error?: (context: M, name: string) => void; -} - type GlobalInvocationTarget = (this: any, ...args: any[]) => any; export type GlobalInvocationInterceptor = ( @@ -88,55 +29,6 @@ export interface GlobalInvocationHook { ): ReturnType; } -export type GlobalTraceOperator = - | "traceCallback" - | "tracePromise" - | "traceSync"; - -export interface GlobalTracingChannel - extends GlobalTracingChannelCollection, GlobalInvocationHook { - readonly start: GlobalHookChannel; - readonly end: GlobalHookChannel; - readonly asyncStart: GlobalHookChannel; - readonly asyncEnd: GlobalHookChannel; - readonly error: GlobalHookChannel; - readonly hasSubscribers: boolean; - subscribe(handlers: GlobalHookHandlers): void; - unsubscribe(handlers: GlobalHookHandlers): boolean; - traceSync any>( - fn: F, - message?: M, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType; - tracePromise PromiseLike>( - fn: F, - message?: M, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType; - traceCallback any>( - fn: F, - position?: number, - message?: M, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType; - traceInvocation( - operator: GlobalTraceOperator, - target: F, - thisArg: ThisParameterType, - args: Parameters, - additional: A, - callbackIndex?: number, - ): ReturnType; -} - -type StoreEntry = [ - GlobalHookAsyncLocalStorage, - GlobalHookTransformFunction | undefined, -]; - let errorReporter: ((error: unknown) => void) | undefined; export function setGlobalHookErrorReporter( @@ -159,198 +51,13 @@ function reportError(error: unknown): void { } } -function setContextValue( - context: Record, - key: "error" | "result", - value: unknown, -): void { - try { - context[key] = value; - } catch (error) { - reportError(error); - } -} - -function wrapStoreRun( - store: GlobalHookAsyncLocalStorage, - message: M, - next: () => unknown, - transform?: GlobalHookTransformFunction, -): () => unknown { - return () => { - let context: unknown; - try { - context = transform ? transform(message) : message; - } catch (error) { - reportError(error); - return next(); - } - - let called = false; - let result: unknown; - let providerError: unknown; - let providerThrew = false; - const runNext = () => { - if (called) { - reportError( - new Error( - "Instrumentation store invoked its callback more than once", - ), - ); - if (providerThrew) { - throw providerError; - } - return result; - } - - called = true; - try { - result = next(); - return result; - } catch (error) { - providerThrew = true; - providerError = error; - throw error; - } - }; - - try { - store.run(context, runNext); - } catch (error) { - if (!providerThrew || error !== providerError) { - reportError(error); - } - } - - if (!called) { - reportError( - new Error("Instrumentation store did not invoke its callback"), - ); - return runNext(); - } - if (providerThrew) { - throw providerError; - } - return result; - }; -} - -class HookChannel< - M, - N extends string | symbol = string, -> implements GlobalHookChannel { - private subscribers: GlobalHookMessageFunction[] = []; - private stores = new Map< - GlobalHookAsyncLocalStorage, - GlobalHookTransformFunction | undefined - >(); - - constructor(readonly name: N) {} - - get hasSubscribers(): boolean { - return this.subscribers.length > 0 || this.stores.size > 0; - } - - subscribe(subscription: GlobalHookMessageFunction): void { - if (typeof subscription !== "function") { - throw new TypeError("subscription must be a function"); - } - this.subscribers = [...this.subscribers, subscription]; - } - - unsubscribe(subscription: GlobalHookMessageFunction): boolean { - const index = this.subscribers.indexOf(subscription); - if (index === -1) { - return false; - } - this.subscribers = [ - ...this.subscribers.slice(0, index), - ...this.subscribers.slice(index + 1), - ]; - return true; - } - - bindStore( - store: GlobalHookAsyncLocalStorage, - transform?: GlobalHookTransformFunction, - ): void { - if (!store || typeof store.run !== "function") { - throw new TypeError("store must have a run method"); - } - this.stores.set( - store as GlobalHookAsyncLocalStorage, - transform as GlobalHookTransformFunction | undefined, - ); - } - - unbindStore(store: GlobalHookAsyncLocalStorage): boolean { - return this.stores.delete(store as GlobalHookAsyncLocalStorage); - } - - publish(message: M): void { - const subscribers = this.subscribers; - for (const subscriber of subscribers) { - try { - subscriber(message, this.name); - } catch (error) { - reportError(error); - } - } - } - - runStores any>( - message: M, - fn: F, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType { - let run = () => { - this.publish(message); - return Reflect.apply(fn, thisArg, args); - }; - for (const [store, transform] of this.stores.entries() as Iterable< - StoreEntry - >) { - run = wrapStoreRun(store, message, run, transform); - } - return run() as ReturnType; - } -} - -const traceEvents = [ - "start", - "end", - "asyncStart", - "asyncEnd", - "error", -] as const; - -function traceInvocation( - hook: GlobalTracingChannel, - operator: GlobalTraceOperator, - target: F, - thisArg: ThisParameterType, - args: Parameters, - additional: unknown, - callbackIndex = -1, -): ReturnType { - const context = { - ...(additional as object), - arguments: args, - self: thisArg, - }; - const invoke = () => hook.invoke(target, thisArg, args, additional); - - if (operator === "traceCallback") { - return hook.traceCallback(invoke, callbackIndex, context) as ReturnType; - } - if (operator === "tracePromise") { - return hook.tracePromise(invoke, context) as ReturnType; +class InvocationHook implements GlobalInvocationHook { + constructor() { + Object.defineProperty(this, invocationHookBrand, { + value: GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, + }); } - return hook.traceSync(invoke, context) as ReturnType; -} -class InvocationHook implements GlobalInvocationHook { private interceptors: GlobalInvocationInterceptor[] = []; get hasInterceptors(): boolean { @@ -402,394 +109,18 @@ class InvocationHook implements GlobalInvocationHook { } } -class TracingHook implements GlobalTracingChannel { - readonly start: GlobalHookChannel; - readonly end: GlobalHookChannel; - readonly asyncStart: GlobalHookChannel; - readonly asyncEnd: GlobalHookChannel; - readonly error: GlobalHookChannel; - private readonly invocationHook = new InvocationHook(); - - constructor(nameOrChannels: string | GlobalTracingChannelCollection) { - Object.defineProperty(this, hookBrand, { - configurable: false, - enumerable: false, - value: GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, - writable: false, - }); - Object.defineProperty(this, invocationHookBrand, { - configurable: false, - enumerable: false, - value: GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, - writable: false, - }); - - if (typeof nameOrChannels === "string") { - this.start = new HookChannel(`tracing:${nameOrChannels}:start`); - this.end = new HookChannel(`tracing:${nameOrChannels}:end`); - this.asyncStart = new HookChannel(`tracing:${nameOrChannels}:asyncStart`); - this.asyncEnd = new HookChannel(`tracing:${nameOrChannels}:asyncEnd`); - this.error = new HookChannel(`tracing:${nameOrChannels}:error`); - return; - } - - this.start = nameOrChannels.start ?? new HookChannel("tracing:start"); - this.end = nameOrChannels.end ?? new HookChannel("tracing:end"); - this.asyncStart = - nameOrChannels.asyncStart ?? new HookChannel("tracing:asyncStart"); - this.asyncEnd = - nameOrChannels.asyncEnd ?? new HookChannel("tracing:asyncEnd"); - this.error = nameOrChannels.error ?? new HookChannel("tracing:error"); - } - - get hasSubscribers(): boolean { - return ( - this.start.hasSubscribers || - this.end.hasSubscribers || - this.asyncStart.hasSubscribers || - this.asyncEnd.hasSubscribers || - this.error.hasSubscribers - ); - } - - get hasInterceptors(): boolean { - return this.invocationHook.hasInterceptors; - } - - intercept(interceptor: GlobalInvocationInterceptor): () => void { - return this.invocationHook.intercept(interceptor); - } - +type HookRegistry = Map; +const inertInvocationHook: GlobalInvocationHook = Object.freeze({ + hasInterceptors: false, + intercept: () => () => {}, invoke( target: F, thisArg: ThisParameterType, args: Parameters, - additional: unknown, - ): ReturnType { - return this.invocationHook.invoke(target, thisArg, args, additional); - } - - traceInvocation( - operator: GlobalTraceOperator, - target: F, - thisArg: ThisParameterType, - args: Parameters, - additional: unknown, - callbackIndex = -1, ): ReturnType { - return traceInvocation( - this, - operator, - target, - thisArg, - args, - additional, - callbackIndex, - ); - } - - subscribe(handlers: GlobalHookHandlers): void { - for (const eventName of traceEvents) { - const handler = handlers[eventName]; - if (handler) { - this[eventName].subscribe(handler); - } - } - } - - unsubscribe(handlers: GlobalHookHandlers): boolean { - let done = true; - for (const eventName of traceEvents) { - const handler = handlers[eventName]; - if (handler && !this[eventName].unsubscribe(handler)) { - done = false; - } - } - return done; - } - - traceSync any>( - fn: F, - message: M = {} as M, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType { - if (!this.hasSubscribers) { - return Reflect.apply(fn, thisArg, args) as ReturnType; - } - - const context = message as Record; - return this.start.runStores(message, () => { - try { - const result = Reflect.apply(fn, thisArg, args); - setContextValue(context, "result", result); - return result; - } catch (error) { - setContextValue(context, "error", error); - this.error.publish(message); - throw error; - } finally { - this.end.publish(message); - } - }); - } - - tracePromise PromiseLike>( - fn: F, - message: M = {} as M, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType { - if (!this.hasSubscribers) { - return Reflect.apply(fn, thisArg, args) as ReturnType; - } - - const context = message as Record; - return this.start.runStores(message, () => { - let result: ReturnType; - try { - result = Reflect.apply(fn, thisArg, args) as ReturnType; - } catch (error) { - setContextValue(context, "error", error); - this.error.publish(message); - this.end.publish(message); - throw error; - } - - this.end.publish(message); - - if ( - !result || - (typeof result !== "object" && typeof result !== "function") - ) { - setContextValue(context, "result", result); - this.asyncStart.publish(message); - this.asyncEnd.publish(message); - return result; - } - - let terminalPublished = false; - const finishUnobservedResult = (error: unknown) => { - reportError(error); - if (!terminalPublished) { - terminalPublished = true; - setContextValue(context, "result", result); - this.asyncStart.publish(message); - this.asyncEnd.publish(message); - } - return result; - }; - - let then: unknown; - try { - then = result.then; - } catch (error) { - return finishUnobservedResult(error); - } - - if (typeof then !== "function") { - setContextValue(context, "result", result); - this.asyncStart.publish(message); - this.asyncEnd.publish(message); - return result; - } - - const resolve = (resolved: unknown) => { - if (!terminalPublished) { - terminalPublished = true; - setContextValue(context, "result", resolved); - this.asyncStart.publish(message); - this.asyncEnd.publish(message); - } - return resolved; - }; - let rejectionThrown = false; - let rejectionError: unknown; - const reject = (error: unknown) => { - if (!terminalPublished) { - terminalPublished = true; - setContextValue(context, "error", error); - this.error.publish(message); - this.asyncStart.publish(message); - this.asyncEnd.publish(message); - } - rejectionThrown = true; - rejectionError = error; - throw error; - }; - - let isPlainPromise: boolean; - try { - isPlainPromise = - result instanceof Promise && result.constructor === Promise; - } catch (error) { - return finishUnobservedResult(error); - } - - try { - if (isPlainPromise) { - return Reflect.apply(then, result, [resolve, reject]); - } - - Reflect.apply(then, result, [ - resolve, - (error: unknown) => { - try { - reject(error); - } catch { - // The original promise-like object is returned below. Keep the - // instrumentation side-chain from changing its rejection behavior. - } - }, - ]); - } catch (error) { - if (rejectionThrown && error === rejectionError) { - return result; - } - return finishUnobservedResult(error); - } - return result; - }) as ReturnType; - } - - traceCallback any>( - fn: F, - position = -1, - message: M = {} as M, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType { - if (!this.hasSubscribers) { - return Reflect.apply(fn, thisArg, args); - } - - const context = message as Record; - const callArgs = - args.length > 0 - ? args - : ((context.arguments as ArrayLike | undefined) ?? args); - const callback = Array.prototype.at.call(callArgs, position); - if (typeof callback !== "function") { - return Reflect.apply(fn, thisArg, args); - } - - const { asyncStart, asyncEnd, error: errorChannel } = this; - function wrappedCallback(this: unknown, error: unknown, result: unknown) { - if (error) { - setContextValue(context, "error", error); - errorChannel.publish(message); - } else { - setContextValue(context, "result", result); - } - - return asyncStart.runStores(message, () => { - try { - return Reflect.apply(callback, this, arguments); - } finally { - asyncEnd.publish(message); - } - }); - } - - Array.prototype.splice.call(callArgs, position, 1, wrappedCallback); - return this.start.runStores(message, () => { - try { - return Reflect.apply(fn, thisArg, args); - } catch (error) { - setContextValue(context, "error", error); - this.error.publish(message); - throw error; - } finally { - this.end.publish(message); - } - }); - } -} - -type HookRegistry = Map; - -const inertChannel: GlobalHookChannel = Object.freeze({ - name: "braintrust:inert", - hasSubscribers: false, - subscribe() {}, - unsubscribe() { - return false; - }, - bindStore() {}, - unbindStore() { - return false; - }, - publish() {}, - runStores any>( - _message: unknown, - fn: F, - thisArg?: ThisParameterType, - ...args: Parameters - ): ReturnType { - return Reflect.apply(fn, thisArg, args) as ReturnType; + return Reflect.apply(target, thisArg, args); }, }); -const inertTracingHook = new TracingHook({ - start: inertChannel, - end: inertChannel, - asyncStart: inertChannel, - asyncEnd: inertChannel, - error: inertChannel, -}); - -function isHookChannel(value: unknown): value is GlobalHookChannel { - if ( - (typeof value !== "object" && typeof value !== "function") || - value === null - ) { - return false; - } - - try { - const channel = value as GlobalHookChannel; - return ( - typeof channel.hasSubscribers === "boolean" && - typeof channel.subscribe === "function" && - typeof channel.unsubscribe === "function" && - typeof channel.bindStore === "function" && - typeof channel.unbindStore === "function" && - typeof channel.publish === "function" && - typeof channel.runStores === "function" - ); - } catch { - return false; - } -} - -function hasTracingHookShape( - value: unknown, -): value is GlobalTracingChannel { - if ( - (typeof value !== "object" && typeof value !== "function") || - value === null - ) { - return false; - } - - try { - const hook = value as GlobalTracingChannel; - return ( - typeof hook.hasSubscribers === "boolean" && - typeof hook.subscribe === "function" && - typeof hook.unsubscribe === "function" && - typeof hook.traceSync === "function" && - typeof hook.tracePromise === "function" && - typeof hook.traceCallback === "function" && - isHookChannel(hook.start) && - isHookChannel(hook.end) && - isHookChannel(hook.asyncStart) && - isHookChannel(hook.asyncEnd) && - isHookChannel(hook.error) - ); - } catch { - return false; - } -} function hasInvocationHookShape( value: unknown, @@ -815,64 +146,6 @@ function hasInvocationHookShape( } } -function installInvocationHook(value: GlobalTracingChannel): void { - const hasInvocationHook = hasInvocationHookShape(value); - const invocationHook = new InvocationHook(); - try { - if (!hasInvocationHook) { - Object.defineProperties(value, { - [invocationHookBrand]: { - configurable: false, - enumerable: false, - value: GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, - writable: false, - }, - hasInterceptors: { - configurable: false, - enumerable: false, - get: () => invocationHook.hasInterceptors, - }, - intercept: { - configurable: false, - enumerable: false, - value: invocationHook.intercept.bind(invocationHook), - writable: false, - }, - invoke: { - configurable: false, - enumerable: false, - value: invocationHook.invoke.bind(invocationHook), - writable: false, - }, - }); - } - if (typeof value.traceInvocation !== "function") { - Object.defineProperty(value, "traceInvocation", { - configurable: false, - enumerable: false, - value: traceInvocation.bind(undefined, value), - writable: false, - }); - } - } catch (error) { - reportError(error); - } -} - -function isCompatibleTracingHook( - value: unknown, -): value is GlobalTracingChannel { - try { - return ( - (value as unknown as Record)[hookBrand] === - GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION && - hasTracingHookShape(value) - ); - } catch { - return false; - } -} - function isCompatibleHookRegistry(value: unknown): value is HookRegistry { try { return ( @@ -951,46 +224,17 @@ function getHookRegistry(): HookRegistry | undefined { } } -export function newGlobalTracingChannel( - nameOrChannels: string | GlobalTracingChannelCollection, -): GlobalTracingChannel { - if (typeof nameOrChannels !== "string") { - return new TracingHook(nameOrChannels); - } - +export function newGlobalInvocationHook( + name: string, +): GlobalInvocationHook { const registry = getHookRegistry(); - if (!registry) { - return inertTracingHook as GlobalTracingChannel; - } - let existing: unknown; - try { - existing = Map.prototype.get.call(registry, nameOrChannels); - } catch (error) { - reportError(error); - return inertTracingHook as GlobalTracingChannel; - } - if (isCompatibleTracingHook(existing)) { - installInvocationHook(existing); - return existing as GlobalTracingChannel; - } + if (!registry) return inertInvocationHook; + const existing = Map.prototype.get.call(registry, name); + if (hasInvocationHookShape(existing)) return existing; if (existing !== undefined) { - reportError( - new Error(`Invalid global instrumentation hook: ${nameOrChannels}`), - ); - try { - Map.prototype.delete.call(registry, nameOrChannels); - } catch (error) { - reportError(error); - return inertTracingHook as GlobalTracingChannel; - } - } - - const hook = new TracingHook(nameOrChannels); - try { - Map.prototype.set.call(registry, nameOrChannels, hook); - } catch (error) { - reportError(error); - return inertTracingHook as GlobalTracingChannel; + reportError(new Error(`Invalid global invocation hook: ${name}`)); } + const hook = new InvocationHook(); + Map.prototype.set.call(registry, name, hook); return hook; } diff --git a/js/src/instrumentation/README.md b/js/src/instrumentation/README.md index dfa3cd042..389579075 100644 --- a/js/src/instrumentation/README.md +++ b/js/src/instrumentation/README.md @@ -1,184 +1,86 @@ # Writing Braintrust Instrumentation Plugins -Braintrust instrumentation plugins wrap provider calls through typed invocation -hooks or consume tracing-compatible events from the internal global registry. -Auto-instrumented provider code and manual wrappers use the same typed channels, -so extraction, stream handling, and span behavior stay aligned. +API wrapping and span instrumentation are separate layers. +Invocation hooks work independently of the SDK; provider plugins explicitly connect them to tracing functions. -## Architecture - -An instrumentation has four parts: - -1. An Orchestrion config identifies the provider function for automatic - transformation. -2. A typed channel defines its arguments, result, extra event fields, and stable - `orchestrion::` identifier. -3. A plugin intercepts that channel, or subscribes to its legacy tracing - lifecycle, and maps the call into Braintrust spans. -4. A manual wrapper invokes the same typed channel when transformation is not - available. - -The global hook transport is internal. New and migrated plugins should prefer -the typed channel's `intercept` API. Existing plugins can continue using -`traceAsyncChannel`, `traceStreamingChannel`, `traceSyncStreamChannel`, or -`BasePlugin` helpers during the gradual migration. - -## Invocation Hooks - -Invocation hooks expose the complete target call and are designed for wrappers -that need to scope the target with `AsyncLocalStorage.run()` or patch arguments, -the receiver, or the returned value: - -```ts -const removeInterceptor = providerChannels.create.intercept( - (target, thisArg, args, additional) => - store.run(additional.context, () => target.apply(thisArg, args)), -); -``` - -Interceptors compose as nested middleware in registration order. Each receives -the next target and may invoke it with different arguments or receiver, replace -its result, call it more than once, or not call it at all. Removing an -interceptor is idempotent. Calling a typed channel through `invoke` only uses -the invocation hook and does not dispatch the legacy tracing lifecycle. -Generated auto-instrumentation wrappers separately retain dual emission during -the migration, with legacy tracing outside the effective intercepted call. - -## Lifecycle - -The event lifecycle is compatible with Node tracing channels: - -- `start`: before the synchronous portion of the target function -- `end`: after the synchronous portion completes -- `asyncStart`: when an asynchronous result begins settling -- `asyncEnd`: after that result settles and before user continuation -- `error`: when the target throws, rejects, or reports a callback error - -Every phase receives the same mutable context object: +## Define wrapping hooks ```ts -interface InstrumentationContext { - arguments: ArrayLike; - self?: unknown; - moduleVersion?: string; - result?: unknown; - error?: unknown; -} +const providerHooks = defineInterceptor("provider-package", { + create: channel<[CreateParams], PromiseLike>({ + channelName: "messages.create", + }), +}); ``` -The generated wrapper passes the target invocation and `moduleVersion` to the -global hook runtime. That runtime creates `arguments` and `self`; tracing -operators add `result` or `error`. +The identifier must match the Orchestrion config: `orchestrion::`. +Definitions describe arguments, return values, and opaque additional data. +They do not contain span names, instrumentation provenance, or tracing methods. -## Defining Typed Channels - -Define the smallest types needed by instrumentation: +Interceptors compose in registration order and may replace arguments, receivers, results, or the entire call. +Removing an interceptor is idempotent. +Wrapping can scope a call without creating spans: ```ts -const providerChannels = defineChannels( - "provider-package", - { - create: channel< - [CreateParams], - CreateResult, - { providerRequestId?: string } - >({ - channelName: "messages.create", - kind: "async", - }), - }, - { instrumentationName: "provider" }, +const remove = providerHooks.create.intercept((target, receiver, args) => + store.run(context, () => Reflect.apply(target, receiver, args)), ); ``` -Channel names must match the Orchestrion config exactly. Do not include the -`orchestrion:` prefix in the transform config; `defineChannels` and Orchestrion -construct it from the package and operation. - -## Subscribing +## Trace independently -Prefer the shared tracing helpers: +Tracing helpers receive a callable and tracing data; they never receive a hook or register interceptors. +Provider plugins explicitly connect the layers: ```ts -this.register( - traceAsyncChannel(providerChannels.create, { - name: "provider.messages.create", - type: "llm", - extractInput(args) { - return { - input: args[0].messages, - metadata: { model: args[0].model }, - }; - }, - extractOutput(result) { - return result.content; - }, - extractMetrics(result) { - return { - prompt_tokens: result.usage.input_tokens, - completion_tokens: result.usage.output_tokens, - }; - }, - }), +this.unsubscribers.push( + providerHooks.create.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + name: "provider.messages.create", + instrumentationName: INSTRUMENTATION_NAMES.PROVIDER, + type: "llm", + extractInput: ([params]) => ({ + input: params.messages, + metadata: { model: params.model }, + }), + extractOutput: (result) => result.content, + extractMetrics: (result) => ({ tokens: result.usage.totalTokens }), + }, + ), + ), ); ``` -The helpers: - -- create and correlate spans with a `WeakMap` keyed by event context -- bind the current span store to `start` for async-context propagation -- contain extraction failures and log them through `debugLogger` -- patch streams without replacing their public semantics -- unsubscribe and unbind stores when a plugin is disabled - -Use raw `IsoChannelHandlers` only when a provider requires lifecycle behavior -that the shared helpers cannot express. - -## Manual Wrappers - -New and migrated manual wrappers pass the original target, receiver, arguments, -and any channel-specific fields to the same typed channel: +The tracing function owns span creation, context propagation, response observation, and finalization. +It preserves the original return value, including Promise subclasses and stream identity. +Do not introduce combined registration APIs such as `interceptAndTrace` or `traceInvocation`. -```ts -return providerChannels.create.invoke(originalCreate, this, [params], { - providerRequestId, -}); -``` +## Manual wrappers -Legacy wrappers can continue calling the tracing-compatible operators until -their plugin is migrated: +Manual and generated wrappers invoke the same hook: ```ts -return providerChannels.create.tracePromise(() => originalCreate(params), { - arguments: [params], -}); +return providerHooks.create.invoke(originalCreate, client, [params], {}); ``` -Do not create spans directly inside wrappers. Keeping span creation in the -plugin prevents auto and manual instrumentation from drifting. - -## Promise and Stream Requirements +Without an interceptor the original function runs directly. +With a tracing plugin enabled the registered tracing function creates spans. +Manual wrappers must not create spans themselves. -Instrumentation is non-invasive: +## Runtime and safety requirements -- Native promises retain normal resolution and rejection behavior. -- Promise subclasses and other thenables are returned unchanged so helper - methods such as `withResponse()` remain available. -- A non-Promise value returned from an `Async` transform remains that value. -- Async iterables and event-emitter streams retain identity and public methods. -- Subscriber or extraction bugs must not alter provider calls. - -Stream patches must be idempotent and preserve cancellation, errors, early -termination, and async context. - -## Event and Span Safety - -- Treat arguments, results, metadata, and headers as untrusted. -- Avoid prototype-sensitive merges and unnecessary mutation of provider data. -- Capture only fields permitted by the instrumentation specification. +- Preserve receivers, arguments, errors, async context, Promise helper methods, and stream cancellation. +- Treat provider inputs and outputs as untrusted and capture only specification-permitted data. +- Keep registration, removal, and stream patching idempotent. +- Contain extraction failures with `debugLogger` without retrying the provider call. - Pass `Error` objects directly to `span.log({ error })`. -- Use narrow vendored provider interfaces shared by wrappers and plugins. -- Keep enable, disable, subscription, and patching behavior idempotent. +- Treat `span.log()` and `span.end()` as non-throwing. + +Invocation protocol version 2 uses an independent global registry. +Rebuild bundles transformed with the previous SDK when upgrading; the old combined tracing protocol is not supported. ## Export Customizers @@ -269,7 +171,7 @@ Clearing customizers is always allowed and silent. Test at the narrowest useful layers: 1. Plugin unit tests for extraction and span handling. -2. Global hook/runtime tests for lifecycle and context behavior. +2. Invocation runtime tests for wrapping and context behavior. 3. Orchestrion transformation tests for generated wrappers. 4. Bundler and loader tests for real transformed execution. 5. Provider e2e tests for wrapped and auto-hook parity. diff --git a/js/src/instrumentation/core/channel-definitions.test.ts b/js/src/instrumentation/core/channel-definitions.test.ts index 383bcdf74..99f70e0e2 100644 --- a/js/src/instrumentation/core/channel-definitions.test.ts +++ b/js/src/instrumentation/core/channel-definitions.test.ts @@ -1,77 +1,27 @@ import { randomUUID } from "node:crypto"; -import { describe, expect, it, vi } from "vitest"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; -import { channel, defineChannels } from "./channel-definitions"; +import { expect, it } from "vitest"; +import { channel, defineInterceptor } from "./channel-definitions"; -describe("typed channel invocation hooks", () => { - it("does not dispatch tracing subscribers from invoke", async () => { - const channels = defineChannels( - `typed-channel-async-${randomUUID()}`, - { - call: channel<[number], number, { label?: string }>({ - channelName: "call", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.BRAINTRUST_JS_LOGGER }, - ); - const lifecycle: string[] = []; - const tracingChannel = channels.call.tracingChannel(); - const handlers = { - start: (event: { arguments: number[] }) => - lifecycle.push(`start:${event.arguments[0]}`), - asyncEnd: (event: { result?: number }) => - lifecycle.push(`asyncEnd:${event.result}`), - }; - tracingChannel.subscribe(handlers); - const removeInterceptor = channels.call.intercept( - async (target, _thisArg, args, additional) => { - lifecycle.push(`intercept:${additional.label}`); - return (await target.apply({ offset: 4 }, [args[0] + 1])) * 2; - }, - ); - const target = vi.fn(async function ( - this: { offset: number }, - value: number, - ) { - lifecycle.push("target"); - return this.offset + value; - }); - - await expect( - channels.call.invoke(target, { offset: 0 }, [2], { label: "test" }), - ).resolves.toBe(14); - expect(tracingChannel.hasSubscribers).toBe(true); - expect(lifecycle).toEqual(["intercept:test", "target"]); - - removeInterceptor(); - tracingChannel.unsubscribe(handlers); - }); - - it("supports receiver, argument, and output patching for sync channels", () => { - const channels = defineChannels( - `typed-channel-sync-${randomUUID()}`, - { - call: channel<[string], string>({ - channelName: "call", - kind: "sync-stream", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.BRAINTRUST_JS_LOGGER }, - ); - channels.call.intercept((target, _thisArg, args) => - target.apply({ prefix: "patched" }, [args[0].toUpperCase()]), - ); - - expect( - channels.call.invoke( - function (this: { prefix: string }, value: string) { - return `${this.prefix}:${value}`; - }, - { prefix: "original" }, - ["value"], - {}, - ), - ).toBe("patched:VALUE"); +it("defines sync and async wrapping contracts without tracing configuration", async () => { + const hooks = defineInterceptor(randomUUID(), { + sync: channel<[number], number>({ channelName: "sync" }), + async: channel<[string], Promise, { suffix: string }>({ + channelName: "async", + }), }); + expect(hooks.sync).not.toHaveProperty("tracingChannel"); + expect(hooks.sync).not.toHaveProperty("instrumentationName"); + const remove = hooks.sync.intercept( + (target, receiver, [n]) => target.call(receiver, n + 1) * 2, + ); + expect(hooks.sync.invoke((n) => n + 3, undefined, [2], {})).toBe(12); + remove(); + hooks.async.intercept((target, receiver, [s], additional) => + target.call(receiver, s + additional.suffix), + ); + expect( + await hooks.async.invoke(async (s) => s.toUpperCase(), undefined, ["a"], { + suffix: "b", + }), + ).toBe("AB"); }); diff --git a/js/src/instrumentation/core/channel-definitions.ts b/js/src/instrumentation/core/channel-definitions.ts index 57626bde8..24a4a4c2f 100644 --- a/js/src/instrumentation/core/channel-definitions.ts +++ b/js/src/instrumentation/core/channel-definitions.ts @@ -1,351 +1,83 @@ -import iso from "../../isomorph"; -import type { IsoTracingChannel } from "../../isomorph"; -import type { SpanInstrumentationName } from "../../span-origin"; -import type { - AsyncEndEventWith, - EndEventWith, - ErrorEventWith, - EventArguments, - StartEventWith, -} from "./types"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; -export type ChannelKind = "async" | "sync-stream"; - -type ChannelTypeInfo< - TArgs extends EventArguments, +type ChannelSpec< + TArgs extends readonly unknown[], TResult, - TExtra extends object = Record, - TChunk = never, - TKind extends ChannelKind = "async", + TExtra extends object, + TChunk, > = { - kind: TKind; + channelName: string; __args?: TArgs; __result?: TResult; __extra?: TExtra; __chunk?: TChunk; }; - -type ChannelSpec< - TArgs extends EventArguments, - TResult, - TExtra extends object = Record, - TChunk = never, - TKind extends ChannelKind = "async", -> = ChannelTypeInfo & { - channelName: string; -}; - -type AnyAsyncChannelSpec = ChannelSpec< - EventArguments, - unknown, - object, - unknown, - "async" ->; - -type AnySyncStreamChannelSpec = ChannelSpec< - EventArguments, - unknown, - object, - unknown, - "sync-stream" ->; - -type AnyChannelSpec = AnyAsyncChannelSpec | AnySyncStreamChannelSpec; - -export type ArgsOf = - TChannel extends ChannelTypeInfo< - infer TArgs, - unknown, - object, - unknown, - ChannelKind - > - ? [...TArgs] - : never; - -export type ResultOf = - TChannel extends ChannelTypeInfo< - EventArguments, - infer TResult, - object, - unknown, - ChannelKind - > - ? TResult - : never; - -export type ExtraOf = - TChannel extends ChannelTypeInfo< - EventArguments, - unknown, - infer TExtra extends object, - unknown, - ChannelKind - > - ? TExtra - : never; - -export type ChunkOf = - TChannel extends ChannelTypeInfo< - EventArguments, - unknown, - object, - infer TChunk, - ChannelKind - > - ? TChunk - : never; - -export type StartOf = StartEventWith< - ArgsOf, - ExtraOf ->; - -export type AsyncEndOf = AsyncEndEventWith< - ResultOf, - ArgsOf, - ExtraOf ->; - -export type EndOf = EndEventWith< - ResultOf, - ArgsOf, - ExtraOf ->; - -export type ErrorOf = ErrorEventWith< - ArgsOf, - ExtraOf ->; - -export type ChannelMessage = - StartOf & - Partial<{ result: ResultOf }> & - Partial, "error">>; - -type InvocationAdditionalOf = - ExtraOf & { moduleVersion?: string }; - -type InvocationResultOf = - TChannel["kind"] extends "async" - ? PromiseLike> - : ResultOf; - -type ChannelInterceptor = ( - target: ( - this: unknown, - ...args: ArgsOf - ) => InvocationResultOf, - thisArg: unknown, - args: ArgsOf, - additional: InvocationAdditionalOf, -) => InvocationResultOf; - -type InterceptMethod = { - intercept(interceptor: ChannelInterceptor): () => void; -}["intercept"]; - -type BaseTypedChannel = TSpec & { - instrumentationName: SpanInstrumentationName; - tracingChannel(): IsoTracingChannel>; - intercept: InterceptMethod; +type AnySpec = ChannelSpec; +export type ArgsOf = T extends { + __args?: infer A extends readonly unknown[]; +} + ? [...A] + : never; +export type ReturnOf = T extends { __result?: infer R } ? R : never; +export type ExtraOf = T extends { __extra?: infer E extends object } + ? E + : never; +export type ChunkOf = T extends { __chunk?: infer R } ? R : never; +export type InvocationAdditionalOf = ExtraOf & { moduleVersion?: string }; +export type Interceptor = ( + target: (this: unknown, ...args: ArgsOf) => ReturnOf, + receiver: unknown, + args: ArgsOf, + additional: InvocationAdditionalOf, +) => ReturnOf; +export type InvocationChannel = T & { + intercept(interceptor: Interceptor): () => void; + invoke ReturnOf>( + target: F, + receiver: ThisParameterType, + args: Parameters | ArgsOf, + additional: InvocationAdditionalOf, + ): ReturnType; }; -export type TypedAsyncChannel = - BaseTypedChannel & { - invoke>>( - target: (this: TThis, ...args: ArgsOf) => TReturn, - thisArg: TThis, - args: ArgsOf, - additional: InvocationAdditionalOf, - ): TReturn; - tracePromise>>( - fn: () => TReturn, - context: StartOf, - ): TReturn; - }; - -export type TypedSyncStreamChannel = - BaseTypedChannel & { - invoke>( - target: (this: TThis, ...args: ArgsOf) => TReturn, - thisArg: TThis, - args: ArgsOf, - additional: InvocationAdditionalOf, - ): TReturn; - traceSync>( - fn: () => TResult, - context: StartOf, - ): TResult; - }; - -export type AnyAsyncChannel = Omit< - TypedAsyncChannel, - "intercept" | "invoke" ->; -export type AnySyncStreamChannel = Omit< - TypedSyncStreamChannel, - "intercept" | "invoke" ->; - -type ChannelSpecMap = Record; - -export function channel< - TArgs extends EventArguments, - TResult, - TExtra extends object = Record, - TChunk = never, ->(spec: { - channelName: string; - kind: "async"; -}): ChannelSpec; +/** Describe a callable. Additional data is opaque to the wrapping runtime. */ export function channel< - TArgs extends EventArguments, + TArgs extends readonly unknown[], TResult, TExtra extends object = Record, TChunk = never, ->(spec: { - channelName: string; - kind: "sync-stream"; -}): ChannelSpec; -export function channel(spec: { - channelName: string; - kind: ChannelKind; -}): AnyChannelSpec { - return spec as AnyChannelSpec; +>(spec: { channelName: string }): ChannelSpec { + return spec; } -type MaterializedChannel = T["kind"] extends "async" - ? TypedAsyncChannel< - ChannelSpec, ResultOf, ExtraOf, ChunkOf, "async"> - > - : TypedSyncStreamChannel< - ChannelSpec, ResultOf, ExtraOf, ChunkOf, "sync-stream"> - >; - -export function defineChannels( +/** Define invocation hooks without initializing the SDK or enabling tracing. */ +export function defineInterceptor>( pkg: string, - channels: T, - options: { instrumentationName: SpanInstrumentationName }, -): { - [K in keyof T]: MaterializedChannel; -} { - const { instrumentationName } = options; + definitions: T, +): { [K in keyof T]: InvocationChannel } { return Object.fromEntries( - Object.entries(channels).map(([key, spec]) => { - const fullChannelName = `orchestrion:${pkg}:${spec.channelName}`; - if (spec.kind === "async") { - const asyncSpec = spec as ChannelSpec< - ArgsOf, - ResultOf, - ExtraOf, - ChunkOf, - "async" - >; - const tracingChannel = () => - iso.newTracingChannel>( - fullChannelName, - ); - const intercept = ( - interceptor: ChannelInterceptor, - ) => { - const hook = tracingChannel(); - return typeof hook.intercept === "function" - ? hook.intercept(interceptor) - : () => {}; - }; - return [ - key, - { - ...asyncSpec, - instrumentationName, - intercept, - invoke: < - TThis, - TReturn extends PromiseLike>, - >( - target: ( - this: TThis, - ...args: ArgsOf - ) => TReturn, - thisArg: TThis, - args: ArgsOf, - additional: InvocationAdditionalOf, - ) => { - const hook = tracingChannel(); - return ( - typeof hook.invoke === "function" - ? hook.invoke(target, thisArg, args, additional) - : Reflect.apply(target, thisArg, args) - ) as TReturn; - }, - tracingChannel, - tracePromise: >>( - fn: () => TReturn, - context: StartOf, - ) => - tracingChannel().tracePromise( - fn, - // eslint-disable-next-line @typescript-eslint/consistent-type-assertions - context as ChannelMessage, - ) as TReturn, - } as AnyAsyncChannel, - ]; - } - - const syncSpec = spec as ChannelSpec< - ArgsOf, - ResultOf, - ExtraOf, - ChunkOf, - "sync-stream" - >; - const tracingChannel = () => - iso.newTracingChannel>( - fullChannelName, - ); - const intercept = ( - interceptor: ChannelInterceptor, - ) => { - const hook = tracingChannel(); - return typeof hook.intercept === "function" - ? hook.intercept(interceptor) - : () => {}; - }; + Object.entries(definitions).map(([key, spec]) => { + const name = `orchestrion:${pkg}:${spec.channelName}`; return [ key, { - ...syncSpec, - instrumentationName, - intercept, - invoke: >( - target: (this: TThis, ...args: ArgsOf) => TResult, - thisArg: TThis, - args: ArgsOf, - additional: InvocationAdditionalOf, - ) => { - const hook = tracingChannel(); - return ( - typeof hook.invoke === "function" - ? hook.invoke(target, thisArg, args, additional) - : Reflect.apply(target, thisArg, args) - ) as TResult; - }, - tracingChannel, - traceSync: ( - fn: () => TResult, - context: StartOf, + ...spec, + intercept: (interceptor: Interceptor) => + newGlobalInvocationHook(name).intercept(interceptor), + invoke: ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + additional: object, ) => - tracingChannel().traceSync( - fn, - // eslint-disable-next-line @typescript-eslint/consistent-type-assertions - context as ChannelMessage, + newGlobalInvocationHook(name).invoke( + target, + receiver, + args, + additional, ), - } as AnySyncStreamChannel, + }, ]; }), - ) as { - [K in keyof T]: MaterializedChannel; - }; + ) as unknown as { [K in keyof T]: InvocationChannel }; } diff --git a/js/src/instrumentation/core/channel-tracing.test.ts b/js/src/instrumentation/core/channel-tracing.test.ts index 55c08f05b..6a08cddf3 100644 --- a/js/src/instrumentation/core/channel-tracing.test.ts +++ b/js/src/instrumentation/core/channel-tracing.test.ts @@ -8,10 +8,10 @@ import { vi, } from "vitest"; import { + NOOP_SPAN, _exportsForTestingOnly, currentSpan, initLogger, - NOOP_SPAN, type Span, type TestBackgroundLogger, } from "../../logger"; @@ -21,25 +21,20 @@ import { withSpanInstrumentationName, } from "../../span-origin"; import { runWithAutoInstrumentationSuppressed } from "../auto-instrumentation-suppression"; -import { channel, defineChannels } from "./channel-definitions"; -import { traceAsyncChannel, traceStreamingChannel } from "./channel-tracing"; - -const testChannels = defineChannels( - "channel-tracing-test", - { - asyncCall: channel<[Record], { ok: true }>({ - channelName: "async.call", - kind: "async", - }), - streamingCall: channel<[Record], { ok: true }>({ - channelName: "streaming.call", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.OPENAI }, -); - -describe("traceAsyncChannel current span binding", () => { +import { channel, defineInterceptor } from "./channel-definitions"; + +import { traceAsyncCall, traceStreamingCall } from "./channel-tracing"; + +const testChannels = defineInterceptor("channel-tracing-test", { + asyncCall: channel<[Record], PromiseLike<{ ok: true }>>({ + channelName: "async.call", + }), + streamingCall: channel<[Record], PromiseLike<{ ok: true }>>({ + channelName: "streaming.call", + }), +}); + +describe("traceAsyncCall current span binding", () => { let backgroundLogger: TestBackgroundLogger; beforeAll(async () => { @@ -60,21 +55,29 @@ describe("traceAsyncChannel current span binding", () => { }); it("binds the created span into the traced async execution context", async () => { - const unsubscribe = traceAsyncChannel(testChannels.asyncCall, { - name: "channel-tracing-test", - type: "function", - extractInput: () => ({ - input: "input", - metadata: undefined, - }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - }); + const unsubscribe = testChannels.asyncCall.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "channel-tracing-test", + type: "function", + extractInput: () => ({ + input: "input", + metadata: undefined, + }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + }, + ), + ); const seenSpanIds: string[] = []; try { - await testChannels.asyncCall.tracePromise( + await testChannels.asyncCall.invoke( async () => { seenSpanIds.push(currentSpan().spanId); await Promise.resolve(); @@ -82,7 +85,9 @@ describe("traceAsyncChannel current span binding", () => { return { ok: true as const }; }, - { arguments: [{}] } as any, + undefined, + ({ arguments: [{}] } as any).arguments, + {}, ); } finally { unsubscribe(); @@ -103,16 +108,24 @@ describe("traceAsyncChannel current span binding", () => { }); it("limits channel provenance to directly instrumented spans", async () => { - const unsubscribe = traceAsyncChannel(testChannels.asyncCall, { - name: "channel-parent", - type: "function", - extractInput: () => ({ input: "input", metadata: undefined }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - }); + const unsubscribe = testChannels.asyncCall.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "channel-parent", + type: "function", + extractInput: () => ({ input: "input", metadata: undefined }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + }, + ), + ); try { - await testChannels.asyncCall.tracePromise( + await testChannels.asyncCall.invoke( async () => { const parent = currentSpan(); parent.startSpan({ name: "user-child" }).end(); @@ -131,7 +144,9 @@ describe("traceAsyncChannel current span binding", () => { .end(); return { ok: true as const }; }, - { arguments: [{}] } as any, + undefined, + ({ arguments: [{}] } as any).arguments, + {}, ); } finally { unsubscribe(); @@ -161,28 +176,36 @@ describe("traceAsyncChannel current span binding", () => { }); it("does not create a span when shouldTrace returns false", async () => { - const unsubscribe = traceAsyncChannel(testChannels.asyncCall, { - name: "channel-tracing-test", - shouldTrace: ([params]) => - !( - typeof params === "object" && - params !== null && - "skip" in params && - params.skip === true + const unsubscribe = testChannels.asyncCall.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "channel-tracing-test", + shouldTrace: ([params]) => + !( + typeof params === "object" && + params !== null && + "skip" in params && + params.skip === true + ), + type: "function", + extractInput: () => ({ + input: "input", + metadata: undefined, + }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + }, ), - type: "function", - extractInput: () => ({ - input: "input", - metadata: undefined, - }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - }); + ); const seenSpanIds: string[] = []; try { - await testChannels.asyncCall.tracePromise( + await testChannels.asyncCall.invoke( async () => { seenSpanIds.push(currentSpan().spanId); await Promise.resolve(); @@ -190,7 +213,9 @@ describe("traceAsyncChannel current span binding", () => { return { ok: true as const }; }, - { arguments: [{ skip: true }] } as any, + undefined, + ({ arguments: [{ skip: true }] } as any).arguments, + {}, ); } finally { unsubscribe(); @@ -207,24 +232,34 @@ describe("traceAsyncChannel current span binding", () => { const consoleErrorSpy = vi .spyOn(console, "error") .mockImplementation(() => {}); - const unsubscribe = traceAsyncChannel(testChannels.asyncCall, { - name: "channel-tracing-test", - shouldTrace: () => { - throw new Error("predicate failed"); - }, - type: "function", - extractInput: () => ({ - input: "input", - metadata: undefined, - }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - }); + const unsubscribe = testChannels.asyncCall.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "channel-tracing-test", + shouldTrace: () => { + throw new Error("predicate failed"); + }, + type: "function", + extractInput: () => ({ + input: "input", + metadata: undefined, + }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + }, + ), + ); try { - await testChannels.asyncCall.tracePromise( + await testChannels.asyncCall.invoke( async () => ({ ok: true as const }), - { arguments: [{}] } as any, + undefined, + ({ arguments: [{}] } as any).arguments, + {}, ); } finally { unsubscribe(); @@ -238,20 +273,28 @@ describe("traceAsyncChannel current span binding", () => { }); it("skips auto instrumentation spans while suppression is active", async () => { - const unsubscribe = traceAsyncChannel(testChannels.asyncCall, { - name: "channel-tracing-test", - type: "function", - extractInput: () => ({ - input: "input", - metadata: undefined, - }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - }); + const unsubscribe = testChannels.asyncCall.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "channel-tracing-test", + type: "function", + extractInput: () => ({ + input: "input", + metadata: undefined, + }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + }, + ), + ); try { await runWithAutoInstrumentationSuppressed(() => - testChannels.asyncCall.tracePromise( + testChannels.asyncCall.invoke( async () => { expect(currentSpan()).toBe(NOOP_SPAN); await Promise.resolve(); @@ -259,7 +302,9 @@ describe("traceAsyncChannel current span binding", () => { return { ok: true as const }; }, - { arguments: [{}] } as any, + undefined, + ({ arguments: [{}] } as any).arguments, + {}, ), ); } finally { @@ -270,40 +315,50 @@ describe("traceAsyncChannel current span binding", () => { expect(spans).toHaveLength(0); }); - it("runs streaming cleanup hooks when span logging fails", async () => { + it("runs completion and failure hooks and ends spans", async () => { const onComplete = vi.fn(); const onError = vi.fn(); const end = vi.fn(); const child = { end, - log: vi.fn(() => { - throw new Error("logging failed"); - }), + log: vi.fn(), } as unknown as Span; - const unsubscribe = traceStreamingChannel(testChannels.streamingCall, { - name: "streaming-channel-test", - startSpan: () => child, - type: "function", - extractInput: () => ({ input: "input", metadata: undefined }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - onComplete, - onError, - }); + const unsubscribe = testChannels.streamingCall.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "streaming-channel-test", + startSpan: () => child, + type: "function", + extractInput: () => ({ input: "input", metadata: undefined }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + onComplete, + onError, + }, + ), + ); try { await expect( - testChannels.streamingCall.tracePromise( + testChannels.streamingCall.invoke( async () => ({ ok: true as const }), - { arguments: [{}] } as any, + undefined, + ({ arguments: [{}] } as any).arguments, + {}, ), ).resolves.toEqual({ ok: true }); await expect( - testChannels.streamingCall.tracePromise( + testChannels.streamingCall.invoke( async () => { throw new Error("call failed"); }, - { arguments: [{}] } as any, + undefined, + ({ arguments: [{}] } as any).arguments, + {}, ), ).rejects.toThrow("call failed"); } finally { @@ -321,15 +376,23 @@ describe("traceAsyncChannel current span binding", () => { end: vi.fn(), log: vi.fn(), } as unknown as Span; - const unsubscribe = traceStreamingChannel(testChannels.streamingCall, { - name: "streaming-channel-test", - startSpan: () => child, - type: "function", - extractInput: () => ({ input: "input", metadata: undefined }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - onError, - }); + const unsubscribe = testChannels.streamingCall.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "streaming-channel-test", + startSpan: () => child, + type: "function", + extractInput: () => ({ input: "input", metadata: undefined }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + onError, + }, + ), + ); const stream = { abort: vi.fn(), async *[Symbol.asyncIterator]() { @@ -338,9 +401,11 @@ describe("traceAsyncChannel current span binding", () => { }; try { - const patched = await testChannels.streamingCall.tracePromise( + const patched = await testChannels.streamingCall.invoke( async () => stream as any, - { arguments: [{}] } as any, + undefined, + ({ arguments: [{}] } as any).arguments, + {}, ); (patched as unknown as typeof stream).abort(); await Promise.resolve(); diff --git a/js/src/instrumentation/core/channel-tracing.ts b/js/src/instrumentation/core/channel-tracing.ts index a8ecc819c..d736bd145 100644 --- a/js/src/instrumentation/core/channel-tracing.ts +++ b/js/src/instrumentation/core/channel-tracing.ts @@ -1,35 +1,29 @@ import { debugLogger } from "../../debug-logger"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; -import { - _internalGetGlobalState, - BRAINTRUST_CURRENT_SPAN_STORE, - startSpan, -} from "../../logger"; -import type { CurrentSpanStore, Span } from "../../logger"; +import type { Span } from "../../logger"; +import { startSpan, withCurrent } from "../../logger"; import { withSpanInstrumentationName, type SpanInstrumentationName, } from "../../span-origin"; import { getCurrentUnixTimestamp, isObject } from "../../util"; +import { isAutoInstrumentationSuppressed } from "../auto-instrumentation-suppression"; +import type { ArgsOf, ChunkOf, ReturnOf } from "./channel-definitions"; +import { + buildStartSpanArgs, + mergeInputMetadata, + type ChannelConfig, +} from "./channel-tracing-utils"; +import { observeResult, runInstrumentation } from "./observe-result"; +import { isAsyncIterable, patchStreamIfNeeded } from "./stream-patcher"; import type { AnyAsyncChannel, AnySyncStreamChannel, - ArgsOf, AsyncEndOf, - ChannelMessage, - ChunkOf, EndOf, ErrorOf, ResultOf, StartOf, -} from "./channel-definitions"; -import { isAsyncIterable, patchStreamIfNeeded } from "./stream-patcher"; -import { - buildStartSpanArgs, - mergeInputMetadata, - type ChannelConfig, -} from "./channel-tracing-utils"; -import { isAutoInstrumentationSuppressed } from "../auto-instrumentation-suppression"; +} from "./tracing-types"; type SpanState = { span: Span; @@ -144,7 +138,7 @@ type SyncStreamChannelSpanConfig = patchResult?: (args: { channelName: string; endEvent: EndOf; - result: ResultOf; + result: ReturnOf; span: Span; startTime: number; }) => boolean; @@ -257,127 +251,6 @@ function shouldTraceEvent< } } -function ensureSpanStateForEvent< - TChannel extends AnyAsyncChannel | AnySyncStreamChannel, ->( - states: WeakMap, - config: ChannelConfig & { - extractInput: ( - args: [...ArgsOf, ...any[]], - event: StartOf, - span: Span, - ) => { - input: unknown; - metadata: unknown; - }; - }, - event: StartOf, - channelName: string, - instrumentationName: SpanInstrumentationName, -): SpanState | undefined { - const key = event as object; - const existing = states.get(key); - if (existing) { - return existing; - } - - if (!shouldTraceEvent(config, event, channelName)) { - return undefined; - } - - const created = startSpanForEvent( - config, - event, - channelName, - instrumentationName, - ); - states.set(key, created); - return created; -} - -function bindCurrentSpanStoreToStart< - TChannel extends AnyAsyncChannel | AnySyncStreamChannel, ->( - tracingChannel: IsoTracingChannel>, - states: WeakMap, - config: ChannelConfig & { - extractInput: ( - args: [...ArgsOf, ...any[]], - event: StartOf, - span: Span, - ) => { - input: unknown; - metadata: unknown; - }; - }, - channelName: string, - instrumentationName: SpanInstrumentationName, -): (() => void) | undefined { - const state = _internalGetGlobalState(); - const startChannel = tracingChannel.start; - const contextManager = state?.contextManager; - const currentSpanStore = contextManager - ? ( - contextManager as { - [BRAINTRUST_CURRENT_SPAN_STORE]?: CurrentSpanStore; - } - )[BRAINTRUST_CURRENT_SPAN_STORE] - : undefined; - - if (!currentSpanStore || !startChannel) { - return undefined; - } - - startChannel.bindStore( - currentSpanStore, - (event: ChannelMessage) => { - if (isAutoInstrumentationSuppressed()) { - return currentSpanStore.getStore(); - } - - const spanState = ensureSpanStateForEvent( - states, - config, - event as StartOf, - channelName, - instrumentationName, - ); - return spanState - ? contextManager!.wrapSpanForStore(spanState.span) - : currentSpanStore.getStore(); - }, - ); - - return () => { - startChannel.unbindStore(currentSpanStore); - }; -} - -function logErrorAndEnd< - TChannel extends AnyAsyncChannel | AnySyncStreamChannel, ->( - states: WeakMap, - event: ErrorOf, - channelName: string, -): void { - const spanData = states.get(event as object); - if (!spanData) { - return; - } - - try { - spanData.span.log({ error: event.error }); - } catch (error) { - debugLogger.error(`Error logging failure for ${channelName}:`, error); - } - try { - spanData.span.end(); - } catch (error) { - debugLogger.error(`Error ending span for ${channelName}:`, error); - } - states.delete(event as object); -} - function runStreamingCompletionHook(args: { channelName: string; config: StreamingChannelSpanConfig; @@ -439,432 +312,419 @@ function runStreamingErrorHook(args: { } } -export function traceAsyncChannel( - channel: TChannel, - config: AsyncChannelSpanConfig, -): () => void { - const tracingChannel = channel.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; - const states = new WeakMap(); - const channelName = channel.channelName; - const unbindCurrentSpanStore = bindCurrentSpanStoreToStart( - tracingChannel, - states, - config, - channelName, - channel.instrumentationName, +export function traceAsyncCall( + call: () => ReturnOf, + event: StartOf, + config: AsyncChannelSpanConfig & { + instrumentationName: SpanInstrumentationName; + }, +): ReturnOf { + const channelName = + typeof config.name === "string" ? config.name : "provider call"; + if ( + isAutoInstrumentationSuppressed() || + !shouldTraceEvent(config, event, channelName) + ) + return call(); + const spanData = runInstrumentation(() => + startSpanForEvent( + config, + event, + channelName, + config.instrumentationName, + ), ); - - const handlers: IsoChannelHandlers> = { - start: (event) => { - if (isAutoInstrumentationSuppressed()) { - return; - } - - ensureSpanStateForEvent( - states, - config, - event as StartOf, - channelName, - channel.instrumentationName, + if (!spanData) return call(); + const { span } = spanData; + const complete = (event: AsyncEndOf) => { + const asyncEndEvent = event as AsyncEndOf; + const { span, startTime } = spanData; + + try { + const output = config.extractOutput(asyncEndEvent.result, asyncEndEvent); + const metrics = config.extractMetrics( + asyncEndEvent.result, + startTime, + asyncEndEvent, + ); + const metadata = config.extractMetadata?.( + asyncEndEvent.result, + asyncEndEvent, ); - }, - asyncEnd: (event) => { - const spanData = states.get(event as object); - if (!spanData) { - return; - } - - const asyncEndEvent = event as AsyncEndOf; - const { span, startTime } = spanData; - - try { - const output = config.extractOutput( - asyncEndEvent.result, - asyncEndEvent, - ); - const metrics = config.extractMetrics( - asyncEndEvent.result, - startTime, - asyncEndEvent, - ); - const metadata = config.extractMetadata?.( - asyncEndEvent.result, - asyncEndEvent, - ); - span.log({ - output, - ...(normalizeMetadata(metadata) !== undefined - ? { metadata: normalizeMetadata(metadata) } - : {}), - metrics, - }); - } catch (error) { - debugLogger.error(`Error extracting output for ${channelName}:`, error); - } finally { - span.end(); - states.delete(event as object); - } - }, - error: (event) => { - logErrorAndEnd(states, event as ErrorOf, channelName); - }, + span.log({ + output, + ...(normalizeMetadata(metadata) !== undefined + ? { metadata: normalizeMetadata(metadata) } + : {}), + metrics, + }); + } catch (error) { + debugLogger.error(`Error extracting output for ${channelName}:`, error); + } finally { + span.end(); + } }; - - tracingChannel.subscribe(handlers); - - return () => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); + const fail = (error: unknown) => { + span.log({ error }); + span.end(); }; + let result: ReturnOf; + try { + result = withCurrent(span, call); + } catch (error) { + fail(error); + throw error; + } + return withCurrent(span, () => + observeResult( + result, + (value) => + complete( + Object.assign(event, { result: value }) as AsyncEndOf, + ), + fail, + ), + ); } -export function traceStreamingChannel( - channel: TChannel, - config: StreamingChannelSpanConfig, -): () => void { - const tracingChannel = channel.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; - const states = new WeakMap(); - const channelName = channel.channelName; - const unbindCurrentSpanStore = bindCurrentSpanStoreToStart( - tracingChannel, - states, - config, - channelName, - channel.instrumentationName, +export function traceStreamingCall( + call: () => ReturnOf, + event: StartOf, + config: StreamingChannelSpanConfig & { + instrumentationName: SpanInstrumentationName; + }, +): ReturnOf { + const channelName = + typeof config.name === "string" ? config.name : "provider call"; + if ( + isAutoInstrumentationSuppressed() || + !shouldTraceEvent(config, event, channelName) + ) + return call(); + const spanData = runInstrumentation(() => + startSpanForEvent( + config, + event, + channelName, + config.instrumentationName, + ), ); + if (!spanData) return call(); + const { span, startTime } = spanData; + const complete = (event: AsyncEndOf) => { + const asyncEndEvent = event as AsyncEndOf; + const { span, startTime } = spanData; + + if (isAsyncIterable(asyncEndEvent.result)) { + let firstChunkTime: number | undefined; + const handleStreamError = (error: Error) => { + try { + span.log({ error }); + } catch (loggingError) { + debugLogger.error( + `Error logging failure for ${channelName}:`, + loggingError, + ); + } + span.end(); - const handlers: IsoChannelHandlers> = { - start: (event) => { - if (isAutoInstrumentationSuppressed()) { - return; - } + runStreamingErrorHook({ + channelName, + config, + error, + event: asyncEndEvent, + span, + startTime, + }); + }; - ensureSpanStateForEvent( - states, - config, - event as StartOf, - channelName, - channel.instrumentationName, - ); - }, - asyncEnd: (event) => { - const spanData = states.get(event as object); - if (!spanData) { - return; - } + patchStreamIfNeeded(asyncEndEvent.result, { + onChunk: () => { + if (firstChunkTime === undefined) { + firstChunkTime = getCurrentUnixTimestamp(); + } + }, + onComplete: (chunks: ChunkOf[]) => { + let completion: + | { + metadata?: Record; + metrics: Record; + output: unknown; + } + | undefined; + try { + let output: unknown; + let metrics: Record; + let metadata: Record | undefined; + + if (config.aggregateChunks) { + const aggregated = config.aggregateChunks( + chunks, + asyncEndEvent.result, + asyncEndEvent, + startTime, + ); + output = aggregated.output; + metrics = aggregated.metrics; + metadata = aggregated.metadata; + } else { + output = config.extractOutput( + chunks as unknown as StreamingResult, + asyncEndEvent, + ); + metrics = config.extractMetrics( + chunks as unknown as StreamingResult, + startTime, + asyncEndEvent, + ); + } - const asyncEndEvent = event as AsyncEndOf; - const { span, startTime } = spanData; + if ( + metrics.time_to_first_token === undefined && + firstChunkTime !== undefined + ) { + metrics.time_to_first_token = firstChunkTime - startTime; + } else if ( + metrics.time_to_first_token === undefined && + chunks.length > 0 + ) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } - if (isAsyncIterable(asyncEndEvent.result)) { - let firstChunkTime: number | undefined; - const handleStreamError = (error: Error) => { - try { - span.log({ error }); - } catch (loggingError) { + completion = { + ...(metadata !== undefined ? { metadata } : {}), + metrics, + output, + }; + span.log({ + output, + ...(metadata !== undefined ? { metadata } : {}), + metrics, + }); + } catch (error) { debugLogger.error( - `Error logging failure for ${channelName}:`, - loggingError, + `Error extracting output for ${channelName}:`, + error, ); - } - try { + } finally { span.end(); - } catch (endingError) { - debugLogger.error( - `Error ending span for ${channelName}:`, - endingError, - ); } - states.delete(event as object); - runStreamingErrorHook({ - channelName, - config, - error, - event: asyncEndEvent, - span, - startTime, - }); - }; - - patchStreamIfNeeded(asyncEndEvent.result, { - onChunk: () => { - if (firstChunkTime === undefined) { - firstChunkTime = getCurrentUnixTimestamp(); - } - }, - onComplete: (chunks: ChunkOf[]) => { - let completion: - | { - metadata?: Record; - metrics: Record; - output: unknown; - } - | undefined; - try { - let output: unknown; - let metrics: Record; - let metadata: Record | undefined; - - if (config.aggregateChunks) { - const aggregated = config.aggregateChunks( - chunks, - asyncEndEvent.result, - asyncEndEvent, - startTime, - ); - output = aggregated.output; - metrics = aggregated.metrics; - metadata = aggregated.metadata; - } else { - output = config.extractOutput( - chunks as unknown as StreamingResult, - asyncEndEvent, - ); - metrics = config.extractMetrics( - chunks as unknown as StreamingResult, - startTime, - asyncEndEvent, - ); - } + if (completion) { + runStreamingCompletionHook({ + channelName, + chunks, + config, + endEvent: asyncEndEvent, + ...(completion.metadata !== undefined + ? { metadata: completion.metadata } + : {}), + metrics: completion.metrics, + output: completion.output, + result: asyncEndEvent.result as StreamingResult, + span, + startTime, + }); + } + }, + onCancel: () => { + const error = new Error("Stream cancelled before completion"); + error.name = "AbortError"; + handleStreamError(error); + }, + onError: handleStreamError, + }); + return; + } + + if ( + config.patchResult?.({ + channelName, + endEvent: asyncEndEvent, + result: asyncEndEvent.result as StreamingResult, + span, + startTime, + }) + ) { + return; + } + + let completion: + | { + metadata?: Record; + metrics: Record; + output: unknown; + } + | undefined; + try { + const output = config.extractOutput( + asyncEndEvent.result as StreamingResult, + asyncEndEvent, + ); + const metrics = config.extractMetrics( + asyncEndEvent.result as StreamingResult, + startTime, + asyncEndEvent, + ); + const metadata = config.extractMetadata?.( + asyncEndEvent.result as StreamingResult, + asyncEndEvent, + ); - if ( - metrics.time_to_first_token === undefined && - firstChunkTime !== undefined - ) { - metrics.time_to_first_token = firstChunkTime - startTime; - } else if ( - metrics.time_to_first_token === undefined && - chunks.length > 0 - ) { - metrics.time_to_first_token = - getCurrentUnixTimestamp() - startTime; - } + completion = { + ...(normalizeMetadata(metadata) !== undefined + ? { metadata: normalizeMetadata(metadata) } + : {}), + metrics, + output, + }; + span.log({ + output, + ...(normalizeMetadata(metadata) !== undefined + ? { metadata: normalizeMetadata(metadata) } + : {}), + metrics, + }); + } catch (error) { + debugLogger.error(`Error extracting output for ${channelName}:`, error); + } finally { + span.end(); + } + if (completion) { + runStreamingCompletionHook({ + channelName, + config, + endEvent: asyncEndEvent, + ...(completion.metadata !== undefined + ? { metadata: completion.metadata } + : {}), + metrics: completion.metrics, + output: completion.output, + result: asyncEndEvent.result as StreamingResult, + span, + startTime, + }); + } + }; + const fail = (error: unknown) => { + span.log({ error }); + span.end(); + runStreamingErrorHook({ + channelName, + config, + error: error as Error, + event: { ...event, error } as ErrorOf, + span, + startTime, + }); + }; + let result: ReturnOf; + try { + result = withCurrent(span, call); + } catch (error) { + fail(error); + throw error; + } + return withCurrent(span, () => + observeResult( + result, + (value) => + complete( + Object.assign(event, { result: value }) as AsyncEndOf, + ), + fail, + ), + ); +} - completion = { - ...(metadata !== undefined ? { metadata } : {}), - metrics, - output, - }; - span.log({ - output, - ...(metadata !== undefined ? { metadata } : {}), - metrics, - }); - } catch (error) { - debugLogger.error( - `Error extracting output for ${channelName}:`, - error, - ); - } finally { - try { - span.end(); - } catch (error) { - debugLogger.error( - `Error ending span for ${channelName}:`, - error, - ); - } - states.delete(event as object); - } - if (completion) { - runStreamingCompletionHook({ - channelName, - chunks, - config, - endEvent: asyncEndEvent, - ...(completion.metadata !== undefined - ? { metadata: completion.metadata } - : {}), - metrics: completion.metrics, - output: completion.output, - result: asyncEndEvent.result as StreamingResult, - span, - startTime, - }); - } - }, - onCancel: () => { - const error = new Error("Stream cancelled before completion"); - error.name = "AbortError"; - handleStreamError(error); - }, - onError: handleStreamError, - }); - return; - } +export function traceSyncStreamCall( + call: () => ReturnOf, + event: StartOf, + config: SyncStreamChannelSpanConfig & { + instrumentationName: SpanInstrumentationName; + }, +): ReturnOf { + const channelName = + typeof config.name === "string" ? config.name : "provider call"; + if ( + isAutoInstrumentationSuppressed() || + !shouldTraceEvent(config, event, channelName) + ) + return call(); + const spanData = runInstrumentation(() => + startSpanForEvent( + config, + event, + channelName, + config.instrumentationName, + ), + ); + if (!spanData) return call(); + const { span } = spanData; + const complete = (event: EndOf) => { + const { span, startTime } = spanData; + const endEvent = event as EndOf; + const handleResolvedResult = (result: ReturnOf) => { + const resolvedEndEvent = { + ...endEvent, + result, + } as EndOf; if ( config.patchResult?.({ channelName, - endEvent: asyncEndEvent, - result: asyncEndEvent.result as StreamingResult, + endEvent: resolvedEndEvent, + result, span, startTime, }) ) { - states.delete(event as object); return; } - let completion: - | { - metadata?: Record; - metrics: Record; - output: unknown; - } - | undefined; - try { - const output = config.extractOutput( - asyncEndEvent.result as StreamingResult, - asyncEndEvent, - ); - const metrics = config.extractMetrics( - asyncEndEvent.result as StreamingResult, - startTime, - asyncEndEvent, - ); - const metadata = config.extractMetadata?.( - asyncEndEvent.result as StreamingResult, - asyncEndEvent, - ); - - completion = { - ...(normalizeMetadata(metadata) !== undefined - ? { metadata: normalizeMetadata(metadata) } - : {}), - metrics, - output, - }; - span.log({ - output, - ...(normalizeMetadata(metadata) !== undefined - ? { metadata: normalizeMetadata(metadata) } - : {}), - metrics, - }); - } catch (error) { - debugLogger.error(`Error extracting output for ${channelName}:`, error); - } finally { - try { - span.end(); - } catch (error) { - debugLogger.error(`Error ending span for ${channelName}:`, error); - } - states.delete(event as object); - } - if (completion) { - runStreamingCompletionHook({ - channelName, - config, - endEvent: asyncEndEvent, - ...(completion.metadata !== undefined - ? { metadata: completion.metadata } - : {}), - metrics: completion.metrics, - output: completion.output, - result: asyncEndEvent.result as StreamingResult, - span, - startTime, - }); - } - }, - error: (event) => { - const spanData = states.get(event as object); - logErrorAndEnd(states, event as ErrorOf, channelName); - if (spanData) { - runStreamingErrorHook({ - channelName, - config, - error: (event as ErrorOf).error, - event: event as ErrorOf, - span: spanData.span, - startTime: spanData.startTime, - }); - } - }, - }; - - tracingChannel.subscribe(handlers); + const stream = result; - return () => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }; -} - -export function traceSyncStreamChannel( - channel: TChannel, - config: SyncStreamChannelSpanConfig, -): () => void { - const tracingChannel = channel.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; - const states = new WeakMap(); - const channelName = channel.channelName; - const unbindCurrentSpanStore = bindCurrentSpanStoreToStart( - tracingChannel, - states, - config, - channelName, - channel.instrumentationName, - ); + if (!isSyncStreamLike>(stream)) { + span.end(); - const handlers: IsoChannelHandlers> = { - start: (event) => { - if (isAutoInstrumentationSuppressed()) { return; } - ensureSpanStateForEvent( - states, - config, - event as StartOf, - channelName, - channel.instrumentationName, - ); - }, - end: (event) => { - const spanData = states.get(event as object); - if (!spanData) { - return; - } + let first = true; - const { span, startTime } = spanData; - const endEvent = event as EndOf; - const handleResolvedResult = (result: ResultOf) => { - const resolvedEndEvent = { - ...endEvent, - result, - } as EndOf; - - if ( - config.patchResult?.({ - channelName, - endEvent: resolvedEndEvent, - result, - span, - startTime, - }) - ) { - return; + stream.on("chunk", () => { + if (first) { + span.log({ + metrics: { + time_to_first_token: getCurrentUnixTimestamp() - startTime, + }, + }); + first = false; } + }); - const stream = result; + stream.on("chatCompletion", (completion) => { + try { + if (hasChoices(completion)) { + span.log({ + output: completion.choices, + }); + } + } catch (error) { + debugLogger.error( + `Error extracting chatCompletion for ${channelName}:`, + error, + ); + } + }); - if (!isSyncStreamLike>(stream)) { - span.end(); - states.delete(event as object); + stream.on("event", (streamEvent) => { + if (!config.extractFromEvent) { return; } - let first = true; - - stream.on("chunk", () => { + try { if (first) { span.log({ metrics: { @@ -873,79 +733,49 @@ export function traceSyncStreamChannel( }); first = false; } - }); - - stream.on("chatCompletion", (completion) => { - try { - if (hasChoices(completion)) { - span.log({ - output: completion.choices, - }); - } - } catch (error) { - debugLogger.error( - `Error extracting chatCompletion for ${channelName}:`, - error, - ); - } - }); - - stream.on("event", (streamEvent) => { - if (!config.extractFromEvent) { - return; - } - - try { - if (first) { - span.log({ - metrics: { - time_to_first_token: getCurrentUnixTimestamp() - startTime, - }, - }); - first = false; - } - const extracted = config.extractFromEvent(streamEvent); - if (extracted && Object.keys(extracted).length > 0) { - span.log(extracted); - } - } catch (error) { - debugLogger.error( - `Error extracting event for ${channelName}:`, - error, - ); + const extracted = config.extractFromEvent(streamEvent); + if (extracted && Object.keys(extracted).length > 0) { + span.log(extracted); } - }); + } catch (error) { + debugLogger.error( + `Error extracting event for ${channelName}:`, + error, + ); + } + }); - stream.on("end", () => { - span.end(); - states.delete(event as object); - }); + stream.on("end", () => { + span.end(); + }); - stream.on("error", (error: Error) => { - span.log({ - error: error.message, - }); - span.end(); - states.delete(event as object); + stream.on("error", (error: Error) => { + span.log({ + error: error.message, }); - }; + span.end(); + }); + }; - handleResolvedResult(endEvent.result); - }, - error: (event) => { - logErrorAndEnd(states, event as ErrorOf, channelName); - }, + handleResolvedResult(endEvent.result); }; - - tracingChannel.subscribe(handlers); - - return () => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); + const fail = (error: unknown) => { + span.log({ error }); + span.end(); }; + let result: ReturnOf; + try { + result = withCurrent(span, call); + } catch (error) { + fail(error); + throw error; + } + withCurrent(span, () => + runInstrumentation(() => complete({ ...event, result } as EndOf)), + ); + return result; } - export function unsubscribeAll( unsubscribers: Array<() => void>, ): Array<() => void> { diff --git a/js/src/instrumentation/core/index.ts b/js/src/instrumentation/core/index.ts index 6e2d2d61c..06c91c9b1 100644 --- a/js/src/instrumentation/core/index.ts +++ b/js/src/instrumentation/core/index.ts @@ -9,19 +9,18 @@ * bundler subpaths, such as `braintrust/vite`. */ -export { BasePlugin } from "./plugin"; -export { toLoggedError } from "./logging"; export { createChannelName, - parseChannelName, isValidChannelName, + parseChannelName, } from "./channel"; +export { toLoggedError } from "./logging"; +export { BasePlugin } from "./plugin"; export type { + AsyncEndEvent, + AsyncStartEvent, BaseContext, - StartEvent, EndEvent, ErrorEvent, - AsyncStartEvent, - AsyncEndEvent, - ChannelHandlers, + StartEvent, } from "./types"; diff --git a/js/src/instrumentation/core/observe-result.test.ts b/js/src/instrumentation/core/observe-result.test.ts new file mode 100644 index 000000000..b04ee7e60 --- /dev/null +++ b/js/src/instrumentation/core/observe-result.test.ts @@ -0,0 +1,58 @@ +import { describe, expect, it, vi } from "vitest"; +import { observeResult } from "./observe-result"; + +vi.mock("../../debug-logger", () => ({ debugLogger: { error: vi.fn() } })); + +describe("observeResult", () => { + it("preserves Promise subclasses and their helper methods", async () => { + class ProviderPromise extends Promise { + requestId = "request-1"; + } + const promise = new ProviderPromise((resolve) => resolve("response")); + const fulfilled = vi.fn(); + expect(observeResult(promise, fulfilled, vi.fn())).toBe(promise); + await promise; + expect(promise.requestId).toBe("request-1"); + expect(fulfilled).toHaveBeenCalledExactlyOnceWith("response"); + }); + + it("does not replace a provider rejection when an observer throws", async () => { + const error = new Error("provider failed"); + const promise = Promise.reject(error); + expect( + observeResult(promise, vi.fn(), () => { + throw new Error("observer failed"); + }), + ).toBe(promise); + await expect(promise).rejects.toBe(error); + }); + + it("contains asynchronous observer failures", async () => { + const promise = Promise.resolve("response"); + observeResult( + promise, + async () => { + throw new Error("observer failed"); + }, + vi.fn(), + ); + await expect(promise).resolves.toBe("response"); + await Promise.resolve(); + }); + + it("contains a throwing then getter without altering the return value", () => { + const result = { + get then(): never { + throw new Error("not a promise"); + }, + }; + expect(observeResult(result, vi.fn(), vi.fn())).toBe(result); + }); + + it("observes synchronous results without making them asynchronous", () => { + const result = { content: "response" }; + const fulfilled = vi.fn(); + expect(observeResult(result, fulfilled, vi.fn())).toBe(result); + expect(fulfilled).toHaveBeenCalledExactlyOnceWith(result); + }); +}); diff --git a/js/src/instrumentation/core/observe-result.ts b/js/src/instrumentation/core/observe-result.ts new file mode 100644 index 000000000..32636dbd8 --- /dev/null +++ b/js/src/instrumentation/core/observe-result.ts @@ -0,0 +1,44 @@ +import { debugLogger } from "../../debug-logger"; + +/** Contain observation failures without retrying or replacing the provider call. */ +export function runInstrumentation(observe: () => T): T | undefined { + try { + const value = observe(); + if (value instanceof Promise) + void value.catch((error) => { + debugLogger.error("Error observing provider call:", error); + }); + return value; + } catch (error) { + debugLogger.error("Error observing provider call:", error); + return undefined; + } +} + +/** Observe settlement while preserving the original value and Promise helpers. */ +export function observeResult( + result: T, + fulfilled: (value: Awaited) => void, + rejected: (error: unknown) => void, +): T { + runInstrumentation(() => { + const then = + result != null && + (typeof result === "object" || typeof result === "function") + ? Reflect.get(result, "then") + : undefined; + if (typeof then !== "function") { + fulfilled(result as Awaited); + return; + } + // Call then directly: Promise.resolve would defer observing foreign thenables. + const observed = Reflect.apply(then, result, [ + (value: Awaited) => runInstrumentation(() => fulfilled(value)), + (error: unknown) => runInstrumentation(() => rejected(error)), + ]); + if (observed != null && typeof observed.then === "function") { + observed.then(undefined, () => {}); + } + }); + return result; +} diff --git a/js/src/instrumentation/core/plugin.ts b/js/src/instrumentation/core/plugin.ts index ae09a43f1..9aacf97fa 100644 --- a/js/src/instrumentation/core/plugin.ts +++ b/js/src/instrumentation/core/plugin.ts @@ -1,27 +1,9 @@ -import iso from "../../isomorph"; -import type { IsoChannelHandlers } from "../../isomorph"; -import { isAsyncIterable, patchStreamIfNeeded } from "./stream-patcher"; -import type { StartEvent } from "./types"; -import { startSpan } from "../../logger"; -import type { Span } from "../../logger"; -import { getCurrentUnixTimestamp } from "../../util"; -import { - buildStartSpanArgs, - mergeInputMetadata, -} from "./channel-tracing-utils"; - -/** - * Base class for creating instrumentation plugins. - * - * Plugins subscribe to global instrumentation hook events and convert them - * into spans, logs, or other observability data. - */ export abstract class BasePlugin { protected enabled = false; protected unsubscribers: Array<() => void> = []; /** - * Enables the plugin. Must be called before the plugin will receive events. + * Enables the plugin. Registers the plugin’s invocation interceptors. */ enable(): void { if (this.enabled) { @@ -32,7 +14,7 @@ export abstract class BasePlugin { } /** - * Disables the plugin. After this, the plugin will no longer receive events. + * Disables the plugin. Removes the plugin’s invocation interceptors. */ disable(): void { if (!this.enabled) { @@ -44,462 +26,13 @@ export abstract class BasePlugin { /** * Called when the plugin is enabled. - * Override this to set up subscriptions. + * Override this to register interceptors. */ protected abstract onEnable(): void; /** * Called when the plugin is disabled. - * Override this to clean up subscriptions. + * Override this to remove interceptors. */ protected abstract onDisable(): void; - - /** - * Helper to subscribe to a channel with raw handlers. - * - * @param channelName - The channel name to subscribe to - * @param handlers - Event handlers - */ - protected subscribe(channelName: string, handlers: IsoChannelHandlers): void { - const channel = iso.newTracingChannel(channelName); - channel.subscribe(handlers); - } - - /** - * Subscribe to a channel for async methods (non-streaming). - * Creates a span and logs input/output/metrics. - */ - protected subscribeToChannel( - channelName: string, - config: { - name: string; - type: string; - extractInput: (args: any[]) => { input: any; metadata: any }; - extractOutput: (result: any, endEvent?: any) => any; - extractMetadata?: (result: any, endEvent?: any) => any; - extractMetrics: ( - result: any, - startTime?: number, - endEvent?: any, - ) => Record; - }, - ): void { - const channel = iso.newTracingChannel(channelName); - - const spans = new WeakMap(); - - const handlers = { - start: (event: StartEvent) => { - const { name, spanAttributes, spanInfoMetadata } = buildStartSpanArgs( - config, - event, - ); - const span = startSpan({ - name, - spanAttributes, - }); - - const startTime = getCurrentUnixTimestamp(); - spans.set(event, { span, startTime }); - - try { - const { input, metadata } = config.extractInput(event.arguments); - span.log({ - input, - metadata: mergeInputMetadata(metadata, spanInfoMetadata), - }); - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error(`Error extracting input for ${channelName}:`, error); - } - }, - - asyncEnd: (event: any) => { - const spanData = spans.get(event); - if (!spanData) { - return; - } - - const { span, startTime } = spanData; - - try { - const output = config.extractOutput(event.result, event); - const metrics = config.extractMetrics(event.result, startTime, event); - const metadata = config.extractMetadata?.(event.result, event); - - span.log({ - output, - ...(metadata !== undefined ? { metadata } : {}), - metrics, - }); - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error(`Error extracting output for ${channelName}:`, error); - } finally { - span.end(); - spans.delete(event); - } - }, - - error: (event: any) => { - const spanData = spans.get(event); - if (!spanData) { - return; - } - - const { span } = spanData; - - span.log({ - error: event.error.message, - }); - span.end(); - spans.delete(event); - }, - }; - - channel.subscribe(handlers); - - // Store unsubscribe function - this.unsubscribers.push(() => { - channel.unsubscribe(handlers); - }); - } - - /** - * Subscribe to a channel for async methods that may return streams. - * Handles both streaming and non-streaming responses. - */ - protected subscribeToStreamingChannel( - channelName: string, - config: { - name: string; - type: string; - extractInput: (args: any[]) => { input: any; metadata: any }; - extractOutput: (result: any, endEvent?: any) => any; - extractMetadata?: (result: any, endEvent?: any) => any; - extractMetrics: ( - result: any, - startTime?: number, - endEvent?: any, - ) => Record; - aggregateChunks?: ( - chunks: any[], - result?: any, - endEvent?: any, - ) => { - output: any; - metrics: Record; - metadata?: any; - }; - }, - ): void { - const channel = iso.newTracingChannel(channelName); - - const spans = new WeakMap(); - - const handlers = { - start: (event: StartEvent) => { - const { name, spanAttributes, spanInfoMetadata } = buildStartSpanArgs( - config, - event, - ); - const span = startSpan({ - name, - spanAttributes, - }); - - const startTime = getCurrentUnixTimestamp(); - spans.set(event, { span, startTime }); - - try { - const { input, metadata } = config.extractInput(event.arguments); - span.log({ - input, - metadata: mergeInputMetadata(metadata, spanInfoMetadata), - }); - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error(`Error extracting input for ${channelName}:`, error); - } - }, - - asyncEnd: (event: any) => { - const spanData = spans.get(event); - if (!spanData) { - return; - } - - const { span, startTime } = spanData; - - // Check if result is a stream - if (isAsyncIterable(event.result)) { - let firstChunkTime: number | undefined; - - // Patch the stream to collect chunks - patchStreamIfNeeded(event.result, { - onChunk: () => { - if (firstChunkTime === undefined) { - firstChunkTime = getCurrentUnixTimestamp(); - } - }, - onComplete: (chunks: any[]) => { - try { - let output: any; - let metrics: Record; - let metadata: any; - - if (config.aggregateChunks) { - const aggregated = config.aggregateChunks( - chunks, - event.result, - event, - ); - output = aggregated.output; - metrics = aggregated.metrics; - metadata = aggregated.metadata; - } else { - output = config.extractOutput(chunks, event); - metrics = config.extractMetrics(chunks, startTime, event); - } - - // Add time_to_first_token if not already present - if ( - metrics.time_to_first_token === undefined && - firstChunkTime !== undefined - ) { - metrics.time_to_first_token = firstChunkTime - startTime; - } else if ( - metrics.time_to_first_token === undefined && - chunks.length > 0 - ) { - metrics.time_to_first_token = - getCurrentUnixTimestamp() - startTime; - } - - span.log({ - output, - ...(metadata !== undefined ? { metadata } : {}), - metrics, - }); - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error( - `Error extracting output for ${channelName}:`, - error, - ); - } finally { - span.end(); - } - }, - onError: (error: Error) => { - span.log({ - error: error.message, - }); - span.end(); - }, - }); - - // Don't delete the span from the map yet - it will be ended by the stream - } else { - // Non-streaming response - try { - const output = config.extractOutput(event.result, event); - const metadata = config.extractMetadata - ? config.extractMetadata(event.result, event) - : undefined; - const metrics = config.extractMetrics( - event.result, - startTime, - event, - ); - - span.log({ - output, - ...(metadata !== undefined ? { metadata } : {}), - metrics, - }); - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error(`Error extracting output for ${channelName}:`, error); - } finally { - span.end(); - spans.delete(event); - } - } - }, - - error: (event: any) => { - const spanData = spans.get(event); - if (!spanData) { - return; - } - - const { span } = spanData; - - span.log({ - error: event.error.message, - }); - span.end(); - spans.delete(event); - }, - }; - - channel.subscribe(handlers); - - // Store unsubscribe function - this.unsubscribers.push(() => { - channel.unsubscribe(handlers); - }); - } - - /** - * Subscribe to a channel for sync methods that return event-based streams. - * Used for methods like beta.chat.completions.stream() and responses.stream(). - */ - protected subscribeToSyncStreamChannel( - channelName: string, - config: { - name: string; - type: string; - extractInput: (args: any[]) => { input: any; metadata: any }; - extractFromEvent?: (event: any) => { - output?: any; - metrics?: Record; - metadata?: any; - }; - }, - ): void { - const channel = iso.newTracingChannel(channelName); - - const spans = new WeakMap(); - - const handlers = { - start: (event: StartEvent) => { - const { name, spanAttributes, spanInfoMetadata } = buildStartSpanArgs( - config, - event, - ); - const span = startSpan({ - name, - spanAttributes, - }); - - const startTime = getCurrentUnixTimestamp(); - spans.set(event, { span, startTime }); - - try { - const { input, metadata } = config.extractInput(event.arguments); - span.log({ - input, - metadata: mergeInputMetadata(metadata, spanInfoMetadata), - }); - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error(`Error extracting input for ${channelName}:`, error); - } - }, - - end: (event: any) => { - const spanData = spans.get(event); - if (!spanData) { - return; - } - - const { span, startTime } = spanData; - const stream = event.result; - - if (!stream || typeof stream.on !== "function") { - // Not a stream, just end the span - span.end(); - spans.delete(event); - return; - } - - let first = true; - - // Listen for stream events - stream.on("chunk", (chunk: any) => { - if (first) { - const now = getCurrentUnixTimestamp(); - span.log({ - metrics: { - time_to_first_token: now - startTime, - }, - }); - first = false; - } - }); - - stream.on("chatCompletion", (completion: any) => { - try { - span.log({ - output: completion.choices, - }); - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error( - `Error extracting chatCompletion for ${channelName}:`, - error, - ); - } - }); - - stream.on("event", (streamEvent: any) => { - if (config.extractFromEvent) { - try { - if (first) { - const now = getCurrentUnixTimestamp(); - span.log({ - metrics: { - time_to_first_token: now - startTime, - }, - }); - first = false; - } - - const extracted = config.extractFromEvent(streamEvent); - if (extracted && Object.keys(extracted).length > 0) { - span.log(extracted); - } - } catch (error) { - // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. - console.error( - `Error extracting event for ${channelName}:`, - error, - ); - } - } - }); - - stream.on("end", () => { - span.end(); - spans.delete(event); - }); - - // Don't delete the span from the map - it will be deleted when the stream ends - }, - - error: (event: any) => { - const spanData = spans.get(event); - if (!spanData) { - return; - } - - const { span } = spanData; - - span.log({ - error: event.error.message, - }); - span.end(); - spans.delete(event); - }, - }; - - channel.subscribe(handlers); - - // Store unsubscribe function - this.unsubscribers.push(() => { - channel.unsubscribe(handlers); - }); - } } diff --git a/js/src/instrumentation/core/tracing-types.ts b/js/src/instrumentation/core/tracing-types.ts new file mode 100644 index 000000000..24bb9e4a5 --- /dev/null +++ b/js/src/instrumentation/core/tracing-types.ts @@ -0,0 +1,39 @@ +import type { ArgsOf, ExtraOf, ReturnOf } from "./channel-definitions"; +import type { + AsyncEndEventWith, + EndEventWith, + ErrorEventWith, + StartEventWith, +} from "./types"; + +// These types describe tracing data, never a registration or invocation API. +export type AnyAsyncChannel = { + __args?: readonly unknown[]; + __result?: unknown; + __extra?: object; + __chunk?: unknown; + channelName: string; +}; +export type AnySyncStreamChannel = AnyAsyncChannel; +export type ResultOf = Awaited>; +export type StartOf = StartEventWith< + ArgsOf, + ExtraOf +>; +export type AsyncEndOf = AsyncEndEventWith< + ResultOf, + ArgsOf, + ExtraOf +>; +export type EndOf = EndEventWith< + ReturnOf, + ArgsOf, + ExtraOf +>; +export type ErrorOf = ErrorEventWith< + ArgsOf, + ExtraOf +>; +export type ChannelMessage = StartOf & + Partial<{ result: ResultOf }> & + Partial, "error">>; diff --git a/js/src/instrumentation/core/types.ts b/js/src/instrumentation/core/types.ts index dc2434dc3..f6b2ab594 100644 --- a/js/src/instrumentation/core/types.ts +++ b/js/src/instrumentation/core/types.ts @@ -1,14 +1,4 @@ -/** - * Standard event types for global hook-based instrumentation. - * - * These types retain the TracingChannel-compatible lifecycle. - * For async functions (tracePromise): - * - start: Called before the synchronous portion executes - * - end: Called after the synchronous portion completes (promise returned) - * - asyncStart: Called when the promise begins to settle - * - asyncEnd: Called when the promise finishes settling (before user code continues) - * - error: Called if the function throws or the promise rejects - */ +/** Per-call tracing data. These types do not define invocation hooks or event subscriptions. */ export type EventArguments = readonly unknown[]; @@ -25,25 +15,22 @@ export type SpanInfoCarrier< }; /** - * Base context object shared across all events in a trace. + * Per-call context passed to tracing callbacks. */ export interface BaseContext { /** * Unique identifier for this trace. - * Can be used to correlate start/end/error events. */ traceId?: string; /** - * Arbitrary data that can be attached by event handlers - * and passed between start/end/error events. + * Additional data used by tracing callbacks. */ [key: string]: unknown; } /** - * Event emitted before the synchronous portion of a function executes. - * This is where you should create spans and extract input data. + * Input context for a tracing function. */ export interface StartEvent extends BaseContext { /** @@ -53,8 +40,7 @@ export interface StartEvent extends BaseContext { } /** - * Event emitted after the synchronous portion completes. - * For async functions, this fires when the promise is returned (not settled). + * Context containing a returned value. */ export interface EndEvent extends BaseContext { /** @@ -70,7 +56,7 @@ export interface EndEvent extends BaseContext { } /** - * Event emitted when a function throws or a promise rejects. + * Context containing a thrown error or rejected promise. */ export interface ErrorEvent extends BaseContext { /** @@ -85,8 +71,7 @@ export interface ErrorEvent extends BaseContext { } /** - * Event emitted when a promise begins to settle. - * This fires after the synchronous portion and when the async continuation starts. + * Input context preserving the callable argument tuple. */ export interface TypedStartEvent< TArguments extends EventArguments = unknown[], @@ -116,9 +101,7 @@ export interface TypedErrorEvent< export interface AsyncStartEvent extends StartEvent {} /** - * Event emitted when a promise finishes settling. - * This fires BEFORE control returns to user code after await. - * This is where you should extract output data and finalize spans. + * Context containing a resolved value for output extraction and finalization. */ // eslint-disable-next-line @typescript-eslint/no-empty-object-type export interface AsyncEndEvent extends EndEvent {} @@ -144,43 +127,3 @@ export type ErrorEventWith< TArguments extends EventArguments = unknown[], TExtra extends object = Record, > = TypedErrorEvent & TExtra; - -/** - * Subscription handlers for a tracing-compatible global hook. - * - * Common usage pattern: - * - Use start to create spans and extract input - * - Use asyncEnd to extract output and finalize spans - * - Use error to handle failures - */ -export interface ChannelHandlers { - /** - * Called before the synchronous portion of a function executes. - * Use this to create spans and extract input data. - */ - start?: (event: StartEvent) => void; - - /** - * Called after the synchronous portion completes (promise returned). - * Usually not needed for typical instrumentation. - */ - end?: (event: EndEvent) => void; - - /** - * Called when a promise begins to settle. - * Usually not needed for typical instrumentation. - */ - asyncStart?: (event: AsyncStartEvent) => void; - - /** - * Called when a promise finishes settling, before user code continues. - * Use this to extract output, patch streams, and finalize spans. - */ - asyncEnd?: (event: AsyncEndEvent) => void; - - /** - * Called when a function throws or promise rejects. - * Use this to log errors and clean up spans. - */ - error?: (event: ErrorEvent) => void; -} diff --git a/js/src/instrumentation/index.ts b/js/src/instrumentation/index.ts index 833786cc5..cc7f85e8b 100644 --- a/js/src/instrumentation/index.ts +++ b/js/src/instrumentation/index.ts @@ -1,8 +1,8 @@ /** * Instrumentation APIs for auto-instrumentation. * - * This module provides the core plugin infrastructure for converting global - * instrumentation hook events into Braintrust spans. + * This module provides the core plugin infrastructure for registering invocation + * interceptors that trace provider calls. * * Following the OpenTelemetry pattern, BasePlugin (like InstrumentationBase) * lives in the core SDK, while individual instrumentation implementations @@ -14,35 +14,34 @@ * @module instrumentation */ -export { BasePlugin } from "./core"; export { BraintrustPlugin } from "./braintrust-plugin"; export type { BraintrustPluginConfig } from "./braintrust-plugin"; -export { OpenAIAgentsTraceProcessor } from "./plugins/openai-agents-trace-processor"; -export type { OpenAIAgentsTraceProcessorOptions } from "./plugins/openai-agents-trace-processor"; +export { BasePlugin } from "./core"; +export { braintrustEveInstrumentation } from "./plugins/eve-instrumentation"; +export { braintrustEveHook } from "./plugins/eve-plugin"; export { braintrustFlueInstrumentation, braintrustFlueObserver, } from "./plugins/flue-plugin"; -export { braintrustEveHook } from "./plugins/eve-plugin"; -export { braintrustEveInstrumentation } from "./plugins/eve-instrumentation"; +export { OpenAIAgentsTraceProcessor } from "./plugins/openai-agents-trace-processor"; +export type { OpenAIAgentsTraceProcessorOptions } from "./plugins/openai-agents-trace-processor"; // Re-export core types for external instrumentation packages +export { + createChannelName, + isValidChannelName, + parseChannelName, +} from "./core"; export type { + AsyncEndEvent, + AsyncStartEvent, BaseContext, - StartEvent, EndEvent, ErrorEvent, - AsyncStartEvent, - AsyncEndEvent, - ChannelHandlers, -} from "./core"; -export { - createChannelName, - parseChannelName, - isValidChannelName, + StartEvent, } from "./core"; // Configuration API +export type { SpanCustomizer, SpanExportData } from "./config"; export { configureInstrumentation } from "./registry"; export type { InstrumentationConfig } from "./registry"; -export type { SpanCustomizer, SpanExportData } from "./config"; diff --git a/js/src/instrumentation/plugins/ai-sdk-channels.ts b/js/src/instrumentation/plugins/ai-sdk-channels.ts index 0842105a0..0531b4e0e 100644 --- a/js/src/instrumentation/plugins/ai-sdk-channels.ts +++ b/js/src/instrumentation/plugins/ai-sdk-channels.ts @@ -1,13 +1,12 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; -import type { ChannelSpanInfo } from "../core/types"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { AISDK, - AISDKEvaluateParams, - AISDKEvaluationResult, AISDKCallParams, AISDKEmbedParams, AISDKEmbeddingResult, + AISDKEvaluateParams, + AISDKEvaluationResult, AISDKGenerateImageParams, AISDKHarnessAgentCallParams, AISDKHarnessAgentCreateSessionParams, @@ -22,6 +21,7 @@ import type { AISDKV7CreateTelemetryDispatcherArgs, AISDKV7TelemetryDispatcher, } from "../../vendor-sdk-types/ai-sdk-v7-telemetry"; +import type { ChannelSpanInfo } from "../core/types"; type AISDKStreamResult = AISDKResult | AsyncIterable; type AISDKChannelContext = { @@ -40,231 +40,200 @@ export const BRAINTRUST_WRAPPED_AI_SDK_MODEL = Symbol.for( "braintrust.ai-sdk.wrapped-model", ); -export const aiSDKChannels = defineChannels( - "ai", - { - modelGenerate: channel< - [AISDKCallParams], - AISDKResult, - AISDKModelChannelContext - >({ - channelName: "model.doGenerate", - kind: "async", - }), - modelStream: channel< - [AISDKCallParams], - AISDKResult & { stream: ReadableStream }, - AISDKModelChannelContext - >({ - channelName: "model.doStream", - kind: "async", - }), - evaluate: channel< - [AISDKEvaluateParams], - AISDKEvaluationResult, - AISDKChannelContext - >({ - channelName: "evaluate", - kind: "async", - }), - generateText: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "generateText", - kind: "async", - }), - generateImage: channel< - [AISDKGenerateImageParams], - AISDKResult, - AISDKChannelContext - >({ - channelName: "generateImage", - kind: "async", - }), - streamText: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "streamText", - kind: "async", - }), - streamTextSync: channel< - [AISDKCallParams], - AISDKResult, - AISDKChannelContext, - unknown - >({ - channelName: "streamText.sync", - kind: "sync-stream", - }), - generateObject: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "generateObject", - kind: "async", - }), - streamObject: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "streamObject", - kind: "async", - }), - streamObjectSync: channel< - [AISDKCallParams], - AISDKResult, - AISDKChannelContext, - unknown - >({ - channelName: "streamObject.sync", - kind: "sync-stream", - }), - embed: channel< - [AISDKEmbedParams], - AISDKEmbeddingResult, - AISDKChannelContext - >({ - channelName: "embed", - kind: "async", - }), - embedMany: channel< - [AISDKEmbedParams], - AISDKEmbeddingResult, - AISDKChannelContext - >({ - channelName: "embedMany", - kind: "async", - }), - rerank: channel< - [AISDKRerankParams], - AISDKRerankResult, - AISDKChannelContext - >({ - channelName: "rerank", - kind: "async", - }), - agentGenerate: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "Agent.generate", - kind: "async", - }), - agentStream: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "Agent.stream", - kind: "async", - }), - agentStreamSync: channel< - [AISDKCallParams], - AISDKResult, - AISDKChannelContext, - unknown - >({ - channelName: "Agent.stream.sync", - kind: "sync-stream", - }), - toolLoopAgentGenerate: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "ToolLoopAgent.generate", - kind: "async", - }), - toolLoopAgentStream: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "ToolLoopAgent.stream", - kind: "async", - }), - workflowAgentStream: channel< - [AISDKCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "WorkflowAgent.stream", - kind: "async", - }), - v7CreateTelemetryDispatcher: channel< - [AISDKV7CreateTelemetryDispatcherArgs], - AISDKV7TelemetryDispatcher - >({ - channelName: "createTelemetryDispatcher", - kind: "sync-stream", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.AI_SDK }, -); +export const aiSDKChannels = defineInterceptor("ai", { + modelGenerate: channel< + [AISDKCallParams], + PromiseLike, + AISDKModelChannelContext + >({ + channelName: "model.doGenerate", + }), + modelStream: channel< + [AISDKCallParams], + PromiseLike< + AISDKResult & { stream: ReadableStream } + >, + AISDKModelChannelContext + >({ + channelName: "model.doStream", + }), + evaluate: channel< + [AISDKEvaluateParams], + PromiseLike, + AISDKChannelContext + >({ + channelName: "evaluate", + }), + generateText: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "generateText", + }), + generateImage: channel< + [AISDKGenerateImageParams], + PromiseLike, + AISDKChannelContext + >({ + channelName: "generateImage", + }), + streamText: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "streamText", + }), + streamTextSync: channel< + [AISDKCallParams], + AISDKResult, + AISDKChannelContext, + unknown + >({ + channelName: "streamText.sync", + }), + generateObject: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "generateObject", + }), + streamObject: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "streamObject", + }), + streamObjectSync: channel< + [AISDKCallParams], + AISDKResult, + AISDKChannelContext, + unknown + >({ + channelName: "streamObject.sync", + }), + embed: channel< + [AISDKEmbedParams], + PromiseLike, + AISDKChannelContext + >({ + channelName: "embed", + }), + embedMany: channel< + [AISDKEmbedParams], + PromiseLike, + AISDKChannelContext + >({ + channelName: "embedMany", + }), + rerank: channel< + [AISDKRerankParams], + PromiseLike, + AISDKChannelContext + >({ + channelName: "rerank", + }), + agentGenerate: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "Agent.generate", + }), + agentStream: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "Agent.stream", + }), + agentStreamSync: channel< + [AISDKCallParams], + AISDKResult, + AISDKChannelContext, + unknown + >({ + channelName: "Agent.stream.sync", + }), + toolLoopAgentGenerate: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "ToolLoopAgent.generate", + }), + toolLoopAgentStream: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "ToolLoopAgent.stream", + }), + workflowAgentStream: channel< + [AISDKCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "WorkflowAgent.stream", + }), + v7CreateTelemetryDispatcher: channel< + [AISDKV7CreateTelemetryDispatcherArgs], + AISDKV7TelemetryDispatcher + >({ + channelName: "createTelemetryDispatcher", + }), +}); -export const harnessAgentChannels = defineChannels( - "@ai-sdk/harness", - { - createSession: channel< - [AISDKHarnessAgentCreateSessionParams?], - AISDKHarnessAgentSession, - AISDKChannelContext - >({ - channelName: "HarnessAgent.createSession", - kind: "async", - }), - generate: channel< - [AISDKHarnessAgentCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "HarnessAgent.generate", - kind: "async", - }), - stream: channel< - [AISDKHarnessAgentCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "HarnessAgent.stream", - kind: "async", - }), - continueGenerate: channel< - [AISDKHarnessAgentCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "HarnessAgent.continueGenerate", - kind: "async", - }), - continueStream: channel< - [AISDKHarnessAgentCallParams], - AISDKStreamResult, - AISDKChannelContext, - unknown - >({ - channelName: "HarnessAgent.continueStream", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.AI_SDK }, -); +export const harnessAgentChannels = defineInterceptor("@ai-sdk/harness", { + createSession: channel< + [AISDKHarnessAgentCreateSessionParams?], + PromiseLike, + AISDKChannelContext + >({ + channelName: "HarnessAgent.createSession", + }), + generate: channel< + [AISDKHarnessAgentCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "HarnessAgent.generate", + }), + stream: channel< + [AISDKHarnessAgentCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "HarnessAgent.stream", + }), + continueGenerate: channel< + [AISDKHarnessAgentCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "HarnessAgent.continueGenerate", + }), + continueStream: channel< + [AISDKHarnessAgentCallParams], + PromiseLike | AISDKStreamResult, + AISDKChannelContext, + unknown + >({ + channelName: "HarnessAgent.continueStream", + }), +}); diff --git a/js/src/instrumentation/plugins/ai-sdk-plugin.streaming.test.ts b/js/src/instrumentation/plugins/ai-sdk-plugin.streaming.test.ts index 170bda408..9ea0c1eb9 100644 --- a/js/src/instrumentation/plugins/ai-sdk-plugin.streaming.test.ts +++ b/js/src/instrumentation/plugins/ai-sdk-plugin.streaming.test.ts @@ -1,5 +1,6 @@ /* eslint-disable @typescript-eslint/no-explicit-any */ /* eslint-disable @typescript-eslint/consistent-type-assertions */ +import * as ai from "ai"; import { afterEach, beforeAll, @@ -8,13 +9,12 @@ import { expect, test, } from "vitest"; -import * as ai from "ai"; -import { configureNode } from "../../node/config"; import { + TestBackgroundLogger, _exportsForTestingOnly, initLogger, - TestBackgroundLogger, } from "../../logger"; +import { configureNode } from "../../node/config"; import { wrapAISDK, wrapAgentClass } from "../../wrappers/ai-sdk"; import { BraintrustMiddleware } from "../../wrappers/ai-sdk/deprecated/BraintrustMiddleware"; import { @@ -183,9 +183,11 @@ describe("AI SDK streaming instrumentation", () => { }); const params = { model, prompt: "Say hello" }; - await aiSDKChannels.generateText.tracePromise( + await aiSDKChannels.generateText.invoke( () => model.doGenerate(params), - { arguments: [params] } as any, + undefined, + [params], + {}, ); const spans = (await backgroundLogger.drain()) as any[]; @@ -282,7 +284,7 @@ describe("AI SDK streaming instrumentation", () => { }, }; - await aiSDKChannels.generateText.tracePromise( + await aiSDKChannels.generateText.invoke( async () => { const result = params.tools.get_weather.execute(input, options); if (Symbol.asyncIterator in result) { @@ -296,7 +298,9 @@ describe("AI SDK streaming instrumentation", () => { } return { text: "done" }; }, - { arguments: [params] } as any, + undefined, + [params], + {}, ); const spans = (await backgroundLogger.drain()) as any[]; @@ -336,11 +340,11 @@ describe("AI SDK streaming instrumentation", () => { prompt: "Say hello.", maxOutputTokens: 16, }; - const result = (await aiSDKChannels.generateText.tracePromise( + const result = (await aiSDKChannels.generateText.invoke( async () => params.model.doGenerate(params), - { - arguments: [params], - } as any, + undefined, + [params], + {}, )) as any; expect(result.text).toBe("hello"); @@ -397,11 +401,11 @@ describe("AI SDK streaming instrumentation", () => { prompt: "Say hello.", maxOutputTokens: 16, }; - const result = (await aiSDKChannels.streamText.tracePromise( + const result = (await aiSDKChannels.streamText.invoke( async () => params.model.doStream(params), - { - arguments: [params], - } as any, + undefined, + [params], + {}, )) as any; for await (const _chunk of result.stream) { @@ -460,11 +464,11 @@ describe("AI SDK streaming instrumentation", () => { prompt: "Say hello.", maxOutputTokens: 16, }; - const result = (await aiSDKChannels.streamText.tracePromise( + const result = (await aiSDKChannels.streamText.invoke( async () => params.model.doStream(params), - { - arguments: [params], - } as any, + undefined, + [params], + {}, )) as any; const chunks: any[] = []; @@ -1093,7 +1097,7 @@ describe("AI SDK streaming instrumentation", () => { const contentDelayMs = 80; let sentContent = false; - const result = (await aiSDKChannels.streamText.tracePromise( + const result = (await aiSDKChannels.streamText.invoke( async () => ({ baseStream: new ReadableStream({ start(controller) { @@ -1122,14 +1126,14 @@ describe("AI SDK streaming instrumentation", () => { }, }), }), - { - arguments: [ - { - model: "mock-tool-model", - prompt: "Call the lookup tool.", - }, - ], - } as any, + undefined, + [ + { + model: "mock-tool-model", + prompt: "Call the lookup tool.", + }, + ], + {}, )) as any; const reader = result.baseStream.getReader(); @@ -1158,7 +1162,7 @@ describe("AI SDK streaming instrumentation", () => { try { let chunkSent = false; - const result = (await aiSDKChannels.streamText.tracePromise( + const result = (await aiSDKChannels.streamText.invoke( async () => { const resultRecord = { baseStream: new ReadableStream({ @@ -1199,14 +1203,14 @@ describe("AI SDK streaming instrumentation", () => { return resultRecord; }, - { - arguments: [ - { - model: "mock-stream-model", - prompt: "Reply with fresh.", - }, - ], - } as any, + undefined, + [ + { + model: "mock-stream-model", + prompt: "Reply with fresh.", + }, + ], + {}, )) as any; expect( @@ -1237,7 +1241,7 @@ describe("AI SDK streaming instrumentation", () => { plugin.enable(); try { - const result = (await aiSDKChannels.streamText.tracePromise( + const result = (await aiSDKChannels.streamText.invoke( async () => { const resultRecord = { stream: new ReadableStream({ @@ -1265,14 +1269,14 @@ describe("AI SDK streaming instrumentation", () => { return resultRecord; }, - { - arguments: [ - { - model: "mock-v7-stream-model", - prompt: "Reply with v7.", - }, - ], - } as any, + undefined, + [ + { + model: "mock-v7-stream-model", + prompt: "Reply with v7.", + }, + ], + {}, )) as any; expect(result.stream.pipeThrough).toEqual(expect.any(Function)); diff --git a/js/src/instrumentation/plugins/ai-sdk-plugin.test.ts b/js/src/instrumentation/plugins/ai-sdk-plugin.test.ts index 015b3763e..2e0d64ba5 100644 --- a/js/src/instrumentation/plugins/ai-sdk-plugin.test.ts +++ b/js/src/instrumentation/plugins/ai-sdk-plugin.test.ts @@ -1,4 +1,11 @@ -import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const telemetryMocks = vi.hoisted(() => ({ braintrustAISDKTelemetry: vi.fn(), @@ -11,34 +18,33 @@ const telemetryMocks = vi.hoisted(() => ({ }, })); -// Mock iso's newTracingChannel - must be before any imports that use it +// Mock platform context independently of invocation hooks. vi.mock("../../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(), - }, + default: {}, })); vi.mock("../../wrappers/ai-sdk/telemetry", () => ({ braintrustAISDKTelemetry: telemetryMocks.braintrustAISDKTelemetry, })); +import { BRAINTRUST_AI_SDK_V7_OPERATION_KEY as AI_SDK_V7_OPERATION_KEY } from "../../vendor-sdk-types/ai-sdk-v7-telemetry"; +import { serializeAISDKToolsForLogging } from "../../wrappers/ai-sdk/tool-serialization"; import { AISDKPlugin, DEFAULT_DENY_OUTPUT_PATHS, + extractTokenMetrics, processAISDKCallInput, processAISDKGenerateImageInput, + processAISDKGenerateImageOutput, + processAISDKOutput as processAISDKOutputActual, processAISDKWorkflowAgentCallInput, processAISDKWorkflowAgentModelCallInput, - processAISDKOutput as processAISDKOutputActual, - processAISDKGenerateImageOutput, - extractTokenMetrics, serializeModelWithProvider, } from "./ai-sdk-plugin"; -import iso from "../../isomorph"; -import { serializeAISDKToolsForLogging } from "../../wrappers/ai-sdk/tool-serialization"; -import { BRAINTRUST_AI_SDK_V7_OPERATION_KEY as AI_SDK_V7_OPERATION_KEY } from "../../vendor-sdk-types/ai-sdk-v7-telemetry"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; type MockTracingChannel = { handlers: any[]; hasSubscribers: boolean; @@ -67,7 +73,7 @@ describe("AISDKPlugin", () => { telemetryMocks.braintrustAISDKTelemetry.mockReturnValue( telemetryMocks.telemetry, ); - mockNewTracingChannel.mockImplementation((name: string) => { + mockNewInvocationHook.mockImplementation((name: string) => { const channel: MockTracingChannel = { handlers: [], hasSubscribers: false, @@ -249,7 +255,7 @@ describe("AISDKPlugin", () => { plugin.enable(); expect( - mockChannels.get("orchestrion:ai:generateImage")?.subscribe, + mockChannels.get("orchestrion:ai:generateImage")?.intercept, ).toHaveBeenCalledTimes(1); }); diff --git a/js/src/instrumentation/plugins/ai-sdk-plugin.ts b/js/src/instrumentation/plugins/ai-sdk-plugin.ts index 989216a67..8b8b5028d 100644 --- a/js/src/instrumentation/plugins/ai-sdk-plugin.ts +++ b/js/src/instrumentation/plugins/ai-sdk-plugin.ts @@ -1,70 +1,38 @@ -import { BasePlugin, toLoggedError } from "../core"; +import type { ReturnOf } from "../core/channel-definitions"; import { debugLogger } from "../../debug-logger"; +import { BasePlugin, toLoggedError } from "../core"; import { - traceAsyncChannel, - traceStreamingChannel, - traceSyncStreamChannel, + traceAsyncCall, + traceStreamingCall, + traceSyncStreamCall, unsubscribeAll, } from "../core/channel-tracing"; -import type { ChannelMessage } from "../core/channel-definitions"; -import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; -import type { IsoChannelHandlers } from "../../isomorph"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + import { SpanTypeAttribute, isObject, isPromiseLike, } from "../../../util/index"; -import { getCurrentUnixTimestamp } from "../../util"; import { - _internalStartSpanWithInitialMerge, Attachment, + _internalStartSpanWithInitialMerge, currentSpan, startSpan, - type Span, withCurrent, + type Span, } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { - convertDataToBlob, - getExtensionFromMediaType, -} from "../../wrappers/attachment-utils"; -import { normalizeAISDKLoggedOutput } from "../../wrappers/ai-sdk/normalize-logged-output"; -import { serializeAISDKToolsForLogging } from "../../wrappers/ai-sdk/tool-serialization"; -import { braintrustAISDKTelemetry } from "../../wrappers/ai-sdk/telemetry"; -import { - bindHarnessTurnParentToStart, - captureHarnessCreateSessionParent, - endHarnessTurn, - harnessContinuationParent, - registerHarnessSessionParent, - registerHarnessTurnSpan, - updateHarnessTurn, - type HarnessTurnParent, -} from "../../wrappers/ai-sdk/harness-agent-context"; -import { - registerWorkflowAgentWrapperSpan, - unregisterWorkflowAgentWrapperSpan, -} from "../../wrappers/ai-sdk/workflow-agent-context"; -import { zodToJsonSchema } from "../../zod/utils"; -import { - aiSDKChannels, - BRAINTRUST_WRAPPED_AI_SDK_MODEL, - harnessAgentChannels, -} from "./ai-sdk-channels"; -import { extractTokenMetrics } from "./ai-sdk-metrics"; -import { currentCloudflareThinkSpan } from "./cloudflare-think-context"; -import { - isAutoInstrumentationSuppressed, - runWithAutoInstrumentationSuppressed, -} from "../auto-instrumentation-suppression"; +import { getCurrentUnixTimestamp } from "../../util"; import type { AISDK, AISDKCallParams, AISDKEmbedParams, AISDKEmbeddingResult, + AISDKGenerateImageParams, AISDKGeneratedFile, AISDKHarnessAgentCallParams, AISDKHarnessAgentSettings, @@ -75,7 +43,6 @@ import type { AISDKOutputResponseFormat, AISDKRerankParams, AISDKRerankResult, - AISDKGenerateImageParams, AISDKResult, AISDKTool, AISDKTools, @@ -86,6 +53,41 @@ import type { AISDKV7TelemetryOptions, } from "../../vendor-sdk-types/ai-sdk-v7-telemetry"; import { BRAINTRUST_AI_SDK_V7_OPERATION_KEY as AI_SDK_V7_OPERATION_KEY } from "../../vendor-sdk-types/ai-sdk-v7-telemetry"; +import { + captureHarnessCreateSessionParent, + endHarnessTurn, + harnessContinuationParent, + registerHarnessSessionParent, + registerHarnessTurnSpan, + runWithHarnessTurnParent, + updateHarnessTurn, + type HarnessTurnParent, +} from "../../wrappers/ai-sdk/harness-agent-context"; +import { normalizeAISDKLoggedOutput } from "../../wrappers/ai-sdk/normalize-logged-output"; +import { braintrustAISDKTelemetry } from "../../wrappers/ai-sdk/telemetry"; +import { serializeAISDKToolsForLogging } from "../../wrappers/ai-sdk/tool-serialization"; +import { + registerWorkflowAgentWrapperSpan, + unregisterWorkflowAgentWrapperSpan, +} from "../../wrappers/ai-sdk/workflow-agent-context"; +import { + convertDataToBlob, + getExtensionFromMediaType, +} from "../../wrappers/attachment-utils"; +import { zodToJsonSchema } from "../../zod/utils"; +import { + isAutoInstrumentationSuppressed, + runWithAutoInstrumentationSuppressed, +} from "../auto-instrumentation-suppression"; +import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; +import type { ChannelMessage } from "../core/tracing-types"; +import { + BRAINTRUST_WRAPPED_AI_SDK_MODEL, + aiSDKChannels, + harnessAgentChannels, +} from "./ai-sdk-channels"; +import { extractTokenMetrics } from "./ai-sdk-metrics"; +import { currentCloudflareThinkSpan } from "./cloudflare-think-context"; interface AISDKPluginConfig { /** @@ -213,482 +215,658 @@ export class AISDKPlugin extends BasePlugin { subscribeToHarnessAgentCreateSession(), ); this.unsubscribers.push( - subscribeToHarnessContinuation( - harnessAgentChannels.continueGenerate, - denyOutputPaths, + harnessAgentChannels.continueGenerate.intercept( + (target, receiver, args, additional) => + traceHarnessContinuation( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + denyOutputPaths, + ), ), - subscribeToHarnessContinuation( - harnessAgentChannels.continueStream, - denyOutputPaths, + harnessAgentChannels.continueStream.intercept( + (target, receiver, args, additional) => + traceHarnessContinuation( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + denyOutputPaths, + ), ), ); // generateText - async function that may return streams this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.generateText, { - name: "generateText", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths), - extractOutput: (result, endEvent) => { - finalizeAISDKChildTracing(endEvent as { [key: string]: unknown }); - return processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ); - }, - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - aggregateChunks: aggregateAISDKChunks, - }), + aiSDKChannels.generateText.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "generateText", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths), + extractOutput: (result, endEvent) => { + finalizeAISDKChildTracing( + endEvent as { [key: string]: unknown }, + ); + return processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ); + }, + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + aggregateChunks: aggregateAISDKChunks, + }, + ), + ), ); // generateImage - async image generation function (v5 experimental, v6+ stable) this.unsubscribers.push( - traceAsyncChannel(aiSDKChannels.generateImage, { - name: "generateImage", - type: SpanTypeAttribute.LLM, - extractInput: ([params], event) => - prepareAISDKGenerateImageInput(params, event.self), - extractOutput: (result, endEvent) => - processAISDKGenerateImageOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), + aiSDKChannels.generateImage.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "generateImage", + type: SpanTypeAttribute.LLM, + extractInput: ([params], event) => + prepareAISDKGenerateImageInput(params, event.self), + extractOutput: (result, endEvent) => + processAISDKGenerateImageOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result) => extractTokenMetrics(result), + }, ), - extractMetrics: (result) => extractTokenMetrics(result), - }), + ), ); // streamText - function returning stream this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.streamText, { - name: "streamText", - type: SpanTypeAttribute.FUNCTION, - shouldTrace: () => currentCloudflareThinkSpan() === undefined, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ), - extractMetrics: (result, startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent, startTime), - aggregateChunks: aggregateAISDKChunks, - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - span, - startTime, - }), - }), + aiSDKChannels.streamText.intercept((target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "streamText", + type: SpanTypeAttribute.FUNCTION, + shouldTrace: () => currentCloudflareThinkSpan() === undefined, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent, startTime), + aggregateChunks: aggregateAISDKChunks, + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + span, + startTime, + }), + }, + ), + ), ); // streamText - sync function returning stream (v4+, used by auto-hook) this.unsubscribers.push( - traceSyncStreamChannel(aiSDKChannels.streamTextSync, { - name: "streamText", - type: SpanTypeAttribute.FUNCTION, - shouldTrace: () => currentCloudflareThinkSpan() === undefined, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths), - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - span, - startTime, - }), - }), + aiSDKChannels.streamTextSync.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "streamText", + type: SpanTypeAttribute.FUNCTION, + shouldTrace: () => currentCloudflareThinkSpan() === undefined, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths), + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + span, + startTime, + }), + }, + ), + ), ); // generateObject - async function that may return streams this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.generateObject, { - name: "generateObject", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths), - extractOutput: (result, endEvent) => { - finalizeAISDKChildTracing(endEvent as { [key: string]: unknown }); - return processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ); - }, - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - aggregateChunks: aggregateAISDKChunks, - }), + aiSDKChannels.generateObject.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "generateObject", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths), + extractOutput: (result, endEvent) => { + finalizeAISDKChildTracing( + endEvent as { [key: string]: unknown }, + ); + return processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ); + }, + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + aggregateChunks: aggregateAISDKChunks, + }, + ), + ), ); // streamObject - function returning stream this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.streamObject, { - name: "streamObject", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), + aiSDKChannels.streamObject.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "streamObject", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent, startTime), + aggregateChunks: aggregateAISDKChunks, + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + span, + startTime, + }), + }, ), - extractMetrics: (result, startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent, startTime), - aggregateChunks: aggregateAISDKChunks, - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - span, - startTime, - }), - }), + ), ); // streamObject - sync function returning stream (v4+, used by auto-hook) this.unsubscribers.push( - traceSyncStreamChannel(aiSDKChannels.streamObjectSync, { - name: "streamObject", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths), - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - span, - startTime, - }), - }), + aiSDKChannels.streamObjectSync.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "streamObject", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths), + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + span, + startTime, + }), + }, + ), + ), ); // embed - async embedding function this.unsubscribers.push( - traceAsyncChannel(aiSDKChannels.embed, { - name: "embed", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event) => - prepareAISDKEmbedInput(params, event.self), - extractOutput: (result, endEvent) => - processAISDKEmbeddingOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ), - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - }), + aiSDKChannels.embed.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "embed", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event) => + prepareAISDKEmbedInput(params, event.self), + extractOutput: (result, endEvent) => + processAISDKEmbeddingOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + }, + ), + ), ); // embedMany - async embedding batch function this.unsubscribers.push( - traceAsyncChannel(aiSDKChannels.embedMany, { - name: "embedMany", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event) => - prepareAISDKEmbedInput(params, event.self), - extractOutput: (result, endEvent) => - processAISDKEmbeddingOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ), - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - }), + aiSDKChannels.embedMany.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "embedMany", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event) => + prepareAISDKEmbedInput(params, event.self), + extractOutput: (result, endEvent) => + processAISDKEmbeddingOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + }, + ), + ), ); // rerank - async reranking function this.unsubscribers.push( - traceAsyncChannel(aiSDKChannels.rerank, { - name: "rerank", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event) => - prepareAISDKRerankInput(params, event.self), - extractOutput: (result, endEvent) => - processAISDKRerankOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ), - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - }), + aiSDKChannels.rerank.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "rerank", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event) => + prepareAISDKRerankInput(params, event.self), + extractOutput: (result, endEvent) => + processAISDKRerankOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + }, + ), + ), ); // Agent.generate - async method this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.agentGenerate, { - name: "Agent.generate", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths, { - agentOwner: true, - }), - extractOutput: (result, endEvent) => { - finalizeAISDKChildTracing(endEvent as { [key: string]: unknown }); - return processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ); - }, - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - aggregateChunks: aggregateAISDKChunks, - }), + aiSDKChannels.agentGenerate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "Agent.generate", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths, { + agentOwner: true, + }), + extractOutput: (result, endEvent) => { + finalizeAISDKChildTracing( + endEvent as { [key: string]: unknown }, + ); + return processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ); + }, + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + aggregateChunks: aggregateAISDKChunks, + }, + ), + ), ); // Agent.stream - async method returning stream (v5, used by wrapAISDK) this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.agentStream, { - name: "Agent.stream", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths, { - agentOwner: true, - }), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), + aiSDKChannels.agentStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "Agent.stream", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths, { + agentOwner: true, + }), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent, startTime), + aggregateChunks: aggregateAISDKChunks, + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + span, + startTime, + }), + }, ), - extractMetrics: (result, startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent, startTime), - aggregateChunks: aggregateAISDKChunks, - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - span, - startTime, - }), - }), + ), ); // Agent.stream - sync method returning stream (v5, used by auto-hook) this.unsubscribers.push( - traceSyncStreamChannel(aiSDKChannels.agentStreamSync, { - name: "Agent.stream", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths, { - agentOwner: true, - }), - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - span, - startTime, - }), - }), + aiSDKChannels.agentStreamSync.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "Agent.stream", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths, { + agentOwner: true, + }), + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + span, + startTime, + }), + }, + ), + ), ); // HarnessAgent.generate - one task span per agent turn this.unsubscribers.push( - traceStreamingChannel(harnessAgentChannels.generate, { - name: "HarnessAgent.generate", - startSpan: _internalStartSpanWithInitialMerge, - type: SpanTypeAttribute.TASK, - extractInput: ([params], event, span) => - prepareAISDKHarnessAgentInput(params, event.self, span), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), + harnessAgentChannels.generate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "HarnessAgent.generate", + startSpan: _internalStartSpanWithInitialMerge, + type: SpanTypeAttribute.TASK, + extractInput: ([params], event, span) => + prepareAISDKHarnessAgentInput(params, event.self, span), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result) => extractTokenMetrics(result), + aggregateChunks: aggregateAISDKChunks, + }, ), - extractMetrics: (result) => extractTokenMetrics(result), - aggregateChunks: aggregateAISDKChunks, - }), + ), ); // HarnessAgent.stream - async method returning an AI SDK stream result this.unsubscribers.push( - traceStreamingChannel(harnessAgentChannels.stream, { - name: "HarnessAgent.stream", - startSpan: _internalStartSpanWithInitialMerge, - type: SpanTypeAttribute.TASK, - extractInput: ([params], event, span) => - prepareAISDKHarnessAgentInput(params, event.self, span), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ), - extractMetrics: (result, startTime) => ({ - ...extractTokenMetrics(result), - ...(startTime === undefined - ? {} - : { - time_to_first_token: getCurrentUnixTimestamp() - startTime, + harnessAgentChannels.stream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "HarnessAgent.stream", + startSpan: _internalStartSpanWithInitialMerge, + type: SpanTypeAttribute.TASK, + extractInput: ([params], event, span) => + prepareAISDKHarnessAgentInput(params, event.self, span), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, startTime) => ({ + ...extractTokenMetrics(result), + ...(startTime === undefined + ? {} + : { + time_to_first_token: + getCurrentUnixTimestamp() - startTime, + }), }), - }), - aggregateChunks: aggregateAISDKChunks, - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - resolvePromiseUsage: true, - span, - startTime, - }), - }), + aggregateChunks: aggregateAISDKChunks, + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + resolvePromiseUsage: true, + span, + startTime, + }), + }, + ), + ), ); // Trace a continuation as its own task only when its original turn cannot // be recovered. Known continuations extend the original Harness task. this.unsubscribers.push( - traceStreamingChannel(harnessAgentChannels.continueGenerate, { - name: "HarnessAgent.continueGenerate", - shouldTrace: (args) => - !harnessContinuationParent(harnessSessionFromArguments(args)), - startSpan: _internalStartSpanWithInitialMerge, - type: SpanTypeAttribute.TASK, - extractInput: ([params], event, span) => - prepareAISDKHarnessAgentInput(params, event.self, span), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), + harnessAgentChannels.continueGenerate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "HarnessAgent.continueGenerate", + shouldTrace: (args) => + !harnessContinuationParent(harnessSessionFromArguments(args)), + startSpan: _internalStartSpanWithInitialMerge, + type: SpanTypeAttribute.TASK, + extractInput: ([params], event, span) => + prepareAISDKHarnessAgentInput(params, event.self, span), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result) => extractTokenMetrics(result), + aggregateChunks: aggregateAISDKChunks, + }, ), - extractMetrics: (result) => extractTokenMetrics(result), - aggregateChunks: aggregateAISDKChunks, - }), + ), ); this.unsubscribers.push( - traceStreamingChannel(harnessAgentChannels.continueStream, { - name: "HarnessAgent.continueStream", - shouldTrace: (args) => - !harnessContinuationParent(harnessSessionFromArguments(args)), - startSpan: _internalStartSpanWithInitialMerge, - type: SpanTypeAttribute.TASK, - extractInput: ([params], event, span) => - prepareAISDKHarnessAgentInput(params, event.self, span), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ), - extractMetrics: (result, startTime) => ({ - ...extractTokenMetrics(result), - ...(startTime === undefined - ? {} - : { - time_to_first_token: getCurrentUnixTimestamp() - startTime, + harnessAgentChannels.continueStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "HarnessAgent.continueStream", + shouldTrace: (args) => + !harnessContinuationParent(harnessSessionFromArguments(args)), + startSpan: _internalStartSpanWithInitialMerge, + type: SpanTypeAttribute.TASK, + extractInput: ([params], event, span) => + prepareAISDKHarnessAgentInput(params, event.self, span), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, startTime) => ({ + ...extractTokenMetrics(result), + ...(startTime === undefined + ? {} + : { + time_to_first_token: + getCurrentUnixTimestamp() - startTime, + }), }), - }), - aggregateChunks: aggregateAISDKChunks, - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - resolvePromiseUsage: true, - span, - startTime, - }), - }), + aggregateChunks: aggregateAISDKChunks, + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + resolvePromiseUsage: true, + span, + startTime, + }), + }, + ), + ), ); // ToolLoopAgent.generate - async method this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.toolLoopAgentGenerate, { - name: "ToolLoopAgent.generate", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths, { - agentOwner: true, - }), - extractOutput: (result, endEvent) => { - finalizeAISDKChildTracing(endEvent as { [key: string]: unknown }); - return processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ); - }, - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - aggregateChunks: aggregateAISDKChunks, - }), + aiSDKChannels.toolLoopAgentGenerate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "ToolLoopAgent.generate", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths, { + agentOwner: true, + }), + extractOutput: (result, endEvent) => { + finalizeAISDKChildTracing( + endEvent as { [key: string]: unknown }, + ); + return processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ); + }, + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + aggregateChunks: aggregateAISDKChunks, + }, + ), + ), ); // ToolLoopAgent.stream - async method returning stream this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.toolLoopAgentStream, { - name: "ToolLoopAgent.stream", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKCallInput(params, event, span, denyOutputPaths, { - agentOwner: true, - }), - extractOutput: (result, endEvent) => - processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), + aiSDKChannels.toolLoopAgentStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "ToolLoopAgent.stream", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKCallInput(params, event, span, denyOutputPaths, { + agentOwner: true, + }), + extractOutput: (result, endEvent) => + processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ), + extractMetrics: (result, startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent, startTime), + aggregateChunks: aggregateAISDKChunks, + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + result, + span, + startTime, + }), + }, ), - extractMetrics: (result, startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent, startTime), - aggregateChunks: aggregateAISDKChunks, - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - result, - span, - startTime, - }), - }), + ), ); // WorkflowAgent.stream - async method returning stream this.unsubscribers.push( - traceStreamingChannel(aiSDKChannels.workflowAgentStream, { - name: "WorkflowAgent.stream", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params], event, span) => - prepareAISDKWorkflowAgentStreamInput( - params, - event, - span, - denyOutputPaths, + aiSDKChannels.workflowAgentStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.AI_SDK, + name: "WorkflowAgent.stream", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params], event, span) => + prepareAISDKWorkflowAgentStreamInput( + params, + event, + span, + denyOutputPaths, + ), + extractOutput: (result, endEvent) => { + finalizeAISDKChildTracing( + endEvent as { [key: string]: unknown }, + ); + return processAISDKOutput( + result, + resolveDenyOutputPaths(endEvent, denyOutputPaths), + ); + }, + extractMetrics: (result, _startTime, endEvent) => + extractTopLevelAISDKMetrics(result, endEvent), + aggregateChunks: aggregateAISDKChunks, + onComplete: ({ span }) => { + unregisterWorkflowAgentWrapperSpan(span); + }, + onError: ({ event, span }) => { + finalizeAISDKChildTracing(event as { [key: string]: unknown }); + unregisterWorkflowAgentWrapperSpan(span); + }, + patchResult: ({ endEvent, result, span, startTime }) => + patchAISDKStreamingResult({ + defaultDenyOutputPaths: denyOutputPaths, + endEvent, + onComplete: () => unregisterWorkflowAgentWrapperSpan(span), + onCancel: () => unregisterWorkflowAgentWrapperSpan(span), + onError: () => unregisterWorkflowAgentWrapperSpan(span), + result, + span, + startTime, + }), + }, ), - extractOutput: (result, endEvent) => { - finalizeAISDKChildTracing(endEvent as { [key: string]: unknown }); - return processAISDKOutput( - result, - resolveDenyOutputPaths(endEvent, denyOutputPaths), - ); - }, - extractMetrics: (result, _startTime, endEvent) => - extractTopLevelAISDKMetrics(result, endEvent), - aggregateChunks: aggregateAISDKChunks, - onComplete: ({ span }) => { - unregisterWorkflowAgentWrapperSpan(span); - }, - onError: ({ event, span }) => { - finalizeAISDKChildTracing(event as { [key: string]: unknown }); - unregisterWorkflowAgentWrapperSpan(span); - }, - patchResult: ({ endEvent, result, span, startTime }) => - patchAISDKStreamingResult({ - defaultDenyOutputPaths: denyOutputPaths, - endEvent, - onComplete: () => unregisterWorkflowAgentWrapperSpan(span), - onCancel: () => unregisterWorkflowAgentWrapperSpan(span), - onError: () => unregisterWorkflowAgentWrapperSpan(span), - result, - span, - startTime, - }), - }), + ), ); } } @@ -828,28 +1006,58 @@ function addEvaluationIds( } function subscribeToHarnessAgentCreateSession(): () => void { - const channel = harnessAgentChannels.createSession.tracingChannel(); + const channel = harnessAgentChannels.createSession; const parents = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - const parent = captureHarnessCreateSessionParent(event.arguments?.[0]); - if (parent) { - parents.set(event, parent); + + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + const parent = captureHarnessCreateSessionParent(event.arguments?.[0]); + if (parent) { + parents.set(event, parent); + } + }; + const resolved = ( + event: ChannelMessage, + ) => { + registerHarnessSessionParent(event.result, parents.get(event)); + parents.delete(event); + }; + const failed = ( + event: ChannelMessage, + ) => { + parents.delete(event); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); }, - asyncEnd: (event) => { - registerHarnessSessionParent(event.result, parents.get(event)); - parents.delete(event); - }, - error: (event) => { - parents.delete(event); - }, - }; - - channel.subscribe(handlers); - return () => channel.unsubscribe(handlers); + ); + return () => removeHandlers(); } type HarnessContinuationChannel = typeof harnessAgentChannels.continueGenerate; @@ -864,146 +1072,159 @@ function harnessContinuationParentFromEvent( } } -function subscribeToHarnessContinuation( - continuationChannel: HarnessContinuationChannel, +function traceHarnessContinuation( + call: () => ReturnOf, + event: ChannelMessage, defaultDenyOutputPaths: string[], -): () => void { - const channel = continuationChannel.tracingChannel(); +): ReturnOf { const parents = new WeakMap(); const startTimes = new WeakMap(); - const unbindParentStore = bindHarnessTurnParentToStart( - channel, - harnessContinuationParentFromEvent, - ); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - const parent = harnessContinuationParentFromEvent(event); - if (!parent) { - return; - } + const prepare = (event: ChannelMessage) => { + const parent = harnessContinuationParentFromEvent(event); + if (!parent) { + return; + } - parents.set(event, parent); - startTimes.set(event, getCurrentUnixTimestamp()); - try { - const params = event.arguments?.[0]; - if (params) { - // Harness reads telemetry from its settings after this event starts. - // The original task already contains the turn input, so only install - // telemetry here; continuation work is logged onto that task. - prepareAISDKHarnessAgentInput(params, event.self); - } - } catch (error) { - debugLogger.error( - "Error preparing Harness continuation telemetry:", - error, - ); - } - }, - asyncEnd: (event) => { - const parent = parents.get(event); - const startTime = startTimes.get(event) ?? getCurrentUnixTimestamp(); - parents.delete(event); - startTimes.delete(event); - if (!parent) { - return; + parents.set(event, parent); + startTimes.set(event, getCurrentUnixTimestamp()); + try { + const params = event.arguments?.[0]; + if (params) { + // Harness reads telemetry from its settings after this event starts. + // The original task already contains the turn input, so only install + // telemetry here; continuation work is logged onto that task. + prepareAISDKHarnessAgentInput(params, event.self); } + } catch (error) { + debugLogger.error( + "Error preparing Harness continuation telemetry:", + error, + ); + } + }; + const resolved = (event: ChannelMessage) => { + const parent = parents.get(event); + const startTime = startTimes.get(event) ?? getCurrentUnixTimestamp(); + parents.delete(event); + startTimes.delete(event); + if (!parent) { + return; + } - const endEvent = event as ChannelMessage & { - result: AISDKResult | AsyncIterable; - }; - const span = { - end: () => endHarnessTurn(parent), - log: (update: Parameters[0]) => - updateHarnessTurn( - parent, - Object.prototype.hasOwnProperty.call(update, "error") && - !Object.prototype.hasOwnProperty.call(update, "output") - ? { ...update, output: null } - : update, - ), - }; + const endEvent = event as ChannelMessage & { + result: AISDKResult | AsyncIterable; + }; + const span = { + end: () => endHarnessTurn(parent), + log: (update: Parameters[0]) => + updateHarnessTurn( + parent, + Object.prototype.hasOwnProperty.call(update, "error") && + !Object.prototype.hasOwnProperty.call(update, "output") + ? { ...update, output: null } + : update, + ), + }; - try { - if (isAsyncIterable(endEvent.result)) { - patchStreamIfNeeded(endEvent.result, { - onComplete: (chunks) => { - try { - const { metadata, metrics, output } = aggregateAISDKChunks( - chunks, - endEvent.result, - endEvent, - ); - span.log({ - ...(metadata ? { metadata } : {}), - metrics, - output, - }); - } catch (error) { - debugLogger.error( - "Error aggregating Harness continuation stream:", - error, - ); - } finally { - span.end(); - } - }, - onError: (error) => { - span.log({ error: toLoggedError(error), output: null }); + try { + if (isAsyncIterable(endEvent.result)) { + patchStreamIfNeeded(endEvent.result, { + onComplete: (chunks) => { + try { + const { metadata, metrics, output } = aggregateAISDKChunks( + chunks, + endEvent.result, + endEvent, + ); + span.log({ + ...(metadata ? { metadata } : {}), + metrics, + output, + }); + } catch (error) { + debugLogger.error( + "Error aggregating Harness continuation stream:", + error, + ); + } finally { span.end(); - }, - }); - return; - } - - if ( - patchAISDKStreamingResult({ - defaultDenyOutputPaths, - endEvent, - result: endEvent.result, - resolvePromiseUsage: true, - span, - startTime, - }) - ) { - return; - } - - finalizeAISDKChildTracing(endEvent); - span.log({ - metrics: extractTokenMetrics(endEvent.result), - output: processAISDKOutput( - endEvent.result, - resolveDenyOutputPaths(endEvent, defaultDenyOutputPaths), - ), + } + }, + onError: (error) => { + span.log({ error: toLoggedError(error), output: null }); + span.end(); + }, }); - span.end(); - } catch (error) { - debugLogger.error("Error tracing Harness continuation:", error); - span.end(); + return; } - }, - error: (event) => { - const parent = parents.get(event); - parents.delete(event); - startTimes.delete(event); - if (!parent) { + + if ( + patchAISDKStreamingResult({ + defaultDenyOutputPaths, + endEvent, + result: endEvent.result, + resolvePromiseUsage: true, + span, + startTime, + }) + ) { return; } - updateHarnessTurn(parent, { - error: toLoggedError(event.error), - output: null, + + finalizeAISDKChildTracing(endEvent); + span.log({ + metrics: extractTokenMetrics(endEvent.result), + output: processAISDKOutput( + endEvent.result, + resolveDenyOutputPaths(endEvent, defaultDenyOutputPaths), + ), }); - endHarnessTurn(parent); - }, + span.end(); + } catch (error) { + debugLogger.error("Error tracing Harness continuation:", error); + span.end(); + } }; + const failed = (event: ChannelMessage) => { + const parent = parents.get(event); + parents.delete(event); + startTimes.delete(event); + if (!parent) { + return; + } + updateHarnessTurn(parent, { + error: toLoggedError(event.error), + output: null, + }); + endHarnessTurn(parent); + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = call(); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } - channel.subscribe(handlers); - return () => { - unbindParentStore(); - channel.unsubscribe(handlers); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); }; + return runWithHarnessTurnParent( + harnessContinuationParentFromEvent(event), + invoke, + ); } function interceptAISDKV7TelemetryDispatcher(): () => void { diff --git a/js/src/instrumentation/plugins/anthropic-channels.ts b/js/src/instrumentation/plugins/anthropic-channels.ts index e1a8793ec..ca017ceb5 100644 --- a/js/src/instrumentation/plugins/anthropic-channels.ts +++ b/js/src/instrumentation/plugins/anthropic-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { AnthropicCreateParams, AnthropicMessage, @@ -15,52 +15,43 @@ import type { type AnthropicResult = AnthropicMessage | AnthropicMessageStream; -export const anthropicChannels = defineChannels( - "@anthropic-ai/sdk", - { - messagesCreate: channel< - [AnthropicCreateParams], - AnthropicResult, - Record, - AnthropicStreamEvent - >({ - channelName: "messages.create", - kind: "async", - }), - betaMessagesCreate: channel< - [AnthropicCreateParams], - AnthropicResult, - Record, - AnthropicStreamEvent - >({ - channelName: "beta.messages.create", - kind: "async", - }), - betaMessagesToolRunner: channel< - [AnthropicToolRunnerParams], - AnthropicToolRunner - >({ - channelName: "beta.messages.toolRunner", - kind: "sync-stream", - }), - betaSessionsEventsStream: channel< - [string, AnthropicSessionEventStreamParams?], - AnthropicSessionEventStream, - Record, - AnthropicSessionEvent - >({ - channelName: "beta.sessions.events.stream", - kind: "async", - }), - betaSessionsThreadsEventsStream: channel< - [string, AnthropicSessionThreadEventStreamParams], - AnthropicSessionEventStream, - Record, - AnthropicSessionEvent - >({ - channelName: "beta.sessions.threads.events.stream", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.ANTHROPIC }, -); +export const anthropicChannels = defineInterceptor("@anthropic-ai/sdk", { + messagesCreate: channel< + [AnthropicCreateParams], + PromiseLike, + Record, + AnthropicStreamEvent + >({ + channelName: "messages.create", + }), + betaMessagesCreate: channel< + [AnthropicCreateParams], + PromiseLike, + Record, + AnthropicStreamEvent + >({ + channelName: "beta.messages.create", + }), + betaMessagesToolRunner: channel< + [AnthropicToolRunnerParams], + AnthropicToolRunner + >({ + channelName: "beta.messages.toolRunner", + }), + betaSessionsEventsStream: channel< + [string, AnthropicSessionEventStreamParams?], + PromiseLike, + Record, + AnthropicSessionEvent + >({ + channelName: "beta.sessions.events.stream", + }), + betaSessionsThreadsEventsStream: channel< + [string, AnthropicSessionThreadEventStreamParams], + PromiseLike, + Record, + AnthropicSessionEvent + >({ + channelName: "beta.sessions.threads.events.stream", + }), +}); diff --git a/js/src/instrumentation/plugins/anthropic-plugin.test.ts b/js/src/instrumentation/plugins/anthropic-plugin.test.ts index a711c3b7a..45ba05661 100644 --- a/js/src/instrumentation/plugins/anthropic-plugin.test.ts +++ b/js/src/instrumentation/plugins/anthropic-plugin.test.ts @@ -1,10 +1,7 @@ import { describe, it, expect, vi } from "vitest"; -// Mock iso's newTracingChannel - must be before any imports that use it vi.mock("../../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(), - }, + default: {}, })); import { diff --git a/js/src/instrumentation/plugins/anthropic-plugin.ts b/js/src/instrumentation/plugins/anthropic-plugin.ts index e184c927d..985565b15 100644 --- a/js/src/instrumentation/plugins/anthropic-plugin.ts +++ b/js/src/instrumentation/plugins/anthropic-plugin.ts @@ -1,29 +1,25 @@ -import { BasePlugin, toLoggedError } from "../core"; -import { traceStreamingChannel, unsubscribeAll } from "../core/channel-tracing"; -import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; +import type { Span } from "../../logger"; import { Attachment, startSpan as startBaseSpan, withCurrent, } from "../../logger"; -import type { Span } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; -import { debugLogger } from "../../debug-logger"; +import { BasePlugin, toLoggedError } from "../core"; +import { traceStreamingCall, unsubscribeAll } from "../core/channel-tracing"; +import { observeResult, runInstrumentation } from "../core/observe-result"; +import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; + import { SpanTypeAttribute, isObject, isPromiseLike, } from "../../../util/index"; -import { isAutoInstrumentationSuppressed } from "../auto-instrumentation-suppression"; +import { debugLogger } from "../../debug-logger"; import { filterFrom, getCurrentUnixTimestamp } from "../../util"; -import { finalizeAnthropicTokens } from "../../wrappers/anthropic-tokens-util"; -import { registerAnthropicSessionStreamCollector } from "../../wrappers/anthropic-session-collector"; -import { anthropicChannels } from "./anthropic-channels"; import type { AnthropicBase64Source, AnthropicCitation, @@ -46,6 +42,11 @@ import type { AnthropicToolRunnerTool, AnthropicUsage, } from "../../vendor-sdk-types/anthropic"; +import { registerAnthropicSessionStreamCollector } from "../../wrappers/anthropic-session-collector"; +import { finalizeAnthropicTokens } from "../../wrappers/anthropic-tokens-util"; +import { isAutoInstrumentationSuppressed } from "../auto-instrumentation-suppression"; +import type { ChannelMessage } from "../core/tracing-types"; +import { anthropicChannels } from "./anthropic-channels"; type AnthropicToolRunnerState = { aggregatedMetrics: Record; @@ -164,95 +165,134 @@ export class AnthropicPlugin extends BasePlugin { // Messages API - supports streaming via stream=true parameter this.unsubscribers.push( - traceStreamingChannel(anthropicChannels.messagesCreate, anthropicConfig), + anthropicChannels.messagesCreate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + ...anthropicConfig, + instrumentationName: INSTRUMENTATION_NAMES.ANTHROPIC, + }, + ), + ), ); // Beta Messages API - supports streaming via stream=true parameter this.unsubscribers.push( - traceStreamingChannel(anthropicChannels.betaMessagesCreate, { - ...anthropicConfig, - name: "anthropic.messages.create", - }), + anthropicChannels.betaMessagesCreate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.ANTHROPIC, + ...anthropicConfig, + name: "anthropic.messages.create", + }, + ), + ), ); } private subscribeToAnthropicToolRunner(): void { - const tracingChannel = - anthropicChannels.betaMessagesToolRunner.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; + const invocationHook = anthropicChannels.betaMessagesToolRunner; const states = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - if (isAutoInstrumentationSuppressed()) { - return; - } + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage< + typeof anthropicChannels.betaMessagesToolRunner + >, + ) => { + if (isAutoInstrumentationSuppressed()) { + return; + } - const params = (event.arguments[0] ?? {}) as AnthropicToolRunnerParams; - const span = startBaseSpan( - withSpanInstrumentationName( - { - name: "anthropic.beta.messages.toolRunner", - spanAttributes: { - type: SpanTypeAttribute.TASK, + const params = (event.arguments[0] ?? + {}) as AnthropicToolRunnerParams; + const span = startBaseSpan( + withSpanInstrumentationName( + { + name: "anthropic.beta.messages.toolRunner", + spanAttributes: { + type: SpanTypeAttribute.TASK, + }, }, - }, - INSTRUMENTATION_NAMES.ANTHROPIC, - ), - ); - - span.log({ - input: processAttachmentsInInput( - coalesceInput(params.messages ?? [], params.system), - ), - metadata: { - ...extractAnthropicToolRunnerMetadata(params), - provider: "anthropic", - }, - }); - - const state = { - aggregatedMetrics: {}, - finalized: false, - iterationCount: 0, - seenMessages: new WeakSet(), - span, - startTime: getCurrentUnixTimestamp(), - } satisfies AnthropicToolRunnerState; - - states.set(event as object, state); - }, + INSTRUMENTATION_NAMES.ANTHROPIC, + ), + ); - end: (event) => { - const state = states.get(event as object); - if (!state) { - return; - } + span.log({ + input: processAttachmentsInInput( + coalesceInput(params.messages ?? [], params.system), + ), + metadata: { + ...extractAnthropicToolRunnerMetadata(params), + provider: "anthropic", + }, + }); + + const state = { + aggregatedMetrics: {}, + finalized: false, + iterationCount: 0, + seenMessages: new WeakSet(), + span, + startTime: getCurrentUnixTimestamp(), + } satisfies AnthropicToolRunnerState; + + states.set(event as object, state); + }; + const returned = ( + event: ChannelMessage< + typeof anthropicChannels.betaMessagesToolRunner + >, + ) => { + const state = states.get(event as object); + if (!state) { + return; + } - patchAnthropicToolRunner({ - runner: event.result as AnthropicToolRunner, - state, - }); - }, + patchAnthropicToolRunner({ + runner: event.result as unknown as AnthropicToolRunner, + state, + }); + }; + const failed = ( + event: ChannelMessage< + typeof anthropicChannels.betaMessagesToolRunner + >, + ) => { + const state = states.get(event as object); + if (!state || !event.error) { + return; + } - error: (event) => { - const state = states.get(event as object); - if (!state || !event.error) { - return; + finalizeAnthropicToolRunnerError(state, event.error); + states.delete(event as object); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } - - finalizeAnthropicToolRunnerError(state, event.error); - states.delete(event as object); + Object.assign(event, { result }); + runInstrumentation(() => returned(event)); + return result; }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToAnthropicSessionStreams(): void { @@ -275,37 +315,62 @@ export class AnthropicPlugin extends BasePlugin { type SessionChannelMessage = ChannelMessage< typeof anthropicChannels.betaSessionsEventsStream >; - const tracingChannel = - channel.tracingChannel() as IsoTracingChannel; + const invocationHook = channel; const pending = new WeakSet(); - const handlers: IsoChannelHandlers = { - start: (event) => { - if (isAutoInstrumentationSuppressed()) { - return; - } - pending.add(event as object); - }, - asyncEnd: (event) => { - if (!pending.delete(event as object)) { - return; - } + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as SessionChannelMessage; + + const prepare = (event: SessionChannelMessage) => { + if (isAutoInstrumentationSuppressed()) { + return; + } + pending.add(event as object); + }; + const resolved = (event: SessionChannelMessage) => { + if (!pending.delete(event as object)) { + return; + } - const stream = event.result as AnthropicSessionEventStream; - if (!isAsyncIterable(stream)) { - return; + const stream = event.result as AnthropicSessionEventStream; + if (!isAsyncIterable(stream)) { + return; + } + registerAnthropicSessionStreamCollector(stream, () => { + wrapAnthropicSessionEventStream(stream, isThread); + }); + }; + const failed = (event: SessionChannelMessage) => { + pending.delete(event as object); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } - registerAnthropicSessionStreamCollector(stream, () => { - wrapAnthropicSessionEventStream(stream, isThread); - }); - }, - error: (event) => { - pending.delete(event as object); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => tracingChannel.unsubscribe(handlers)); + ); + this.unsubscribers.push(removeHandlers); } } diff --git a/js/src/instrumentation/plugins/anthropic-sessions-plugin.test.ts b/js/src/instrumentation/plugins/anthropic-sessions-plugin.test.ts index 01cc32592..32ef8bdba 100644 --- a/js/src/instrumentation/plugins/anthropic-sessions-plugin.test.ts +++ b/js/src/instrumentation/plugins/anthropic-sessions-plugin.test.ts @@ -1,4 +1,12 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +import { invocationController } from "../test-utils/invocation"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const { mockStartSpan, mockWithCurrent } = vi.hoisted(() => ({ mockStartSpan: vi.fn(), @@ -12,7 +20,6 @@ vi.mock("../../isomorph", () => ({ getStore: vi.fn(() => undefined), run: vi.fn((_store: unknown, callback: () => unknown) => callback()), })), - newTracingChannel: vi.fn(), }, })); @@ -25,11 +32,12 @@ vi.mock("../../logger", async (importOriginal) => { }; }); -import iso from "../../isomorph"; import { collectAnthropicSession } from "../../wrappers/anthropic-session-collector"; import { AnthropicPlugin } from "./anthropic-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; describe("AnthropicPlugin Sessions instrumentation", () => { let currentSpan: TestSpan | undefined; @@ -40,9 +48,11 @@ describe("AnthropicPlugin Sessions instrumentation", () => { currentSpan = undefined; handlersByName = new Map(); spans = []; - mockNewTracingChannel.mockImplementation((name: string) => ({ - subscribe: vi.fn((handlers) => handlersByName.set(name, handlers)), - unsubscribe: vi.fn(), + mockNewInvocationHook.mockImplementation((name: string) => ({ + intercept: vi.fn((interceptor) => { + handlersByName.set(name, invocationController(interceptor)); + return vi.fn(); + }), })); mockWithCurrent.mockImplementation( (span: TestSpan, callback: () => unknown) => { @@ -153,8 +163,8 @@ describe("AnthropicPlugin Sessions instrumentation", () => { result?: ReturnType; } = { arguments: ["session-1"] }; - handlers.start(context); - handlers.asyncEnd(Object.assign(context, { result: stream })); + handlers.begin(context); + handlers.resolve(Object.assign(context, { result: stream })); expect(context.result).toBe(stream); expect(spans).toEqual([]); expect(collectAnthropicSession(stream)).toBe(stream); @@ -275,8 +285,8 @@ describe("AnthropicPlugin Sessions instrumentation", () => { ]); const context = { arguments: ["session-1"] }; - handlers.start(context); - handlers.asyncEnd(Object.assign(context, { result: stream })); + handlers.begin(context); + handlers.resolve(Object.assign(context, { result: stream })); for await (const _event of stream) { // Consume the stream without opting into collection. } @@ -314,8 +324,8 @@ describe("AnthropicPlugin Sessions instrumentation", () => { arguments: ["thread-1", { session_id: "session-1" }], }; - handlers.start(context); - handlers.asyncEnd(Object.assign(context, { result: stream })); + handlers.begin(context); + handlers.resolve(Object.assign(context, { result: stream })); expect(collectAnthropicSession(stream)).toBe(stream); for await (const _event of stream) { // Consume the stream. diff --git a/js/src/instrumentation/plugins/bedrock-runtime-channels.ts b/js/src/instrumentation/plugins/bedrock-runtime-channels.ts index 080b21a3a..53e75172a 100644 --- a/js/src/instrumentation/plugins/bedrock-runtime-channels.ts +++ b/js/src/instrumentation/plugins/bedrock-runtime-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { BedrockRuntimeChannelContext, BedrockRuntimeCommandLike, @@ -14,34 +14,21 @@ type BedrockRuntimeStreamEvent = const clientSendChannel = channel< [BedrockRuntimeCommandLike, unknown?], - BedrockRuntimeSendResult, + PromiseLike, BedrockRuntimeChannelContext, BedrockRuntimeStreamEvent >({ channelName: "client.send", - kind: "async", }); -export const bedrockRuntimeChannels = defineChannels( - "aws-bedrock-runtime", - { - clientSend: clientSendChannel, - }, - { instrumentationName: INSTRUMENTATION_NAMES.BEDROCK_RUNTIME }, -); +export const bedrockRuntimeChannels = defineInterceptor("aws-bedrock-runtime", { + clientSend: clientSendChannel, +}); -export const smithyCoreChannels = defineChannels( - "@smithy/core", - { - clientSend: clientSendChannel, - }, - { instrumentationName: INSTRUMENTATION_NAMES.BEDROCK_RUNTIME }, -); +export const smithyCoreChannels = defineInterceptor("@smithy/core", { + clientSend: clientSendChannel, +}); -export const smithyClientChannels = defineChannels( - "@smithy/smithy-client", - { - clientSend: clientSendChannel, - }, - { instrumentationName: INSTRUMENTATION_NAMES.BEDROCK_RUNTIME }, -); +export const smithyClientChannels = defineInterceptor("@smithy/smithy-client", { + clientSend: clientSendChannel, +}); diff --git a/js/src/instrumentation/plugins/bedrock-runtime-plugin.test.ts b/js/src/instrumentation/plugins/bedrock-runtime-plugin.test.ts index fc74a9068..81026acb1 100644 --- a/js/src/instrumentation/plugins/bedrock-runtime-plugin.test.ts +++ b/js/src/instrumentation/plugins/bedrock-runtime-plugin.test.ts @@ -51,9 +51,7 @@ describe("BedrockRuntimePlugin", () => { }); it("traces promise-style Smithy send events and ignores callback overloads", async () => { - const tracingChannel = smithyCoreChannels.clientSend.tracingChannel(); - - await smithyCoreChannels.clientSend.tracePromise( + await smithyCoreChannels.clientSend.invoke( async () => ({ output: { message: { @@ -67,14 +65,14 @@ describe("BedrockRuntimePlugin", () => { totalTokens: 2, }, }), - { - arguments: [ - new ConverseCommand({ - messages: [{ role: "user", content: [{ text: "OK" }] }], - modelId: "us.amazon.nova-lite-v1:0", - }) as any, - ], - }, + undefined, + [ + new ConverseCommand({ + messages: [{ role: "user", content: [{ text: "OK" }] }], + modelId: "us.amazon.nova-lite-v1:0", + }) as any, + ], + {}, ); const callbackEvent: any = { @@ -86,7 +84,7 @@ describe("BedrockRuntimePlugin", () => { () => {}, ], }; - tracingChannel.start!.publish(callbackEvent); + callbackEvent.result = { output: { message: { @@ -95,7 +93,12 @@ describe("BedrockRuntimePlugin", () => { }, }, }; - tracingChannel.asyncEnd!.publish(callbackEvent); + smithyCoreChannels.clientSend.invoke( + () => callbackEvent.result, + undefined, + callbackEvent.arguments, + {}, + ); const spans = await backgroundLogger.drain(); expect(spans).toEqual( @@ -139,19 +142,19 @@ describe("BedrockRuntimePlugin", () => { }; const originalIterator = stream[Symbol.asyncIterator]; - const result = await channel.tracePromise( + const result = await channel.invoke( async () => ({ body: stream, service: "s3", }), - { - arguments: [ - new GetObjectCommand({ - Bucket: "not-bedrock", - Key: "object.txt", - }) as any, - ], - }, + undefined, + [ + new GetObjectCommand({ + Bucket: "not-bedrock", + Key: "object.txt", + }) as any, + ], + {}, ); expect(result.body).toBe(stream); @@ -197,18 +200,18 @@ describe("BedrockRuntimePlugin", () => { }; } - const result = await smithyCoreChannels.clientSend.tracePromise( + const result = await smithyCoreChannels.clientSend.invoke( async () => ({ body: body(), }), - { - arguments: [ - new InvokeModelWithBidirectionalStreamCommand({ - body: undefined, - modelId: "us.amazon.nova-lite-v1:0", - }) as any, - ], - }, + undefined, + [ + new InvokeModelWithBidirectionalStreamCommand({ + body: undefined, + modelId: "us.amazon.nova-lite-v1:0", + }) as any, + ], + {}, ); for await (const _chunk of result.body) { diff --git a/js/src/instrumentation/plugins/bedrock-runtime-plugin.ts b/js/src/instrumentation/plugins/bedrock-runtime-plugin.ts index a759d1ffe..77893de50 100644 --- a/js/src/instrumentation/plugins/bedrock-runtime-plugin.ts +++ b/js/src/instrumentation/plugins/bedrock-runtime-plugin.ts @@ -1,10 +1,13 @@ -import { BasePlugin } from "../core"; -import { traceStreamingChannel, unsubscribeAll } from "../core/channel-tracing"; -import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; +import type { ReturnOf } from "../core/channel-definitions"; +import type { StartOf } from "../core/tracing-types"; import { SpanTypeAttribute, isObject } from "../../../util/index"; -import { getCurrentUnixTimestamp } from "../../util"; import type { Span } from "../../logger"; -import type { AnyAsyncChannel } from "../core/channel-definitions"; +import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { getCurrentUnixTimestamp } from "../../util"; +import { BasePlugin } from "../core"; +import { traceStreamingCall, unsubscribeAll } from "../core/channel-tracing"; +import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; + import type { BedrockRuntimeConverseRequest, BedrockRuntimeConverseResponse, @@ -32,7 +35,14 @@ export class BedrockRuntimePlugin extends BasePlugin { bedrockRuntimeChannels.clientSend, smithyCoreChannels.clientSend, smithyClientChannels.clientSend, - ].map((channel) => traceBedrockRuntimeClientSendChannel(channel)), + ].map((channel) => + channel.intercept((target, receiver, args, additional) => + traceBedrockRuntimeClientSend( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + ), + ), + ), ); } @@ -41,30 +51,36 @@ export class BedrockRuntimePlugin extends BasePlugin { } } -function traceBedrockRuntimeClientSendChannel( - channel: AnyAsyncChannel, -): () => void { - return traceStreamingChannel(channel, { - name: ([command]) => buildBedrockRuntimeSpanInfo(command).name, - shouldTrace: ([command, optionsOrCb, cb]) => - getBedrockRuntimeOperation(command) !== undefined && - typeof optionsOrCb !== "function" && - typeof cb !== "function", - type: SpanTypeAttribute.LLM, - extractInput: ([command]) => extractBedrockRuntimeInput(command), - extractOutput: (result, endEvent) => - extractBedrockRuntimeOutput(endEvent?.arguments?.[0], result), - extractMetadata: (result, endEvent) => - extractBedrockRuntimeResponseMetadata(endEvent?.arguments?.[0], result), - extractMetrics: (result) => extractBedrockRuntimeResponseMetrics(result), - patchResult: ({ endEvent, result, span, startTime }) => - patchBedrockRuntimeStreamingResult({ - command: endEvent.arguments?.[0], - result, - span, - startTime, - }), - }); +function traceBedrockRuntimeClientSend( + call: () => ReturnOf, + event: StartOf, +): ReturnOf { + return traceStreamingCall( + call, + event, + { + instrumentationName: INSTRUMENTATION_NAMES.BEDROCK_RUNTIME, + name: ([command]) => buildBedrockRuntimeSpanInfo(command).name, + shouldTrace: ([command, optionsOrCb, cb]) => + getBedrockRuntimeOperation(command) !== undefined && + typeof optionsOrCb !== "function" && + typeof cb !== "function", + type: SpanTypeAttribute.LLM, + extractInput: ([command]) => extractBedrockRuntimeInput(command), + extractOutput: (result, endEvent) => + extractBedrockRuntimeOutput(endEvent?.arguments?.[0], result), + extractMetadata: (result, endEvent) => + extractBedrockRuntimeResponseMetadata(endEvent?.arguments?.[0], result), + extractMetrics: (result) => extractBedrockRuntimeResponseMetrics(result), + patchResult: ({ endEvent, result, span, startTime }) => + patchBedrockRuntimeStreamingResult({ + command: endEvent.arguments?.[0], + result, + span, + startTime, + }), + }, + ); } function extractBedrockRuntimeInput(command: unknown): { diff --git a/js/src/instrumentation/plugins/claude-agent-sdk-channels.ts b/js/src/instrumentation/plugins/claude-agent-sdk-channels.ts index c58e23144..2532af2fc 100644 --- a/js/src/instrumentation/plugins/claude-agent-sdk-channels.ts +++ b/js/src/instrumentation/plugins/claude-agent-sdk-channels.ts @@ -1,11 +1,11 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { ClaudeAgentSDKMessage, ClaudeAgentSDKQueryParams, } from "../../vendor-sdk-types/claude-agent-sdk"; -export const claudeAgentSDKChannels = defineChannels( +export const claudeAgentSDKChannels = defineInterceptor( "@anthropic-ai/claude-agent-sdk", { query: channel< @@ -15,8 +15,6 @@ export const claudeAgentSDKChannels = defineChannels( ClaudeAgentSDKMessage >({ channelName: "query", - kind: "sync-stream", }), }, - { instrumentationName: INSTRUMENTATION_NAMES.CLAUDE_AGENT_SDK }, ); diff --git a/js/src/instrumentation/plugins/claude-agent-sdk-plugin.test.ts b/js/src/instrumentation/plugins/claude-agent-sdk-plugin.test.ts index 1c35a4a27..c0728dda1 100644 --- a/js/src/instrumentation/plugins/claude-agent-sdk-plugin.test.ts +++ b/js/src/instrumentation/plugins/claude-agent-sdk-plugin.test.ts @@ -1,7 +1,14 @@ -import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; import { AsyncLocalStorage } from "node:async_hooks"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); -// Mock iso's newTracingChannel - must be before any imports that use it +// Mock platform context independently of invocation hooks. const streamPatcherMock = vi.hoisted(() => ({ options: undefined as | { @@ -14,7 +21,6 @@ const streamPatcherMock = vi.hoisted(() => ({ vi.mock("../../isomorph", () => ({ default: { newAsyncLocalStorage: () => new AsyncLocalStorage(), - newTracingChannel: vi.fn(), }, })); @@ -32,11 +38,12 @@ vi.mock("../core/stream-patcher", () => ({ }), })); -import { ClaudeAgentSDKPlugin } from "./claude-agent-sdk-plugin"; -import iso from "../../isomorph"; import { startSpan } from "../../logger"; +import { ClaudeAgentSDKPlugin } from "./claude-agent-sdk-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; // Mock the logger module vi.mock("../../logger", () => ({ @@ -140,7 +147,7 @@ describe("ClaudeAgentSDKPlugin", () => { hasSubscribers: false, }; - mockNewTracingChannel.mockReturnValue(mockChannel); + mockNewInvocationHook.mockReturnValue(mockChannel); plugin = new ClaudeAgentSDKPlugin(); }); @@ -153,7 +160,7 @@ describe("ClaudeAgentSDKPlugin", () => { it("should enable the plugin and subscribe to channels", () => { plugin.enable(); - expect(mockNewTracingChannel).toHaveBeenCalledWith( + expect(mockNewInvocationHook).toHaveBeenCalledWith( "orchestrion:@anthropic-ai/claude-agent-sdk:query", ); expect(mockChannel.intercept).toHaveBeenCalledTimes(1); diff --git a/js/src/instrumentation/plugins/cloudflare-agents-channels.ts b/js/src/instrumentation/plugins/cloudflare-agents-channels.ts index 2f12705e1..5f6aa6d85 100644 --- a/js/src/instrumentation/plugins/cloudflare-agents-channels.ts +++ b/js/src/instrumentation/plugins/cloudflare-agents-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { CloudflareAgentToolClass, CloudflareRunAgentToolOptions, @@ -10,17 +10,12 @@ type CloudflareAgentsChannelContext = { self?: unknown; }; -export const cloudflareAgentsChannels = defineChannels( - "agents", - { - runAgentTool: channel< - [CloudflareAgentToolClass, CloudflareRunAgentToolOptions], - CloudflareRunAgentToolResult, - CloudflareAgentsChannelContext - >({ - channelName: "Agent.runAgentTool", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.CLOUDFLARE_AGENTS }, -); +export const cloudflareAgentsChannels = defineInterceptor("agents", { + runAgentTool: channel< + [CloudflareAgentToolClass, CloudflareRunAgentToolOptions], + PromiseLike, + CloudflareAgentsChannelContext + >({ + channelName: "Agent.runAgentTool", + }), +}); diff --git a/js/src/instrumentation/plugins/cloudflare-agents-plugin.test.ts b/js/src/instrumentation/plugins/cloudflare-agents-plugin.test.ts index d6d2de7de..8f98f079e 100644 --- a/js/src/instrumentation/plugins/cloudflare-agents-plugin.test.ts +++ b/js/src/instrumentation/plugins/cloudflare-agents-plugin.test.ts @@ -7,11 +7,19 @@ import { it, vi, } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; import type { StartSpanArgs } from "../../logger"; import { - getSpanInstrumentationName, INSTRUMENTATION_NAMES, + getSpanInstrumentationName, } from "../../span-origin"; +import { invocationController } from "../test-utils/invocation"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const { mockStartSpan } = vi.hoisted(() => ({ mockStartSpan: vi.fn(), @@ -24,18 +32,18 @@ vi.mock("../../logger", () => ({ vi.mock("../../isomorph", () => ({ default: { getEnv: vi.fn(), - newTracingChannel: vi.fn(), }, })); -import iso from "../../isomorph"; import { CloudflareAgentsPlugin } from "./cloudflare-agents-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; describe("CloudflareAgentsPlugin", () => { let handlers: any; - let subscribe: ReturnType; + let subscribe: ReturnType void>>; let unsubscribe: ReturnType; let spans: Array<{ args: any; @@ -50,7 +58,12 @@ describe("CloudflareAgentsPlugin", () => { handlers = nextHandlers; }); unsubscribe = vi.fn(); - mockNewTracingChannel.mockReturnValue({ subscribe, unsubscribe }); + mockNewInvocationHook.mockReturnValue({ + intercept: (interceptor: any) => { + subscribe(invocationController(interceptor)); + return unsubscribe; + }, + }); mockStartSpan.mockImplementation((args: any, context: any) => { const span = { args, context, end: vi.fn(), log: vi.fn() }; spans.push(span); @@ -67,7 +80,7 @@ describe("CloudflareAgentsPlugin", () => { plugin.enable(); plugin.enable(); - expect(mockNewTracingChannel).toHaveBeenCalledWith( + expect(mockNewInvocationHook).toHaveBeenCalledWith( "orchestrion:agents:Agent.runAgentTool", ); expect(subscribe).toHaveBeenCalledTimes(1); @@ -100,8 +113,8 @@ describe("CloudflareAgentsPlugin", () => { ], }; - handlers.start(event); - handlers.asyncEnd( + handlers.begin(event); + handlers.resolve( Object.assign(event, { result: { status: "completed", @@ -146,8 +159,8 @@ describe("CloudflareAgentsPlugin", () => { class FailingAgent {} const event = { arguments: [FailingAgent, { input: "fail" }] }; - handlers.start(event); - handlers.asyncEnd( + handlers.begin(event); + handlers.resolve( Object.assign(event, { result: { status: "error", @@ -171,10 +184,10 @@ describe("CloudflareAgentsPlugin", () => { const second = { arguments: [SecondAgent, { input: 2 }] }; const rejection = new Error("rejected"); - handlers.start(first); - handlers.start(second); - handlers.error(Object.assign(second, { error: rejection })); - handlers.asyncEnd( + handlers.begin(first); + handlers.begin(second); + handlers.reject(Object.assign(second, { error: rejection })); + handlers.resolve( Object.assign(first, { result: { status: "completed", output: "first" }, }), @@ -202,7 +215,7 @@ describe("CloudflareAgentsPlugin", () => { }, ); - handlers.start({ arguments: [AgentWithGetter, options] }); + handlers.begin({ arguments: [AgentWithGetter, options] }); expect(spans).toHaveLength(0); expect(nameGetter).not.toHaveBeenCalled(); diff --git a/js/src/instrumentation/plugins/cloudflare-agents-plugin.ts b/js/src/instrumentation/plugins/cloudflare-agents-plugin.ts index d9813f419..ffe41e37b 100644 --- a/js/src/instrumentation/plugins/cloudflare-agents-plugin.ts +++ b/js/src/instrumentation/plugins/cloudflare-agents-plugin.ts @@ -1,14 +1,15 @@ +import { SpanTypeAttribute } from "../../../util/index"; import { debugLogger } from "../../debug-logger"; -import type { IsoChannelHandlers } from "../../isomorph"; -import { _internalStartSpanWithContext } from "../../logger"; import type { Span } from "../../logger"; +import { _internalStartSpanWithContext } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { SpanTypeAttribute } from "../../../util/index"; import { BasePlugin } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + +import type { ChannelMessage } from "../core/tracing-types"; import { cloudflareAgentsChannels } from "./cloudflare-agents-channels"; const CLOUDFLARE_WORKERS_CONTEXT = { @@ -19,87 +20,117 @@ const CLOUDFLARE_WORKERS_CONTEXT = { export class CloudflareAgentsPlugin extends BasePlugin { protected onEnable(): void { - const channel = cloudflareAgentsChannels.runAgentTool.tracingChannel(); + const channel = cloudflareAgentsChannels.runAgentTool; const spans = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - try { - const agentClass = event.arguments[0]; - const options = event.arguments[1]; - if (ownValue(options, "detached")) { - return; - } - const name = ownValue(agentClass, "name"); - if (typeof name !== "string" || name.length === 0) { - debugLogger.warn( - "Skipping Cloudflare Agents runAgentTool span because the child agent class has no name.", + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + try { + const agentClass = event.arguments[0]; + const options = event.arguments[1]; + if (ownValue(options, "detached")) { + return; + } + + const name = ownValue(agentClass, "name"); + if (typeof name !== "string" || name.length === 0) { + debugLogger.warn( + "Skipping Cloudflare Agents runAgentTool span because the child agent class has no name.", + ); + return; + } + + const span = _internalStartSpanWithContext( + withSpanInstrumentationName( + { + name, + spanAttributes: { type: SpanTypeAttribute.TOOL }, + event: { + input: ownValue(options, "input"), + }, + }, + INSTRUMENTATION_NAMES.CLOUDFLARE_AGENTS, + ), + CLOUDFLARE_WORKERS_CONTEXT, ); + spans.set(event, span); + } catch (error) { + logInstrumentationError("start", error); + } + }; + const resolved = ( + event: ChannelMessage, + ) => { + const span = spans.get(event); + if (!span) { return; } + spans.delete(event); - const span = _internalStartSpanWithContext( - withSpanInstrumentationName( - { - name, - spanAttributes: { type: SpanTypeAttribute.TOOL }, - event: { - input: ownValue(options, "input"), - }, - }, - INSTRUMENTATION_NAMES.CLOUDFLARE_AGENTS, - ), - CLOUDFLARE_WORKERS_CONTEXT, - ); - spans.set(event, span); - } catch (error) { - logInstrumentationError("start", error); - } - }, - asyncEnd: (event) => { - const span = spans.get(event); - if (!span) { - return; - } - spans.delete(event); - - try { - const status = ownValue(event.result, "status"); - if (status === "completed") { - span.log({ output: ownValue(event.result, "output") }); - } else { - const error = ownValue(event.result, "error"); - if (typeof error === "string") { - span.log({ error }); + try { + const status = ownValue(event.result, "status"); + if (status === "completed") { + span.log({ output: ownValue(event.result, "output") }); + } else { + const error = ownValue(event.result, "error"); + if (typeof error === "string") { + span.log({ error }); + } } + } catch (error) { + logInstrumentationError("completion", error); + } finally { + safelyEndSpan(span); } - } catch (error) { - logInstrumentationError("completion", error); - } finally { - safelyEndSpan(span); - } - }, - error: (event) => { - const span = spans.get(event); - if (!span) { - return; - } - spans.delete(event); + }; + const failed = ( + event: ChannelMessage, + ) => { + const span = spans.get(event); + if (!span) { + return; + } + spans.delete(event); + try { + span.log({ error: event.error }); + } catch (error) { + logInstrumentationError("rejection", error); + } finally { + safelyEndSpan(span); + } + }; + runInstrumentation(() => prepare(event)); + let result; try { - span.log({ error: event.error }); + result = Reflect.apply(target, receiver, args); } catch (error) { - logInstrumentationError("rejection", error); - } finally { - safelyEndSpan(span); + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); }, - }; - - channel.subscribe(handlers); - this.unsubscribers.push(() => channel.unsubscribe(handlers)); + ); + this.unsubscribers.push(removeHandlers); } protected onDisable(): void { diff --git a/js/src/instrumentation/plugins/cloudflare-ai-chat-channels.ts b/js/src/instrumentation/plugins/cloudflare-ai-chat-channels.ts index c23871dbb..33a15aad9 100644 --- a/js/src/instrumentation/plugins/cloudflare-ai-chat-channels.ts +++ b/js/src/instrumentation/plugins/cloudflare-ai-chat-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { CloudflareAIChatResponseResult, CloudflareAIChatTurnCallback, @@ -10,16 +10,15 @@ type CloudflareAIChatChannelContext = { self?: unknown; }; -export const cloudflareAIChatChannels = defineChannels( +export const cloudflareAIChatChannels = defineInterceptor( "@cloudflare/ai-chat", { runExclusiveChatTurn: channel< [string, CloudflareAIChatTurnCallback, CloudflareAIChatTurnOptions?], - unknown, + PromiseLike, CloudflareAIChatChannelContext >({ channelName: "AIChatAgent._runExclusiveChatTurn", - kind: "async", }), onChatResponse: channel< @@ -28,8 +27,6 @@ export const cloudflareAIChatChannels = defineChannels( CloudflareAIChatChannelContext >({ channelName: "AIChatAgent.onChatResponse", - kind: "sync-stream", }), }, - { instrumentationName: INSTRUMENTATION_NAMES.CLOUDFLARE_AI_CHAT }, ); diff --git a/js/src/instrumentation/plugins/cloudflare-ai-chat-instrumentation.ts b/js/src/instrumentation/plugins/cloudflare-ai-chat-instrumentation.ts index 0a109f3e6..c0e775c7e 100644 --- a/js/src/instrumentation/plugins/cloudflare-ai-chat-instrumentation.ts +++ b/js/src/instrumentation/plugins/cloudflare-ai-chat-instrumentation.ts @@ -21,12 +21,11 @@ export function instrumentCloudflareAIChatAgent( this: CloudflareAIChatAgent, ...args: Parameters ): Promise { - return cloudflareAIChatChannels.runExclusiveChatTurn.tracePromise( - () => Reflect.apply(original, this, args), - { - arguments: args, - self: this, - }, + return cloudflareAIChatChannels.runExclusiveChatTurn.invoke( + original, + this, + args, + {}, ); }; wrappedTurnRunners.add(wrapped); @@ -52,12 +51,11 @@ export function instrumentCloudflareAIChatResponseHook( result: CloudflareAIChatResponseResult, ): unknown { const args: [CloudflareAIChatResponseResult] = [result]; - return cloudflareAIChatChannels.onChatResponse.traceSync( - () => Reflect.apply(original, this, args), - { - arguments: args, - self: this, - }, + return cloudflareAIChatChannels.onChatResponse.invoke( + original, + this, + args, + {}, ); }; wrappedResponseHooks.add(wrapped); diff --git a/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.test.ts b/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.test.ts index d20631a2d..f1d869037 100644 --- a/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.test.ts +++ b/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.test.ts @@ -1,4 +1,12 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +import { invocationController } from "../test-utils/invocation"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const { mockInternalGetGlobalState, @@ -15,7 +23,7 @@ const { })); vi.mock("../../isomorph", () => ({ - default: { newTracingChannel: vi.fn() }, + default: {}, })); vi.mock("../../logger", () => ({ @@ -25,14 +33,15 @@ vi.mock("../../logger", () => ({ withCurrent: (...args: unknown[]) => (mockWithCurrent as any)(...args), })); -import iso from "../../isomorph"; import { INSTRUMENTATION_NAMES, INTERNAL_SPAN_INSTRUMENTATION_NAME, } from "../../span-origin"; import { CloudflareAIChatPlugin } from "./cloudflare-ai-chat-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; describe("CloudflareAIChatPlugin", () => { let plugin: CloudflareAIChatPlugin; @@ -40,7 +49,7 @@ describe("CloudflareAIChatPlugin", () => { beforeEach(() => { channels = new Map(); - mockNewTracingChannel.mockImplementation((name: string) => { + mockNewInvocationHook.mockImplementation((name: string) => { const existing = channels.get(name); if (existing) { return existing; @@ -82,7 +91,7 @@ describe("CloudflareAIChatPlugin", () => { self: agent, } as any; - turnHandlers.start?.(event, "start"); + turnHandlers.begin?.(event); await event.arguments[1](); agent.onChatResponse({ message: { @@ -94,7 +103,7 @@ describe("CloudflareAIChatPlugin", () => { requestId: "request-1", status: "completed", }); - turnHandlers.asyncEnd?.(event, "asyncEnd"); + turnHandlers.resolve?.(event); const span = mockStartSpan.mock.results[0].value; expect(mockStartSpan).toHaveBeenCalledWith({ @@ -140,24 +149,21 @@ describe("CloudflareAIChatPlugin", () => { arguments: ["request-error", async () => undefined, undefined], self: agent, } as any; - turnHandlers.start?.(event, "start"); + turnHandlers.begin?.(event); await event.arguments[1](); - responseHandlers.start?.( - { - arguments: [ - { - error: "stream failed", - message: { parts: [{ text: "partial" }], role: "assistant" }, - requestId: "request-error", - status: "error", - }, - ], - self: agent, - } as any, - "start", - ); - turnHandlers.asyncEnd?.(event, "asyncEnd"); + responseHandlers.begin?.({ + arguments: [ + { + error: "stream failed", + message: { parts: [{ text: "partial" }], role: "assistant" }, + requestId: "request-error", + status: "error", + }, + ], + self: agent, + } as any); + turnHandlers.resolve?.(event); const span = mockStartSpan.mock.results[0].value; expect(span.log).toHaveBeenCalledWith({ @@ -186,9 +192,9 @@ describe("CloudflareAIChatPlugin", () => { self: agent, } as any; - handlers.start?.(event, "start"); + handlers.begin?.(event); await event.arguments[1](); - handlers.asyncEnd?.(event, "asyncEnd"); + handlers.resolve?.(event); const span = mockStartSpan.mock.results[0].value; expect(span.end).toHaveBeenCalledTimes(1); @@ -240,8 +246,8 @@ describe("CloudflareAIChatPlugin", () => { self: agent, } as any; - handlers.start?.(event, "start"); - handlers.asyncEnd?.(event, "asyncEnd"); + handlers.begin?.(event); + handlers.resolve?.(event); const span = mockStartSpan.mock.results[0].value; expect(span.end).toHaveBeenCalledTimes(1); @@ -285,7 +291,7 @@ describe("CloudflareAIChatPlugin", () => { self: agent, } as any; - handlers.start?.(event, "start"); + handlers.begin?.(event); await event.arguments[1](); agent.messages[1] = { id: "assistant-1", @@ -298,7 +304,7 @@ describe("CloudflareAIChatPlugin", () => { requestId: "request-continuation", status: "completed", }); - handlers.asyncEnd?.(event, "asyncEnd"); + handlers.resolve?.(event); const span = mockStartSpan.mock.results[0].value; expect(span.log).toHaveBeenCalledWith({ @@ -336,15 +342,15 @@ describe("CloudflareAIChatPlugin", () => { self: agent, } as any; - handlers.start?.(outer, "start"); - handlers.start?.(inner, "start"); - handlers.asyncEnd?.(inner, "asyncEnd"); + handlers.begin?.(outer); + handlers.begin?.(inner); + handlers.resolve?.(inner); const span = mockStartSpan.mock.results[0].value; expect(mockStartSpan).toHaveBeenCalledTimes(1); expect(span.end).not.toHaveBeenCalled(); - handlers.asyncEnd?.(outer, "asyncEnd"); + handlers.resolve?.(outer); expect(span.end).toHaveBeenCalledTimes(1); }); @@ -356,9 +362,9 @@ describe("CloudflareAIChatPlugin", () => { arguments: ["request-1", async () => undefined, undefined], self: { messages: [], onChatResponse() {} }, } as any; - handlers.start?.(failedEvent, "start"); + handlers.begin?.(failedEvent); failedEvent.error = failure; - handlers.error?.(failedEvent, "error"); + handlers.reject?.(failedEvent); const failedSpan = mockStartSpan.mock.results[0].value; expect(failedSpan.log).toHaveBeenCalledWith({ error: failure }); @@ -368,7 +374,7 @@ describe("CloudflareAIChatPlugin", () => { arguments: ["request-2", async () => undefined, undefined], self: { messages: [], onChatResponse() {} }, } as any; - handlers.start?.(pendingEvent, "start"); + handlers.begin?.(pendingEvent); const pendingSpan = mockStartSpan.mock.results[1].value; plugin.disable(); expect(pendingSpan.end).toHaveBeenCalledTimes(1); @@ -388,28 +394,20 @@ describe("CloudflareAIChatPlugin", () => { }); function createMockChannel() { - const subscribed: any[] = []; + let interceptor: any; + let controller: ReturnType; return { - handlers: () => subscribed[0], - hasSubscribers: false, - start: { - bindStore: vi.fn(), - unbindStore: vi.fn(), - }, - subscribe: vi.fn((handlers) => subscribed.push(handlers)), - traceSync: vi.fn((callback, event) => { - subscribed[0]?.start?.(event, "start"); - try { - const result = callback(); - event.result = result; - subscribed[0]?.end?.(event, "end"); - return result; - } catch (error) { - event.error = error; - subscribed[0]?.error?.(event, "error"); - throw error; - } + handlers: () => controller, + intercept: vi.fn((next) => { + interceptor = next; + controller = invocationController(next); + return vi.fn(); }), - unsubscribe: vi.fn(), + invoke: ( + target: any, + receiver: unknown, + args: unknown[], + additional: object, + ) => interceptor(target, receiver, args, additional), }; } diff --git a/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.ts b/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.ts index d875824af..f658d70dd 100644 --- a/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.ts +++ b/js/src/instrumentation/plugins/cloudflare-ai-chat-plugin.ts @@ -1,24 +1,20 @@ import { BasePlugin } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; -import { - BRAINTRUST_CURRENT_SPAN_STORE, - _internalGetGlobalState, - startSpan as startBaseSpan, - withCurrent, -} from "../../logger"; -import type { CurrentSpanStore, Span } from "../../logger"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute } from "../../../util/index"; import { debugLogger } from "../../debug-logger"; +import type { Span } from "../../logger"; +import { startSpan as startBaseSpan, withCurrent } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { SpanTypeAttribute } from "../../../util/index"; import type { CloudflareAIChatAgent, CloudflareAIChatMessage, CloudflareAIChatTurnCallback, } from "../../vendor-sdk-types/cloudflare-ai-chat"; +import type { ChannelMessage } from "../core/tracing-types"; import { cloudflareAIChatChannels } from "./cloudflare-ai-chat-channels"; import { instrumentCloudflareAIChatResponseHook } from "./cloudflare-ai-chat-instrumentation"; @@ -67,111 +63,121 @@ export class CloudflareAIChatPlugin extends BasePlugin { } private subscribeToTurnRunner(): void { - const tracingChannel = - cloudflareAIChatChannels.runExclusiveChatTurn.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; - - const unbindCurrentSpanStore = this.bindCurrentSpanStore(tracingChannel); - const handlers: IsoChannelHandlers> = { - start: (event) => { - this.ensureEventState(event); - }, - asyncEnd: (event) => { - this.finishEvent(event); - }, - error: (event) => { - this.finishEvent(event, event.error); - }, - }; + const invocationHook = cloudflareAIChatChannels.runExclusiveChatTurn; + + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const spanState = runInstrumentation(() => + this.ensureEventState(event), + ); + const prepare = (event: ChannelMessage) => { + this.ensureEventState(event); + }; + const resolved = (event: ChannelMessage) => { + this.finishEvent(event); + }; + const failed = (event: ChannelMessage) => { + this.finishEvent(event, event.error); + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); + }; + return spanState ? withCurrent(spanState.span, invoke) : invoke(); + }, + ); + this.unsubscribers.push(removeHandlers); } private subscribeToResponseHook(): void { - const tracingChannel = - cloudflareAIChatChannels.onChatResponse.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; - const handlers: IsoChannelHandlers> = { - start: (event) => { - let state: TurnState | undefined; - try { - const agent = asObject(event.self); - const result = event.arguments[0]; - state = agent - ? this.findResponseState( - agent, - stringValue(ownValue(result, "requestId")), - ) - : undefined; - if (!state) { - return; + const invocationHook = cloudflareAIChatChannels.onChatResponse; + + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + + const prepare = (event: ChannelMessage) => { + let state: TurnState | undefined; + try { + const agent = asObject(event.self); + const result = event.arguments[0]; + state = agent + ? this.findResponseState( + agent, + stringValue(ownValue(result, "requestId")), + ) + : undefined; + if (!state) { + return; + } + state.responseObserved = true; + + const output = serializeMessage(ownValue(result, "message")); + const status = stringValue(ownValue(result, "status")); + const error = ownValue(result, "error"); + const input = + ownValue(result, "continuation") === true + ? state.input + : serializeMessages(readProperty(agent, "messages"))?.filter( + (message) => + typeof output?.id !== "string" || + message.id !== output.id, + ); + state.span.log({ + ...(input !== undefined ? { input } : {}), + ...(output !== undefined ? { output } : {}), + ...(status === "error" && error !== undefined ? { error } : {}), + }); + } catch (error) { + debugLogger.debug( + "Failed to process @cloudflare/ai-chat response hook:", + error, + ); + } finally { + if (state?.settled) { + this.cleanupState(state); + } } - state.responseObserved = true; - - const output = serializeMessage(ownValue(result, "message")); - const status = stringValue(ownValue(result, "status")); - const error = ownValue(result, "error"); - const input = - ownValue(result, "continuation") === true - ? state.input - : serializeMessages(readProperty(agent, "messages"))?.filter( - (message) => - typeof output?.id !== "string" || message.id !== output.id, - ); - state.span.log({ - ...(input !== undefined ? { input } : {}), - ...(output !== undefined ? { output } : {}), - ...(status === "error" && error !== undefined ? { error } : {}), - }); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); } catch (error) { - debugLogger.debug( - "Failed to process @cloudflare/ai-chat response hook:", - error, - ); - } finally { - if (state?.settled) { - this.cleanupState(state); - } + throw error; } + return result; }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => tracingChannel.unsubscribe(handlers)); - } - - private bindCurrentSpanStore( - tracingChannel: IsoTracingChannel>, - ): (() => void) | undefined { - const globalState = _internalGetGlobalState(); - const contextManager = globalState?.contextManager; - const startChannel = tracingChannel.start; - const currentSpanStore = contextManager - ? ( - contextManager as { - [BRAINTRUST_CURRENT_SPAN_STORE]?: CurrentSpanStore; - } - )[BRAINTRUST_CURRENT_SPAN_STORE] - : undefined; - - if (!startChannel || !currentSpanStore || !contextManager) { - return undefined; - } - - startChannel.bindStore(currentSpanStore, (event) => { - const state = this.ensureEventState(event); - return state - ? contextManager.wrapSpanForStore(state.span) - : currentSpanStore.getStore(); - }); - - return () => startChannel.unbindStore(currentSpanStore); + ); + this.unsubscribers.push(removeHandlers); } private ensureEventState( diff --git a/js/src/instrumentation/plugins/cloudflare-think-channels.ts b/js/src/instrumentation/plugins/cloudflare-think-channels.ts index e7b50113c..c8728018c 100644 --- a/js/src/instrumentation/plugins/cloudflare-think-channels.ts +++ b/js/src/instrumentation/plugins/cloudflare-think-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { CloudflareThinkInstance, CloudflareThinkStreamableResult, @@ -11,17 +11,12 @@ type CloudflareThinkChannelContext = { moduleVersion?: string; }; -export const cloudflareThinkChannels = defineChannels( - "@cloudflare/think", - { - runInferenceLoop: channel< - [CloudflareThinkTurnInput], - CloudflareThinkStreamableResult, - CloudflareThinkChannelContext - >({ - channelName: "Think.runInferenceLoop", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.CLOUDFLARE_THINK }, -); +export const cloudflareThinkChannels = defineInterceptor("@cloudflare/think", { + runInferenceLoop: channel< + [CloudflareThinkTurnInput], + PromiseLike, + CloudflareThinkChannelContext + >({ + channelName: "Think.runInferenceLoop", + }), +}); diff --git a/js/src/instrumentation/plugins/cloudflare-think-plugin.ts b/js/src/instrumentation/plugins/cloudflare-think-plugin.ts index 8a29eb16c..b7416cd05 100644 --- a/js/src/instrumentation/plugins/cloudflare-think-plugin.ts +++ b/js/src/instrumentation/plugins/cloudflare-think-plugin.ts @@ -1,24 +1,24 @@ +import { withCurrent } from "../../logger"; import { BasePlugin } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers } from "../../isomorph"; -import { - BRAINTRUST_CURRENT_SPAN_STORE, - _internalGetGlobalState, - startSpan, -} from "../../logger"; -import type { CurrentSpanStore, Span } from "../../logger"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute, isObject } from "../../../util/index"; +import { debugLogger } from "../../debug-logger"; +import type { Span } from "../../logger"; +import { _internalGetGlobalState, startSpan } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { debugLogger } from "../../debug-logger"; import { getCurrentUnixTimestamp } from "../../util"; -import { SpanTypeAttribute, isObject } from "../../../util/index"; import { isAutoInstrumentationSuppressed } from "../auto-instrumentation-suppression"; +import type { ChannelMessage } from "../core/tracing-types"; // Think delegates inference and tool execution to AI SDK's streamText. Its // events keep the task open through stream consumption and provide the model, // tool, output, and usage data that the outer Think call does not expose. +import type { AISDKResult } from "../../vendor-sdk-types/ai-sdk"; +import type { CloudflareThinkMessage } from "../../vendor-sdk-types/cloudflare-think"; import { aiSDKChannels } from "./ai-sdk-channels"; import { DEFAULT_DENY_OUTPUT_PATHS, @@ -31,8 +31,6 @@ import { registerCloudflareThinkSpan, unregisterCloudflareThinkSpan, } from "./cloudflare-think-context"; -import type { AISDKResult } from "../../vendor-sdk-types/ai-sdk"; -import type { CloudflareThinkMessage } from "../../vendor-sdk-types/cloudflare-think"; type ThinkRunState = { aiEvent?: Record; @@ -72,18 +70,8 @@ export class CloudflareThinkPlugin extends BasePlugin { } private subscribeToThinkRuns(): void { - const channel = cloudflareThinkChannels.runInferenceLoop.tracingChannel(); + const channel = cloudflareThinkChannels.runInferenceLoop; const states = new WeakMap(); - const state = _internalGetGlobalState(); - const contextManager = state?.contextManager; - const currentSpanStore = contextManager - ? ( - contextManager as { - [BRAINTRUST_CURRENT_SPAN_STORE]?: CurrentSpanStore; - } - )[BRAINTRUST_CURRENT_SPAN_STORE] - : undefined; - const ensureState = ( event: ChannelMessage, ): ThinkRunState | undefined => { @@ -130,83 +118,163 @@ export class CloudflareThinkPlugin extends BasePlugin { return runState; }; - if (contextManager && currentSpanStore && channel.start) { - channel.start.bindStore(currentSpanStore, (event) => { - const runState = ensureState(event); - return runState - ? contextManager.wrapSpanForStore(runState.span) - : currentSpanStore.getStore(); - }); - this.unsubscribers.push(() => - channel.start?.unbindStore(currentSpanStore), - ); - } + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage< + typeof cloudflareThinkChannels.runInferenceLoop + >, + ) => { + ensureState(event); + }; + const resolved = ( + event: ChannelMessage< + typeof cloudflareThinkChannels.runInferenceLoop + >, + ) => { + const runState = states.get(event); + states.delete(event); + if (!runState || runState.finalized || runState.aiResultPatched) { + return; + } + this.finishState(runState, undefined, event.result); + }; + const failed = ( + event: ChannelMessage< + typeof cloudflareThinkChannels.runInferenceLoop + >, + ) => { + const runState = states.get(event); + states.delete(event); + if (runState) { + this.finishState(runState, event.error); + } + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - ensureState(event); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); + }; + const spanState = runInstrumentation(() => ensureState(event)); + return spanState ? withCurrent(spanState.span, invoke) : invoke(); }, - asyncEnd: (event) => { - const runState = states.get(event); - states.delete(event); - if (!runState || runState.finalized || runState.aiResultPatched) { - return; - } - this.finishState(runState, undefined, event.result); - }, - error: (event) => { - const runState = states.get(event); - states.delete(event); - if (runState) { - this.finishState(runState, event.error); - } - }, - }; - - channel.subscribe(handlers); - this.unsubscribers.push(() => channel.unsubscribe(handlers)); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToAISDKStreamTextSync(): void { - const channel = aiSDKChannels.streamTextSync.tracingChannel(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - this.startAISDKStream(event); - }, - end: (event) => { - this.endAISDKStream(event); - }, - error: (event) => { - this.errorAISDKStream(event); - }, - }; + const channel = aiSDKChannels.streamTextSync; - channel.subscribe(handlers); - this.unsubscribers.push(() => channel.unsubscribe(handlers)); + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + this.startAISDKStream(event); + }; + const returned = ( + event: ChannelMessage, + ) => { + this.endAISDKStream(event); + }; + const failed = ( + event: ChannelMessage, + ) => { + this.errorAISDKStream(event); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } + Object.assign(event, { result }); + runInstrumentation(() => returned(event)); + return result; + }, + ); + this.unsubscribers.push(removeHandlers); } private subscribeToAISDKStreamTextAsync(): void { - const channel = aiSDKChannels.streamText.tracingChannel(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - this.startAISDKStream(event); - }, - asyncEnd: (event) => { - this.endAISDKStream(event); - }, - error: (event) => { - this.errorAISDKStream(event); - }, - }; + const channel = aiSDKChannels.streamText; - channel.subscribe(handlers); - this.unsubscribers.push(() => channel.unsubscribe(handlers)); + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + this.startAISDKStream(event); + }; + const resolved = ( + event: ChannelMessage, + ) => { + this.endAISDKStream(event); + }; + const failed = ( + event: ChannelMessage, + ) => { + this.errorAISDKStream(event); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); + }, + ); + this.unsubscribers.push(removeHandlers); } private startAISDKStream(event: AISDKStreamEvent): void { diff --git a/js/src/instrumentation/plugins/cohere-channels.ts b/js/src/instrumentation/plugins/cohere-channels.ts index 54a145f5b..978284c4a 100644 --- a/js/src/instrumentation/plugins/cohere-channels.ts +++ b/js/src/instrumentation/plugins/cohere-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { CohereChatRequest, CohereChatResponse, @@ -11,33 +11,25 @@ import type { CohereRerankResponse, } from "../../vendor-sdk-types/cohere"; -export const cohereChannels = defineChannels( - "cohere-ai", - { - chat: channel<[CohereChatRequest], CohereChatResponse>({ - channelName: "chat", - kind: "async", - }), +export const cohereChannels = defineInterceptor("cohere-ai", { + chat: channel<[CohereChatRequest], PromiseLike>({ + channelName: "chat", + }), - chatStream: channel< - [CohereChatRequest], - CohereChatStreamResult, - Record, - CohereChatStreamEvent - >({ - channelName: "chatStream", - kind: "async", - }), + chatStream: channel< + [CohereChatRequest], + PromiseLike, + Record, + CohereChatStreamEvent + >({ + channelName: "chatStream", + }), - embed: channel<[CohereEmbedRequest], CohereEmbedResponse>({ - channelName: "embed", - kind: "async", - }), + embed: channel<[CohereEmbedRequest], PromiseLike>({ + channelName: "embed", + }), - rerank: channel<[CohereRerankRequest], CohereRerankResponse>({ - channelName: "rerank", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.COHERE }, -); + rerank: channel<[CohereRerankRequest], PromiseLike>({ + channelName: "rerank", + }), +}); diff --git a/js/src/instrumentation/plugins/cohere-plugin.ts b/js/src/instrumentation/plugins/cohere-plugin.ts index afaea628c..61c6d93d4 100644 --- a/js/src/instrumentation/plugins/cohere-plugin.ts +++ b/js/src/instrumentation/plugins/cohere-plugin.ts @@ -1,13 +1,6 @@ -import { BasePlugin } from "../core"; -import { - traceAsyncChannel, - traceStreamingChannel, - unsubscribeAll, -} from "../core/channel-tracing"; import { SpanTypeAttribute, isObject } from "../../../util/index"; -import { processInputAttachments } from "../../wrappers/attachment-utils"; +import { INSTRUMENTATION_NAMES } from "../../span-origin"; import { getCurrentUnixTimestamp } from "../../util"; -import { cohereChannels } from "./cohere-channels"; import type { CohereChatResponse, CohereChatStreamEvent, @@ -15,6 +8,14 @@ import type { CohereToolCall, CohereUsageLike, } from "../../vendor-sdk-types/cohere"; +import { processInputAttachments } from "../../wrappers/attachment-utils"; +import { BasePlugin } from "../core"; +import { + traceAsyncCall, + traceStreamingCall, + unsubscribeAll, +} from "../core/channel-tracing"; +import { cohereChannels } from "./cohere-channels"; export class CoherePlugin extends BasePlugin { protected onEnable(): void { @@ -27,67 +28,97 @@ export class CoherePlugin extends BasePlugin { private subscribeToCohereChannels(): void { this.unsubscribers.push( - traceStreamingChannel(cohereChannels.chat, { - name: "cohere.chat", - type: SpanTypeAttribute.LLM, - extractInput: extractChatInputWithMetadata, - extractOutput: (result) => extractCohereChatOutput(result), - extractMetadata: (result) => extractCohereResponseMetadata(result), - extractMetrics: (result, startTime) => { - const metrics = parseCohereMetricsFromUsage(result); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - }), + cohereChannels.chat.intercept((target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.COHERE, + name: "cohere.chat", + type: SpanTypeAttribute.LLM, + extractInput: extractChatInputWithMetadata, + extractOutput: (result) => extractCohereChatOutput(result), + extractMetadata: (result) => extractCohereResponseMetadata(result), + extractMetrics: (result, startTime) => { + const metrics = parseCohereMetricsFromUsage(result); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(cohereChannels.chatStream, { - name: "cohere.chatStream", - type: SpanTypeAttribute.LLM, - extractInput: extractChatInputWithMetadata, - extractOutput: () => undefined, - extractMetadata: () => undefined, - extractMetrics: () => ({}), - aggregateChunks: aggregateCohereChatStreamChunks, - }), + cohereChannels.chatStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.COHERE, + name: "cohere.chatStream", + type: SpanTypeAttribute.LLM, + extractInput: extractChatInputWithMetadata, + extractOutput: () => undefined, + extractMetadata: () => undefined, + extractMetrics: () => ({}), + aggregateChunks: aggregateCohereChatStreamChunks, + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(cohereChannels.embed, { - name: "cohere.embed", - type: SpanTypeAttribute.LLM, - extractInput: extractEmbedInputWithMetadata, - extractOutput: extractCohereEmbeddingOutput, - extractMetadata: (result) => extractCohereResponseMetadata(result), - extractMetrics: (result) => parseCohereMetricsFromUsage(result), - }), + cohereChannels.embed.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.COHERE, + name: "cohere.embed", + type: SpanTypeAttribute.LLM, + extractInput: extractEmbedInputWithMetadata, + extractOutput: extractCohereEmbeddingOutput, + extractMetadata: (result) => extractCohereResponseMetadata(result), + extractMetrics: (result) => parseCohereMetricsFromUsage(result), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(cohereChannels.rerank, { - name: "cohere.rerank", - type: SpanTypeAttribute.LLM, - extractInput: extractRerankInputWithMetadata, - extractOutput: (result) => { - if (!isObject(result) || !Array.isArray(result.results)) { - return undefined; - } - - return result.results.slice(0, 100).map((item) => ({ - index: isObject(item) ? item.index : undefined, - relevance_score: isObject(item) - ? ((typeof item.relevanceScore === "number" - ? item.relevanceScore - : item.relevance_score) ?? null) - : null, - })); - }, - extractMetadata: (result) => extractCohereResponseMetadata(result), - extractMetrics: (result) => parseCohereMetricsFromUsage(result), - }), + cohereChannels.rerank.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.COHERE, + name: "cohere.rerank", + type: SpanTypeAttribute.LLM, + extractInput: extractRerankInputWithMetadata, + extractOutput: (result) => { + if (!isObject(result) || !Array.isArray(result.results)) { + return undefined; + } + + return result.results.slice(0, 100).map((item) => ({ + index: isObject(item) ? item.index : undefined, + relevance_score: isObject(item) + ? ((typeof item.relevanceScore === "number" + ? item.relevanceScore + : item.relevance_score) ?? null) + : null, + })); + }, + extractMetadata: (result) => extractCohereResponseMetadata(result), + extractMetrics: (result) => parseCohereMetricsFromUsage(result), + }, + ), + ), ); } } diff --git a/js/src/instrumentation/plugins/cursor-sdk-channels.ts b/js/src/instrumentation/plugins/cursor-sdk-channels.ts index 9ee604a1a..af9474763 100644 --- a/js/src/instrumentation/plugins/cursor-sdk-channels.ts +++ b/js/src/instrumentation/plugins/cursor-sdk-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { CursorSDKAgent, CursorSDKAgentOptions, @@ -9,44 +9,34 @@ import type { CursorSDKUserMessage, } from "../../vendor-sdk-types/cursor-sdk"; -export const cursorSDKChannels = defineChannels( - "@cursor/sdk", - { - create: channel< - [CursorSDKAgentOptions], - CursorSDKAgent, - Record - >({ +export const cursorSDKChannels = defineInterceptor("@cursor/sdk", { + create: channel<[CursorSDKAgentOptions], PromiseLike, object>( + { channelName: "Agent.create", - kind: "async", - }), - resume: channel< - [string, Partial | undefined], - CursorSDKAgent, - Record - >({ - channelName: "Agent.resume", - kind: "async", - }), - prompt: channel< - [string | CursorSDKUserMessage, CursorSDKAgentOptions | undefined], - CursorSDKRunResult, - Record - >({ - channelName: "Agent.prompt", - kind: "async", - }), - send: channel< - [string | CursorSDKUserMessage, CursorSDKSendOptions | undefined], - CursorSDKRun, - { - agent?: CursorSDKAgent; - operation?: "send"; - } - >({ - channelName: "agent.send", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.CURSOR_SDK }, -); + }, + ), + resume: channel< + [string, Partial | undefined], + PromiseLike, + object + >({ + channelName: "Agent.resume", + }), + prompt: channel< + [string | CursorSDKUserMessage, CursorSDKAgentOptions | undefined], + PromiseLike, + object + >({ + channelName: "Agent.prompt", + }), + send: channel< + [string | CursorSDKUserMessage, CursorSDKSendOptions | undefined], + PromiseLike, + { + agent?: CursorSDKAgent; + operation?: "send"; + } + >({ + channelName: "agent.send", + }), +}); diff --git a/js/src/instrumentation/plugins/cursor-sdk-plugin.test.ts b/js/src/instrumentation/plugins/cursor-sdk-plugin.test.ts index b22052c8c..7005d5b97 100644 --- a/js/src/instrumentation/plugins/cursor-sdk-plugin.test.ts +++ b/js/src/instrumentation/plugins/cursor-sdk-plugin.test.ts @@ -1,23 +1,30 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +import { invocationController } from "../test-utils/invocation"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const { mockStartSpan } = vi.hoisted(() => ({ mockStartSpan: vi.fn(), })); vi.mock("../../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(), - }, + default: {}, })); vi.mock("../../logger", () => ({ startSpan: (...args: unknown[]) => mockStartSpan(...args), })); -import iso from "../../isomorph"; import { CursorSDKPlugin } from "./cursor-sdk-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; describe("CursorSDKPlugin", () => { let handlersByName: Map; @@ -31,9 +38,12 @@ describe("CursorSDKPlugin", () => { beforeEach(() => { handlersByName = new Map(); spans = []; - mockNewTracingChannel.mockImplementation((name: string) => ({ - subscribe: vi.fn((handlers) => handlersByName.set(name, handlers)), - tracePromise: vi.fn((fn) => fn()), + mockNewInvocationHook.mockImplementation((name: string) => ({ + intercept: vi.fn((interceptor) => { + handlersByName.set(name, invocationController(interceptor)); + return vi.fn(); + }), + invoke: vi.fn((fn, receiver, args) => Reflect.apply(fn, receiver, args)), unsubscribe: vi.fn(), })); mockStartSpan.mockImplementation((args: any) => { @@ -89,7 +99,7 @@ describe("CursorSDKPlugin", () => { send: originalSend, }; - createHandlers.asyncEnd({ + createHandlers.call({ arguments: [{ local: { cwd: "/tmp/repo" } }], result: agent, }); @@ -103,8 +113,8 @@ describe("CursorSDKPlugin", () => { arguments: ["use a tool", {}], result: run, }; - sendHandlers.start(sendEvent); - sendHandlers.asyncEnd(sendEvent); + sendHandlers.begin(sendEvent); + sendHandlers.resolve(sendEvent); await run.wait(); @@ -168,7 +178,7 @@ describe("CursorSDKPlugin", () => { result: run, }; - sendHandlers.start(event); + sendHandlers.begin(event); await (event.arguments[1] as any).onDelta({ update: { type: "turn-ended", @@ -180,7 +190,7 @@ describe("CursorSDKPlugin", () => { }, }, }); - sendHandlers.asyncEnd(event); + sendHandlers.resolve(event); const chunks = []; for await (const chunk of run.stream()) { @@ -225,12 +235,13 @@ describe("CursorSDKPlugin", () => { ); const promptEvent = { arguments: ["hello", { local: { cwd: "/tmp" } }] }; - promptHandlers.start(promptEvent); - sendHandlers.start({ arguments: ["nested", {}] }); - promptHandlers.asyncEnd({ - ...promptEvent, - result: { id: "run-1", result: "done", status: "finished" }, - }); + promptHandlers.begin(promptEvent); + sendHandlers.begin({ arguments: ["nested", {}] }); + promptHandlers.resolve( + Object.assign(promptEvent, { + result: { id: "run-1", result: "done", status: "finished" }, + }), + ); expect(spans.filter((span) => span.name === "Cursor Agent")).toHaveLength( 1, diff --git a/js/src/instrumentation/plugins/cursor-sdk-plugin.ts b/js/src/instrumentation/plugins/cursor-sdk-plugin.ts index f01db5987..e1dd1f61d 100644 --- a/js/src/instrumentation/plugins/cursor-sdk-plugin.ts +++ b/js/src/instrumentation/plugins/cursor-sdk-plugin.ts @@ -1,16 +1,15 @@ import { BasePlugin, toLoggedError } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers } from "../../isomorph"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute } from "../../../util/index"; import { debugLogger } from "../../debug-logger"; -import { startSpan as startBaseSpan } from "../../logger"; import type { Span } from "../../logger"; +import { startSpan as startBaseSpan } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; import { getCurrentUnixTimestamp } from "../../util"; -import { SpanTypeAttribute } from "../../../util/index"; -import { cursorSDKChannels } from "./cursor-sdk-channels"; import type { CursorSDKAgent, CursorSDKAgentOptions, @@ -28,6 +27,8 @@ import type { CursorSDKUsage, CursorSDKUserMessage, } from "../../vendor-sdk-types/cursor-sdk"; +import type { ChannelMessage } from "../core/tracing-types"; +import { cursorSDKChannels } from "./cursor-sdk-channels"; const PATCHED_AGENT = Symbol.for("braintrust.cursor-sdk.auto-patched-agent"); const PATCHED_RUN = Symbol.for("braintrust.cursor-sdk.patched-run"); @@ -89,187 +90,238 @@ export class CursorSDKPlugin extends BasePlugin { private subscribeToAgentFactory( channel: typeof cursorSDKChannels.create | typeof cursorSDKChannels.resume, ): void { - const tracingChannel = channel.tracingChannel(); - const handlers: IsoChannelHandlers> = { - asyncEnd: (event) => { - patchCursorAgentInPlace(event.result); - }, - error: () => {}, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - tracingChannel.unsubscribe(handlers); - }); + this.unsubscribers.push( + channel.intercept((target, receiver, args) => + observeResult( + Reflect.apply(target, receiver, args), + patchCursorAgentInPlace, + () => {}, + ), + ), + ); } private subscribeToPrompt(): void { - const channel = cursorSDKChannels.prompt.tracingChannel(); + const channel = cursorSDKChannels.prompt; const states = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - this.promptDepth += 1; - const message = event.arguments[0]; - const options = event.arguments[1]; - const metadata = { - ...extractAgentOptionsMetadata(options), - "cursor_sdk.operation": "Agent.prompt", - provider: "cursor", - ...(event.moduleVersion - ? { "cursor_sdk.version": event.moduleVersion } - : {}), - }; - const span = startBaseSpan( - withSpanInstrumentationName( - { - name: "Cursor Agent", - spanAttributes: { type: SpanTypeAttribute.TASK }, - }, - INSTRUMENTATION_NAMES.CURSOR_SDK, - ), - ); - const startTime = getCurrentUnixTimestamp(); - safeLog(span, { - input: sanitizeUserMessage(message), - metadata, - }); - states.set(event, { metadata, span, startTime }); - }, - asyncEnd: (event) => { - this.promptDepth = Math.max(0, this.promptDepth - 1); - const state = states.get(event); - if (!state) { - return; - } - try { - safeLog(state.span, { - metadata: { - ...state.metadata, - ...extractRunResultMetadata(event.result), - }, - metrics: buildDurationMetrics(state.startTime), - output: event.result?.result ?? event.result, + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + this.promptDepth += 1; + const message = event.arguments[0]; + const options = event.arguments[1]; + const metadata = { + ...extractAgentOptionsMetadata(options), + "cursor_sdk.operation": "Agent.prompt", + provider: "cursor", + ...(event.moduleVersion + ? { "cursor_sdk.version": event.moduleVersion } + : {}), + }; + const span = startBaseSpan( + withSpanInstrumentationName( + { + name: "Cursor Agent", + spanAttributes: { type: SpanTypeAttribute.TASK }, + }, + INSTRUMENTATION_NAMES.CURSOR_SDK, + ), + ); + const startTime = getCurrentUnixTimestamp(); + safeLog(span, { + input: sanitizeUserMessage(message), + metadata, }); - } finally { + states.set(event, { metadata, span, startTime }); + }; + const resolved = ( + event: ChannelMessage, + ) => { + this.promptDepth = Math.max(0, this.promptDepth - 1); + const state = states.get(event); + if (!state) { + return; + } + try { + safeLog(state.span, { + metadata: { + ...state.metadata, + ...extractRunResultMetadata(event.result), + }, + metrics: buildDurationMetrics(state.startTime), + output: event.result?.result ?? event.result, + }); + } finally { + state.span.end(); + states.delete(event); + } + }; + const failed = ( + event: ChannelMessage, + ) => { + this.promptDepth = Math.max(0, this.promptDepth - 1); + const state = states.get(event); + if (!state || !event.error) { + return; + } + safeLog(state.span, { error: event.error.message }); state.span.end(); states.delete(event); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); }, - error: (event) => { - this.promptDepth = Math.max(0, this.promptDepth - 1); - const state = states.get(event); - if (!state || !event.error) { - return; - } - safeLog(state.span, { error: event.error.message }); - state.span.end(); - states.delete(event); - }, - }; - - channel.subscribe(handlers); - this.unsubscribers.push(() => { - channel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToSend(): void { - const channel = cursorSDKChannels.send.tracingChannel(); + const channel = cursorSDKChannels.send; const states = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - if (this.promptDepth > 0) { - return; - } + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + if (this.promptDepth > 0) { + return; + } - const message = event.arguments[0]; - const sendOptions = event.arguments[1]; - const agent = event.agent; - const metadata = { - ...extractSendMetadata(sendOptions), - ...(agent ? extractAgentMetadata(agent) : {}), - "cursor_sdk.operation": "agent.send", - provider: "cursor", - ...(event.moduleVersion - ? { "cursor_sdk.version": event.moduleVersion } - : {}), - }; - const span = startBaseSpan( - withSpanInstrumentationName( - { - name: "Cursor Agent", - spanAttributes: { type: SpanTypeAttribute.TASK }, - }, - INSTRUMENTATION_NAMES.CURSOR_SDK, - ), - ); - const startTime = getCurrentUnixTimestamp(); - safeLog(span, { - input: sanitizeUserMessage(message), - metadata, - }); + const message = event.arguments[0]; + const sendOptions = event.arguments[1]; + const agent = event.agent; + const metadata = { + ...extractSendMetadata(sendOptions), + ...(agent ? extractAgentMetadata(agent) : {}), + "cursor_sdk.operation": "agent.send", + provider: "cursor", + ...(event.moduleVersion + ? { "cursor_sdk.version": event.moduleVersion } + : {}), + }; + const span = startBaseSpan( + withSpanInstrumentationName( + { + name: "Cursor Agent", + spanAttributes: { type: SpanTypeAttribute.TASK }, + }, + INSTRUMENTATION_NAMES.CURSOR_SDK, + ), + ); + const startTime = getCurrentUnixTimestamp(); + safeLog(span, { + input: sanitizeUserMessage(message), + metadata, + }); - const state: CursorRunState = { - activeToolSpans: new Map(), - agent, - conversationText: [], - deltaText: [], - finalized: false, - input: message, - metadata, - metrics: {}, - span, - startTime, - streamMessages: [], - streamText: [], - stepText: [], - taskText: [], + const state: CursorRunState = { + activeToolSpans: new Map(), + agent, + conversationText: [], + deltaText: [], + finalized: false, + input: message, + metadata, + metrics: {}, + span, + startTime, + streamMessages: [], + streamText: [], + stepText: [], + taskText: [], + }; + + if (hasCursorCallbacks(sendOptions)) { + event.arguments[1] = wrapSendOptionsCallbacks(sendOptions, state); + } + states.set(event, state); }; + const resolved = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } - if (hasCursorCallbacks(sendOptions)) { - event.arguments[1] = wrapSendOptionsCallbacks(sendOptions, state); - } - states.set(event, state); - }, - asyncEnd: (event) => { - const state = states.get(event); - if (!state) { - return; - } - - if (!event.result) { - return; - } - state.run = event.result; - state.metadata = { - ...state.metadata, - ...extractRunMetadata(event.result), + if (!event.result) { + return; + } + state.run = event.result; + state.metadata = { + ...state.metadata, + ...extractRunMetadata(event.result), + }; + patchCursorRun(event.result, state); }; - patchCursorRun(event.result, state); - }, - error: (event) => { - const state = states.get(event); - if (!state || !event.error) { - return; + const failed = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state || !event.error) { + return; + } + safeLog(state.span, { error: event.error.message }); + endOpenToolSpans(state, event.error.message); + state.span.end(); + state.finalized = true; + states.delete(event); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } - safeLog(state.span, { error: event.error.message }); - endOpenToolSpans(state, event.error.message); - state.span.end(); - state.finalized = true; - states.delete(event); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); }, - }; - - channel.subscribe(handlers); - this.unsubscribers.push(() => { - channel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } } @@ -303,13 +355,11 @@ function patchCursorAgentInPlace(agent: unknown): void { string | CursorSDKUserMessage, CursorSDKSendOptions | undefined, ]; - return cursorSDKChannels.send.tracePromise( - () => originalSend(...args), - { - agent: agentRecord, - arguments: args, - operation: "send", - } as never, + return cursorSDKChannels.send.invoke( + originalSend, + undefined, + [...args], + { agent: agentRecord, operation: "send" }, ); }, writable: true, diff --git a/js/src/instrumentation/plugins/elevenlabs-channels.ts b/js/src/instrumentation/plugins/elevenlabs-channels.ts index e5ba16be4..bdc843e0c 100644 --- a/js/src/instrumentation/plugins/elevenlabs-channels.ts +++ b/js/src/instrumentation/plugins/elevenlabs-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { ElevenLabsAudio, ElevenLabsSpeechArgs, @@ -8,29 +8,26 @@ import type { ElevenLabsTranscriptionRequest, } from "../../vendor-sdk-types/elevenlabs"; -export const elevenLabsChannels = defineChannels( +export const elevenLabsChannels = defineInterceptor( "@elevenlabs/elevenlabs-js", { - convert: channel({ + convert: channel>({ channelName: "textToSpeech.convert", - kind: "async", }), - stream: channel({ + stream: channel>({ channelName: "textToSpeech.stream", - kind: "async", }), convertWithTimestamps: channel< ElevenLabsSpeechArgs, - ElevenLabsTimestampAudio - >({ channelName: "textToSpeech.convertWithTimestamps", kind: "async" }), + PromiseLike + >({ channelName: "textToSpeech.convertWithTimestamps" }), streamWithTimestamps: channel< ElevenLabsSpeechArgs, - AsyncIterable - >({ channelName: "textToSpeech.streamWithTimestamps", kind: "async" }), + PromiseLike> + >({ channelName: "textToSpeech.streamWithTimestamps" }), transcribe: channel< [ElevenLabsTranscriptionRequest, unknown?], - ElevenLabsTranscription - >({ channelName: "speechToText.convert", kind: "async" }), + PromiseLike + >({ channelName: "speechToText.convert" }), }, - { instrumentationName: INSTRUMENTATION_NAMES.ELEVENLABS }, ); diff --git a/js/src/instrumentation/plugins/elevenlabs-plugin.ts b/js/src/instrumentation/plugins/elevenlabs-plugin.ts index fc3abed71..0773e17c1 100644 --- a/js/src/instrumentation/plugins/elevenlabs-plugin.ts +++ b/js/src/instrumentation/plugins/elevenlabs-plugin.ts @@ -6,11 +6,9 @@ import { withSpanInstrumentationName, } from "../../span-origin"; import type { - ElevenLabsSpeechArgs, ElevenLabsSpeechRequest, ElevenLabsTimestampAudio, ElevenLabsTranscription, - ElevenLabsTranscriptionRequest, } from "../../vendor-sdk-types/elevenlabs"; import { getExtensionFromMediaType } from "../../wrappers/attachment-utils"; import { @@ -35,145 +33,158 @@ export class ElevenLabsPlugin extends BasePlugin { "streamWithTimestamps", ] as const) { this.unsubscribers.push( - interceptCall( - elevenLabsChannels[method], - `elevenlabs.textToSpeech.${method}`, - ([voice, request]) => ({ - input: { - operation: "speech", - prompt: request.text, - parameters: { - voice, - format: request.outputFormat, - language: request.languageCode, - speed: request.voiceSettings?.speed, - }, - }, - metadata: { - provider: "elevenlabs", - model: request.modelId ?? "eleven_multilingual_v2", - }, - }), - (value, args, span, finish, headers, started) => - captureSpeech( - value, - args[1], - method, - span, - finish, - started, - headers, + elevenLabsChannels[method].intercept( + (target, receiver, args, additional) => + traceElevenLabsCall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + `elevenlabs.textToSpeech.${method}`, + ([voice, request]) => ({ + input: { + operation: "speech", + prompt: request.text, + parameters: { + voice, + format: request.outputFormat, + language: request.languageCode, + speed: request.voiceSettings?.speed, + }, + }, + metadata: { + provider: "elevenlabs", + model: request.modelId ?? "eleven_multilingual_v2", + }, + }), + (value, args, span, finish, headers, started) => + captureSpeech( + value, + args[1], + method, + span, + finish, + started, + headers, + ), ), ), ); } this.unsubscribers.push( - interceptCall<[ElevenLabsTranscriptionRequest, unknown?]>( - elevenLabsChannels.transcribe, - "elevenlabs.speechToText.convert", - ([request]) => { - // Webhook requests return an acknowledgement; the transcription arrives - // separately and cannot be captured by this request/response span. - if (request.webhook) return undefined; - const file = request.file; - const filename = - typeof File !== "undefined" && file instanceof File - ? file.name - : "audio"; - const contentType = - file instanceof Blob - ? file.type || "application/octet-stream" - : "application/octet-stream"; - const blob = - file instanceof Blob - ? file - : file instanceof Uint8Array - ? new Blob([new Uint8Array(file)], { type: contentType }) - : file instanceof ArrayBuffer - ? new Blob([file.slice(0)], { type: contentType }) - : undefined; - const fileData = blob - ? new Attachment({ data: blob, filename, contentType }) - : request.cloudStorageUrl; - return { - input: { - operation: "transcribe", - content: fileData - ? [{ type: "file", file: { filename, file_data: fileData } }] - : [], - parameters: { - language: request.languageCode, - timestamp_granularities: request.timestampsGranularity, - }, - }, - metadata: { provider: "elevenlabs", model: request.modelId }, - }; - }, - (value, _args, span, finish) => { - const result = value as ElevenLabsTranscription; - const transcripts = result.transcripts ?? [result]; - span.log({ - output: { - content: transcripts.flatMap((transcript) => - typeof transcript.text === "string" - ? [{ type: "text", text: transcript.text }] - : [], - ), - annotations: { - language: result.languageCode, - words: - result.words ?? - (result.transcripts - ? result.transcripts.flatMap( - (transcript) => transcript.words ?? [], - ) - : undefined), - }, + elevenLabsChannels.transcribe.intercept( + (target, receiver, args, additional) => + traceElevenLabsCall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "elevenlabs.speechToText.convert", + ([request]) => { + // Webhook requests return an acknowledgement; the transcription arrives + // separately and cannot be captured by this request/response span. + if (request.webhook) return undefined; + const file = request.file; + const filename = + typeof File !== "undefined" && file instanceof File + ? file.name + : "audio"; + const contentType = + file instanceof Blob + ? file.type || "application/octet-stream" + : "application/octet-stream"; + const blob = + file instanceof Blob + ? file + : file instanceof Uint8Array + ? new Blob([new Uint8Array(file)], { type: contentType }) + : file instanceof ArrayBuffer + ? new Blob([file.slice(0)], { type: contentType }) + : undefined; + const fileData = blob + ? new Attachment({ data: blob, filename, contentType }) + : request.cloudStorageUrl; + return { + input: { + operation: "transcribe", + content: fileData + ? [ + { + type: "file", + file: { filename, file_data: fileData }, + }, + ] + : [], + parameters: { + language: request.languageCode, + timestamp_granularities: request.timestampsGranularity, + }, + }, + metadata: { provider: "elevenlabs", model: request.modelId }, + }; }, - }); - finish(); - }, - ([request], span) => { - const file = request.file; - if (!isAsyncIterable(file)) return; - const chunks: Uint8Array[] = []; - const filename = - isObject(file) && typeof file.path === "string" - ? file.path.split(/[\\/]/).pop() || "audio" - : "audio"; - observeByteStream(file, { - onChunk: (chunk) => chunks.push(new Uint8Array(chunk)), - onComplete: () => { - const contentType = "application/octet-stream"; - const data = new Blob(chunks as BlobPart[], { - type: contentType, - }); - chunks.length = 0; + (value, _args, span, finish) => { + const result = value as ElevenLabsTranscription; + const transcripts = result.transcripts ?? [result]; span.log({ - input: { - content: [ - { - type: "file", - file: { - filename, - file_data: new Attachment({ - data, - filename, - contentType, - }), - }, - }, - ], + output: { + content: transcripts.flatMap((transcript) => + typeof transcript.text === "string" + ? [{ type: "text", text: transcript.text }] + : [], + ), + annotations: { + language: result.languageCode, + words: + result.words ?? + (result.transcripts + ? result.transcripts.flatMap( + (transcript) => transcript.words ?? [], + ) + : undefined), + }, }, }); + finish(); }, - onCancel: () => { - chunks.length = 0; + ([request], span) => { + const file = request.file; + if (!isAsyncIterable(file)) return; + const chunks: Uint8Array[] = []; + const filename = + isObject(file) && typeof file.path === "string" + ? file.path.split(/[\\/]/).pop() || "audio" + : "audio"; + observeByteStream(file, { + onChunk: (chunk) => chunks.push(new Uint8Array(chunk)), + onComplete: () => { + const contentType = "application/octet-stream"; + const data = new Blob(chunks as BlobPart[], { + type: contentType, + }); + chunks.length = 0; + span.log({ + input: { + content: [ + { + type: "file", + file: { + filename, + file_data: new Attachment({ + data, + filename, + contentType, + }), + }, + }, + ], + }, + }); + }, + onCancel: () => { + chunks.length = 0; + }, + aroundRead: (next) => withCurrent(span, next), + debugLabel: "ElevenLabs audio", + }); }, - aroundRead: (next) => withCurrent(span, next), - debugLabel: "ElevenLabs audio", - }); - }, + ), ), ); } @@ -183,18 +194,10 @@ export class ElevenLabsPlugin extends BasePlugin { } type Finish = (error?: unknown) => void; -type CallChannel = { - intercept( - callback: ( - target: (...args: Args) => PromiseLike, - self: unknown, - args: Args, - ) => PromiseLike, - ): () => void; -}; -function interceptCall( - channel: CallChannel, +function traceElevenLabsCall( + call: () => TResult, + context: { arguments: Args; self: unknown; additional: unknown }, name: string, input: ( args: Args, @@ -208,89 +211,86 @@ function interceptCall( started: number, ) => void, beforeInvoke?: (args: Args, span: Span) => void, -): () => void { - return channel.intercept((target, self, args) => { - const invoke = () => Reflect.apply(target, self, args); - if (isAutoInstrumentationSuppressed()) return invoke(); - const started = Date.now() / 1000; - let span: Span | undefined; - try { - const event = input(args); - if (event !== undefined) - span = startSpan( - withSpanInstrumentationName( - { - name, - spanAttributes: { type: SpanTypeAttribute.LLM }, - event, - }, - INSTRUMENTATION_NAMES.ELEVENLABS, - ), - ); - } catch (error) { - debugLogger.error("Error starting ElevenLabs span", error); - return invoke(); - } - if (!span) return invoke(); - const activeSpan = span; - let ended = false; - const finish: Finish = (error) => { - if (ended) return; - ended = true; - try { - if (error !== undefined) activeSpan.log({ error }); - activeSpan.end(); - } catch (loggingError) { - debugLogger.error("Error ending ElevenLabs span", loggingError); - } - }; - try { - beforeInvoke?.(args, span); - } catch (error) { - debugLogger.error("Error observing ElevenLabs input", error); - } - let result: PromiseLike; - try { - result = withCurrent(span, () => - runWithAutoInstrumentationSuppressed(invoke), +): TResult { + const args = context.arguments; + + if (isAutoInstrumentationSuppressed()) return call(); + const started = Date.now() / 1000; + let span: Span | undefined; + try { + const event = input(args); + if (event !== undefined) + span = startSpan( + withSpanInstrumentationName( + { + name, + spanAttributes: { type: SpanTypeAttribute.LLM }, + event, + }, + INSTRUMENTATION_NAMES.ELEVENLABS, + ), ); - } catch (error) { - finish(error); - throw error; + } catch (error) { + debugLogger.error("Error starting ElevenLabs span", error); + return call(); + } + if (!span) return call(); + const activeSpan = span; + let ended = false; + const finish: Finish = (error) => { + if (ended) return; + ended = true; + try { + if (error !== undefined) activeSpan.log({ error }); + activeSpan.end(); + } catch (loggingError) { + debugLogger.error("Error ending ElevenLabs span", loggingError); } - // Observe the SDK promise without replacing it: withRawResponse() must remain - // available, including when it is the application's only consumption path. - const capture = (value: unknown, headers?: Headers) => { - try { - output(value, args, span, finish, headers, started); - } catch (error) { - debugLogger.error("Error capturing ElevenLabs output", error); - finish(); - } - }; + }; + try { + beforeInvoke?.(args, span); + } catch (error) { + debugLogger.error("Error observing ElevenLabs input", error); + } + let result: TResult; + try { + result = withCurrent(span, () => + runWithAutoInstrumentationSuppressed(call), + ); + } catch (error) { + finish(error); + throw error; + } + const capture = (value: unknown, headers?: Headers) => { try { - if (isObject(result) && typeof result.withRawResponse === "function") { - void result - .withRawResponse() - .then( - ({ - data, - rawResponse, - }: { - data: unknown; - rawResponse?: { headers?: Headers }; - }) => capture(data, rawResponse?.headers), - finish, - ); - } else { - void Promise.resolve(result).then((value) => capture(value), finish); - } + output(value, args, span, finish, headers, started); } catch (error) { - debugLogger.error("Error observing ElevenLabs result", error); + debugLogger.error("Error capturing ElevenLabs output", error); finish(); } - return result; - }); + }; + try { + if (isObject(result) && typeof result.withRawResponse === "function") { + void result + .withRawResponse() + .then( + ({ + data, + rawResponse, + }: { + data: unknown; + rawResponse?: { headers?: Headers }; + }) => capture(data, rawResponse?.headers), + finish, + ); + } else { + void Promise.resolve(result).then((value) => capture(value), finish); + } + } catch (error) { + debugLogger.error("Error observing ElevenLabs result", error); + finish(); + } + return result; } function captureSpeech( diff --git a/js/src/instrumentation/plugins/flue-channels.ts b/js/src/instrumentation/plugins/flue-channels.ts index 308d50aea..f98429c62 100644 --- a/js/src/instrumentation/plugins/flue-channels.ts +++ b/js/src/instrumentation/plugins/flue-channels.ts @@ -1,14 +1,9 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { FlueObservableContext } from "../../vendor-sdk-types/flue"; -export const flueChannels = defineChannels( - "@flue/runtime", - { - createContext: channel<[unknown], FlueObservableContext>({ - channelName: "createFlueContext", - kind: "sync-stream", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.FLUE }, -); +export const flueChannels = defineInterceptor("@flue/runtime", { + createContext: channel<[unknown], FlueObservableContext>({ + channelName: "createFlueContext", + }), +}); diff --git a/js/src/instrumentation/plugins/flue-plugin.test.ts b/js/src/instrumentation/plugins/flue-plugin.test.ts index 6c891c4c0..aa734d3f8 100644 --- a/js/src/instrumentation/plugins/flue-plugin.test.ts +++ b/js/src/instrumentation/plugins/flue-plugin.test.ts @@ -37,40 +37,31 @@ vi.mock("../../debug-logger", () => ({ }, })); -const { mockNewTracingChannel, mockTracingChannels } = vi.hoisted(() => { - const tracingChannels = new Map(); +const { mockNewInvocationHook, mockInvocationHooks } = vi.hoisted(() => { + const invocationHooks = new Map(); - function tracingChannel(name: string) { - const existing = tracingChannels.get(name); + function invocationHook(name: string) { + const existing = invocationHooks.get(name); if (existing) { return existing; } const handlers = new Set(); - const stores = new Map unknown>(); const channel = { __handlers: handlers, - __stores: stores, - start: { - bindStore: vi.fn( - (store: unknown, transform: (message: any) => unknown) => { - stores.set(store, transform); - }, - ), - unbindStore: vi.fn((store: unknown) => stores.delete(store)), - }, - subscribe: vi.fn((handler: any) => { + remove: vi.fn((handler: any) => handlers.delete(handler)), + intercept: vi.fn((handler: any) => { handlers.add(handler); + return () => channel.remove(handler); }), - unsubscribe: vi.fn((handler: any) => handlers.delete(handler)), }; - tracingChannels.set(name, channel); + invocationHooks.set(name, channel); return channel; } return { - mockNewTracingChannel: vi.fn((name: string) => tracingChannel(name)), - mockTracingChannels: tracingChannels, + mockNewInvocationHook: vi.fn((name: string) => invocationHook(name)), + mockInvocationHooks: invocationHooks, }; }); @@ -101,10 +92,11 @@ vi.mock("../../logger", () => ({ }, })); -vi.mock("../../isomorph", () => ({ - default: { - newTracingChannel: mockNewTracingChannel, - }, +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: mockNewInvocationHook, })); import { @@ -166,9 +158,8 @@ describe("Flue observe instrumentation", () => { delete (globalThis as Record)[ Symbol.for("braintrust.flue.observe-bridge") ]; - for (const channel of mockTracingChannels.values()) { + for (const channel of mockInvocationHooks.values()) { channel.__handlers.clear(); - channel.__stores.clear(); } vi.clearAllMocks(); }); @@ -1450,11 +1441,11 @@ describe("Flue observe instrumentation", () => { }; plugin.enable(); - expect(mockNewTracingChannel).toHaveBeenCalledWith( + expect(mockNewInvocationHook).toHaveBeenCalledWith( CREATE_CONTEXT_CHANNEL_NAME, ); expect( - tracingChannel(CREATE_CONTEXT_CHANNEL_NAME).subscribe, + invocationHook(CREATE_CONTEXT_CHANNEL_NAME).intercept, ).toHaveBeenCalledTimes(1); emitCreateContextEnd(context); @@ -1484,7 +1475,7 @@ describe("Flue observe instrumentation", () => { plugin.disable(); expect( - tracingChannel(CREATE_CONTEXT_CHANNEL_NAME).unsubscribe, + invocationHook(CREATE_CONTEXT_CHANNEL_NAME).remove, ).toHaveBeenCalledTimes(1); expect(unsubscribeContext).toHaveBeenCalledTimes(1); }); @@ -1507,19 +1498,19 @@ describe("Flue observe instrumentation", () => { emitCreateContextEnd(context); expect( - tracingChannel(CREATE_CONTEXT_CHANNEL_NAME).subscribe, + invocationHook(CREATE_CONTEXT_CHANNEL_NAME).intercept, ).toHaveBeenCalledTimes(1); expect(context.subscribeEvent).toHaveBeenCalledTimes(1); first.disable(); expect( - tracingChannel(CREATE_CONTEXT_CHANNEL_NAME).unsubscribe, + invocationHook(CREATE_CONTEXT_CHANNEL_NAME).remove, ).not.toHaveBeenCalled(); expect(unsubscribeContext).not.toHaveBeenCalled(); second.disable(); expect( - tracingChannel(CREATE_CONTEXT_CHANNEL_NAME).unsubscribe, + invocationHook(CREATE_CONTEXT_CHANNEL_NAME).remove, ).toHaveBeenCalledTimes(1); contextSubscribers[0]?.({ @@ -1567,14 +1558,14 @@ describe("Flue observe instrumentation", () => { } function emitCreateContextEnd(result: unknown) { - for (const handlers of tracingChannel(CREATE_CONTEXT_CHANNEL_NAME) + for (const handlers of invocationHook(CREATE_CONTEXT_CHANNEL_NAME) .__handlers) { - handlers.end?.({ result }); + handlers(() => result, undefined, [], {}); } } - function tracingChannel(channelName: string) { - const channel = mockTracingChannels.get(channelName); + function invocationHook(channelName: string) { + const channel = mockInvocationHooks.get(channelName); if (!channel) { throw new Error(`Missing mocked tracing channel: ${channelName}`); } diff --git a/js/src/instrumentation/plugins/flue-plugin.ts b/js/src/instrumentation/plugins/flue-plugin.ts index af8ff5153..6f8e64886 100644 --- a/js/src/instrumentation/plugins/flue-plugin.ts +++ b/js/src/instrumentation/plugins/flue-plugin.ts @@ -1,23 +1,21 @@ -import { BasePlugin, toLoggedError } from "../core"; import { debugLogger } from "../../debug-logger"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers } from "../../isomorph"; +import { BasePlugin, toLoggedError } from "../core"; +import { runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute } from "../../../util/index"; +import type { CurrentSpanStore, Span, StartSpanArgs } from "../../logger"; import { BRAINTRUST_CURRENT_SPAN_STORE, NOOP_SPAN, - flush, _internalGetGlobalState, + flush, startSpan as startBaseSpan, withCurrent, } from "../../logger"; -import type { Span, StartSpanArgs } from "../../logger"; -import type { CurrentSpanStore } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { SpanTypeAttribute } from "../../../util/index"; -import { flueChannels } from "./flue-channels"; import type { FlueBaseEvent, FlueCompactionEvent, @@ -43,17 +41,14 @@ import type { FlueTurnEvent, FlueTurnRequestEvent, } from "../../vendor-sdk-types/flue"; +import type { ChannelMessage } from "../core/tracing-types"; +import { flueChannels } from "./flue-channels"; type FlueObserver = (event: unknown, ctx?: unknown) => void; type BraintrustFlueObserver = FlueObserver & FlueInstrumentation; type FlueAutoState = { - createContextChannel?: ReturnType< - typeof flueChannels.createContext.tracingChannel - >; - createContextHandlers?: IsoChannelHandlers< - ChannelMessage - >; + removeCreateContextInterceptor?: () => void; contexts: WeakSet; refCount: number; }; @@ -128,19 +123,33 @@ function enableFlueAutoInstrumentation(): () => void { const state = getAutoState(); state.refCount += 1; - if (!state.createContextHandlers) { - const createContextChannel = flueChannels.createContext.tracingChannel(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - end: (event) => { - subscribeToFlueContext(event.result, state); + if (!state.removeCreateContextInterceptor) { + const createContextChannel = flueChannels.createContext; + + const removeHandlers = createContextChannel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const returned = ( + event: ChannelMessage, + ) => { + subscribeToFlueContext(event.result, state); + }; + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + throw error; + } + Object.assign(event, { result }); + runInstrumentation(() => returned(event)); + return result; }, - }; - - createContextChannel.subscribe(handlers); - state.createContextChannel = createContextChannel; - state.createContextHandlers = handlers; + ); + state.removeCreateContextInterceptor = removeHandlers; } let released = false; @@ -199,9 +208,7 @@ function releaseAutoState(state: FlueAutoState): void { } try { - if (state.createContextChannel && state.createContextHandlers) { - state.createContextChannel.unsubscribe(state.createContextHandlers); - } + state.removeCreateContextInterceptor?.(); } finally { Reflect.deleteProperty(globalThis, FLUE_AUTO_STATE); } diff --git a/js/src/instrumentation/plugins/genkit-channels.ts b/js/src/instrumentation/plugins/genkit-channels.ts index 5f1ac13ff..6574f1f79 100644 --- a/js/src/instrumentation/plugins/genkit-channels.ts +++ b/js/src/instrumentation/plugins/genkit-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { GenkitAction, GenkitEmbedManyParams, @@ -11,59 +11,46 @@ import type { GenkitGenerateStreamResponse, } from "../../vendor-sdk-types/genkit"; -export const genkitChannels = defineChannels( - "@genkit-ai/ai", - { - generate: channel<[GenkitGenerateInput], GenkitGenerateResponse>({ +export const genkitChannels = defineInterceptor("@genkit-ai/ai", { + generate: channel<[GenkitGenerateInput], PromiseLike>( + { channelName: "generate", - kind: "async", - }), + }, + ), - generateStream: channel< - [GenkitGenerateInput], - GenkitGenerateStreamResponse, - Record, - GenkitGenerateResponseChunk - >({ - channelName: "generateStream", - kind: "sync-stream", - }), + generateStream: channel< + [GenkitGenerateInput], + GenkitGenerateStreamResponse, + Record, + GenkitGenerateResponseChunk + >({ + channelName: "generateStream", + }), - embed: channel<[GenkitEmbedParams], GenkitEmbedding[]>({ - channelName: "embed", - kind: "async", - }), + embed: channel<[GenkitEmbedParams], PromiseLike>({ + channelName: "embed", + }), - embedMany: channel<[GenkitEmbedManyParams], unknown>({ - channelName: "embedMany", - kind: "async", - }), + embedMany: channel<[GenkitEmbedManyParams], PromiseLike>({ + channelName: "embedMany", + }), - actionRun: channel<[unknown, unknown?], unknown>({ - channelName: "action.run", - kind: "async", - }), + actionRun: channel<[unknown, unknown?], PromiseLike>({ + channelName: "action.run", + }), - actionStream: channel< - [unknown, unknown?], - ReturnType>, - Record, - unknown - >({ - channelName: "action.stream", - kind: "sync-stream", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.GENKIT }, -); + actionStream: channel< + [unknown, unknown?], + ReturnType>, + Record, + unknown + >({ + channelName: "action.stream", + }), +}); -export const genkitCoreChannels = defineChannels( - "@genkit-ai/core", - { - actionSpan: channel<[unknown, unknown, unknown?], unknown>({ - channelName: "action.span", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.GENKIT }, -); +export const genkitCoreChannels = defineInterceptor("@genkit-ai/core", { + actionSpan: channel<[unknown, unknown, unknown?], PromiseLike>({ + channelName: "action.span", + }), +}); diff --git a/js/src/instrumentation/plugins/genkit-plugin.test.ts b/js/src/instrumentation/plugins/genkit-plugin.test.ts index dcf41c0f7..9da0d4d30 100644 --- a/js/src/instrumentation/plugins/genkit-plugin.test.ts +++ b/js/src/instrumentation/plugins/genkit-plugin.test.ts @@ -1,7 +1,7 @@ import { afterEach, beforeAll, beforeEach, describe, expect, it } from "vitest"; import { _exportsForTestingOnly, initLogger } from "../../logger"; -import { GenkitPlugin } from "./genkit-plugin"; import { genkitChannels } from "./genkit-channels"; +import { GenkitPlugin } from "./genkit-plugin"; function singleQueueStream( chunks: T[], @@ -59,7 +59,7 @@ describe("GenkitPlugin stream patching", () => { plugin.enable(); const stream = singleQueueStream([{ text: "hello" }, { text: " world" }]); - const result = genkitChannels.generateStream.traceSync( + const result = genkitChannels.generateStream.invoke( () => ({ response: Promise.resolve({ text: "hello world", @@ -71,9 +71,9 @@ describe("GenkitPlugin stream patching", () => { }), stream, }), - { arguments: [{ prompt: "Say hello world." }] } as Parameters< - typeof genkitChannels.generateStream.traceSync - >[1], + undefined, + [{ prompt: "Say hello world." }], + {}, ); await drainMicrotasks(); @@ -94,15 +94,14 @@ describe("GenkitPlugin stream patching", () => { }, }); - const result = genkitChannels.actionStream.traceSync( + const result = genkitChannels.actionStream.invoke( () => ({ output: Promise.resolve({ done: true }), stream, }), - { - arguments: [{ input: true }], - self: action, - } as Parameters[1], + action, + [{ input: true }], + {}, ); await drainMicrotasks(); diff --git a/js/src/instrumentation/plugins/genkit-plugin.ts b/js/src/instrumentation/plugins/genkit-plugin.ts index 90edb66e4..7b6f58b32 100644 --- a/js/src/instrumentation/plugins/genkit-plugin.ts +++ b/js/src/instrumentation/plugins/genkit-plugin.ts @@ -1,26 +1,21 @@ +import { withCurrent } from "../../logger"; import { BasePlugin, toLoggedError } from "../core"; import { - traceAsyncChannel, - traceSyncStreamChannel, + traceAsyncCall, + traceSyncStreamCall, unsubscribeAll, } from "../core/channel-tracing"; +import { observeResult, runInstrumentation } from "../core/observe-result"; import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; -import { - _internalGetGlobalState, - BRAINTRUST_CURRENT_SPAN_STORE, - startSpan as startBaseSpan, -} from "../../logger"; -import type { CurrentSpanStore, Span } from "../../logger"; + +import { SpanTypeAttribute } from "../../../util/index"; +import type { Span } from "../../logger"; +import { startSpan as startBaseSpan } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; import { getCurrentUnixTimestamp, isObject } from "../../util"; -import { SpanTypeAttribute } from "../../../util/index"; -import { processInputAttachments } from "../../wrappers/attachment-utils"; -import { genkitChannels, genkitCoreChannels } from "./genkit-channels"; import type { GenkitAction, GenkitActionMetadata, @@ -32,6 +27,9 @@ import type { GenkitGenerateStreamResponse, GenkitUsage, } from "../../vendor-sdk-types/genkit"; +import { processInputAttachments } from "../../wrappers/attachment-utils"; +import type { ChannelMessage } from "../core/tracing-types"; +import { genkitChannels, genkitCoreChannels } from "./genkit-channels"; type SpanState = { span: Span; @@ -49,49 +47,78 @@ export class GenkitPlugin extends BasePlugin { private subscribeToGenkitChannels(): void { this.unsubscribers.push( - traceAsyncChannel(genkitChannels.generate, { - name: "genkit.generate", - type: SpanTypeAttribute.LLM, - extractInput: ([input]) => extractGenerateInput(input), - extractOutput: extractGenerateOutput, - extractMetadata: (result, event) => - extractGenerateResponseMetadata(result, event?.arguments?.[0]), - extractMetrics: (result) => parseGenkitUsageMetrics(result?.usage), - }), + genkitChannels.generate.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GENKIT, + name: "genkit.generate", + type: SpanTypeAttribute.LLM, + extractInput: ([input]) => extractGenerateInput(input), + extractOutput: extractGenerateOutput, + extractMetadata: (result, event) => + extractGenerateResponseMetadata(result, event?.arguments?.[0]), + extractMetrics: (result) => parseGenkitUsageMetrics(result?.usage), + }, + ), + ), ); this.unsubscribers.push( - traceSyncStreamChannel(genkitChannels.generateStream, { - name: "genkit.generateStream", - type: SpanTypeAttribute.LLM, - extractInput: ([input]) => extractGenerateInput(input), - patchResult: ({ result, span, startTime }) => - patchGenerateStreamResult(result, span, startTime), - }), + genkitChannels.generateStream.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GENKIT, + name: "genkit.generateStream", + type: SpanTypeAttribute.LLM, + extractInput: ([input]) => extractGenerateInput(input), + patchResult: ({ result, span, startTime }) => + patchGenerateStreamResult(result, span, startTime), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(genkitChannels.embed, { - name: "genkit.embed", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params]) => extractEmbedInput(params), - extractOutput: (result) => summarizeEmbeddingResult(result), - extractMetadata: (_result, event) => - extractEmbedMetadata(event?.arguments?.[0]), - extractMetrics: () => ({}), - }), + genkitChannels.embed.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GENKIT, + name: "genkit.embed", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params]) => extractEmbedInput(params), + extractOutput: (result) => summarizeEmbeddingResult(result), + extractMetadata: (_result, event) => + extractEmbedMetadata(event?.arguments?.[0]), + extractMetrics: () => ({}), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(genkitChannels.embedMany, { - name: "genkit.embedMany", - type: SpanTypeAttribute.FUNCTION, - extractInput: ([params]) => extractEmbedManyInput(params), - extractOutput: summarizeEmbeddingResult, - extractMetadata: (_result, event) => - extractEmbedMetadata(event?.arguments?.[0]), - extractMetrics: () => ({}), - }), + genkitChannels.embedMany.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GENKIT, + name: "genkit.embedMany", + type: SpanTypeAttribute.FUNCTION, + extractInput: ([params]) => extractEmbedManyInput(params), + extractOutput: summarizeEmbeddingResult, + extractMetadata: (_result, event) => + extractEmbedMetadata(event?.arguments?.[0]), + extractMetrics: () => ({}), + }, + ), + ), ); this.subscribeToActionRun(); @@ -100,126 +127,189 @@ export class GenkitPlugin extends BasePlugin { } private subscribeToActionRun(): void { - const tracingChannel = - genkitChannels.actionRun.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; + const invocationHook = genkitChannels.actionRun; const states = new WeakMap(); - const unbindCurrentSpanStore = bindActionCurrentSpanStoreToStart( - tracingChannel, - states, - (event) => startActionRunSpan(event), - ); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - ensureActionSpanState(states, event as object, () => - startActionRunSpan(event), + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const spanState = runInstrumentation( + () => + states.get(event) ?? ((event) => startActionRunSpan(event))(event), ); - }, - asyncEnd: (event) => { - const state = states.get(event); - if (!state) { - return; - } - - try { - state.span.log({ - output: extractActionOutput(event.result), - metrics: durationMetrics(state.startTime), - }); - } finally { + if (spanState) states.set(event, spanState); + const prepare = ( + event: ChannelMessage, + ) => { + ensureActionSpanState(states, event as object, () => + startActionRunSpan(event), + ); + }; + const resolved = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + + try { + state.span.log({ + output: extractActionOutput(event.result), + metrics: durationMetrics(state.startTime), + }); + } finally { + state.span.end(); + states.delete(event); + } + }; + const failed = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state || !event.error) { + return; + } + state.span.log({ error: event.error.message }); state.span.end(); states.delete(event); - } - }, - error: (event) => { - const state = states.get(event); - if (!state || !event.error) { - return; - } - state.span.log({ error: event.error.message }); - state.span.end(); - states.delete(event); + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } + + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); + }; + return spanState ? withCurrent(spanState.span, invoke) : invoke(); }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToActionSpan(): void { - const tracingChannel = - genkitCoreChannels.actionSpan.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; + const invocationHook = genkitCoreChannels.actionSpan; const states = new WeakMap(); - const unbindCurrentSpanStore = bindActionCurrentSpanStoreToStart( - tracingChannel, - states, - (event) => startActionSpan(event), - ); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - ensureActionSpanState(states, event as object, () => - startActionSpan(event), + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const spanState = runInstrumentation( + () => states.get(event) ?? ((event) => startActionSpan(event))(event), ); - }, - asyncEnd: (event) => { - const state = states.get(event as object); - if (!state) { - return; - } - - try { - state.span.log({ - input: extractActionSpanInput(event.arguments), - output: extractActionOutput(event.result), - metrics: durationMetrics(state.startTime), - }); - } finally { + if (spanState) states.set(event, spanState); + const prepare = ( + event: ChannelMessage, + ) => { + ensureActionSpanState(states, event as object, () => + startActionSpan(event), + ); + }; + const resolved = ( + event: ChannelMessage, + ) => { + const state = states.get(event as object); + if (!state) { + return; + } + + try { + state.span.log({ + input: extractActionSpanInput(event.arguments), + output: extractActionOutput(event.result), + metrics: durationMetrics(state.startTime), + }); + } finally { + state.span.end(); + states.delete(event as object); + } + }; + const failed = ( + event: ChannelMessage, + ) => { + const state = states.get(event as object); + if (!state || !event.error) { + return; + } + state.span.log({ error: event.error.message }); state.span.end(); states.delete(event as object); - } - }, - error: (event) => { - const state = states.get(event as object); - if (!state || !event.error) { - return; - } - state.span.log({ error: event.error.message }); - state.span.end(); - states.delete(event as object); + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } + + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); + }; + return spanState ? withCurrent(spanState.span, invoke) : invoke(); }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToActionStream(): void { this.unsubscribers.push( - traceSyncStreamChannel(genkitChannels.actionStream, { - name: "genkit.action.stream", - type: SpanTypeAttribute.TASK, - extractInput: ([input], event) => ({ - input, - metadata: actionMetadataForLog(extractActionMetadata(event.self)), - }), - patchResult: ({ result, span, startTime }) => - patchActionStreamResult(result, span, startTime), - }), + genkitChannels.actionStream.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GENKIT, + name: "genkit.action.stream", + type: SpanTypeAttribute.TASK, + extractInput: ([input], event) => ({ + input, + metadata: actionMetadataForLog( + extractActionMetadata(event.self), + ), + }), + patchResult: ({ result, span, startTime }) => + patchActionStreamResult(result, span, startTime), + }, + ), + ), ); } } @@ -299,52 +389,6 @@ function ensureActionSpanState( return created; } -function bindActionCurrentSpanStoreToStart< - TChannel extends - | typeof genkitChannels.actionRun - | typeof genkitCoreChannels.actionSpan, ->( - tracingChannel: IsoTracingChannel>, - states: WeakMap, - create: (event: ChannelMessage) => SpanState | undefined, -): (() => void) | undefined { - const state = _internalGetGlobalState(); - const contextManager = state?.contextManager; - const startChannel = tracingChannel.start as - | ({ - bindStore?: ( - store: CurrentSpanStore, - callback: (event: ChannelMessage) => unknown, - ) => void; - unbindStore?: (store: CurrentSpanStore) => void; - } & object) - | undefined; - const currentSpanStore = contextManager - ? ( - contextManager as { - [BRAINTRUST_CURRENT_SPAN_STORE]?: CurrentSpanStore; - } - )[BRAINTRUST_CURRENT_SPAN_STORE] - : undefined; - - if (!startChannel?.bindStore || !currentSpanStore) { - return undefined; - } - - startChannel.bindStore(currentSpanStore, (event) => { - const state = ensureActionSpanState(states, event as object, () => - create(event), - ); - return state - ? contextManager!.wrapSpanForStore(state.span) - : currentSpanStore.getStore(); - }); - - return () => { - startChannel.unbindStore?.(currentSpanStore); - }; -} - function normalizeInput(input: GenkitGenerateInput): GenkitGenerateInput { if (typeof input === "string" || Array.isArray(input)) { return { prompt: input }; diff --git a/js/src/instrumentation/plugins/github-copilot-channels.ts b/js/src/instrumentation/plugins/github-copilot-channels.ts index 94b5abde9..bc003b969 100644 --- a/js/src/instrumentation/plugins/github-copilot-channels.ts +++ b/js/src/instrumentation/plugins/github-copilot-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { GitHubCopilotAssistantMessageEvent, GitHubCopilotMessageOptions, @@ -8,27 +8,23 @@ import type { GitHubCopilotSessionConfig, } from "../../vendor-sdk-types/github-copilot"; -export const gitHubCopilotChannels = defineChannels( - "@github/copilot-sdk", - { - createSession: channel<[GitHubCopilotSessionConfig], GitHubCopilotSession>({ - channelName: "client.createSession", - kind: "async", - }), - resumeSession: channel< - [string, GitHubCopilotResumeSessionConfig], - GitHubCopilotSession - >({ - channelName: "client.resumeSession", - kind: "async", - }), - sendAndWait: channel< - [GitHubCopilotMessageOptions, number?], - GitHubCopilotAssistantMessageEvent | undefined - >({ - channelName: "session.sendAndWait", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.GITHUB_COPILOT }, -); +export const gitHubCopilotChannels = defineInterceptor("@github/copilot-sdk", { + createSession: channel< + [GitHubCopilotSessionConfig], + PromiseLike + >({ + channelName: "client.createSession", + }), + resumeSession: channel< + [string, GitHubCopilotResumeSessionConfig], + PromiseLike + >({ + channelName: "client.resumeSession", + }), + sendAndWait: channel< + [GitHubCopilotMessageOptions, number?], + PromiseLike + >({ + channelName: "session.sendAndWait", + }), +}); diff --git a/js/src/instrumentation/plugins/github-copilot-plugin.ts b/js/src/instrumentation/plugins/github-copilot-plugin.ts index 8fa156cfc..a2175b314 100644 --- a/js/src/instrumentation/plugins/github-copilot-plugin.ts +++ b/js/src/instrumentation/plugins/github-copilot-plugin.ts @@ -1,18 +1,10 @@ -import { BasePlugin } from "../core"; -import type { IsoChannelHandlers } from "../../isomorph"; -import { startSpan as startBaseSpan } from "../../logger"; +import { SpanTypeAttribute } from "../../../util/index"; import type { Span } from "../../logger"; +import { startSpan as startBaseSpan } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { SpanTypeAttribute } from "../../../util/index"; -import { - extractAnthropicCacheTokens, - finalizeAnthropicTokens, -} from "../../wrappers/anthropic-tokens-util"; -import type { AnthropicTokenMetrics } from "../../wrappers/anthropic-tokens-util"; -import { gitHubCopilotChannels } from "./github-copilot-channels"; import type { GitHubCopilotSession, GitHubCopilotSessionConfig, @@ -20,6 +12,14 @@ import type { GitHubCopilotTrackedEvent, GitHubCopilotUsageData, } from "../../vendor-sdk-types/github-copilot"; +import type { AnthropicTokenMetrics } from "../../wrappers/anthropic-tokens-util"; +import { + extractAnthropicCacheTokens, + finalizeAnthropicTokens, +} from "../../wrappers/anthropic-tokens-util"; +import { BasePlugin } from "../core"; +import { observeResult, runInstrumentation } from "../core/observe-result"; +import { gitHubCopilotChannels } from "./github-copilot-channels"; const ROOT_AGENT_KEY = "__root__"; @@ -676,84 +676,105 @@ function isGitHubCopilotSession(value: unknown): value is GitHubCopilotSession { // --------------------------------------------------------------------------- // eslint-disable-next-line @typescript-eslint/no-explicit-any -function makeSessionHandlers( - sessionStates: WeakMap, +function traceSessionCreation( + call: () => T, + args: unknown[], configArgIndex: number, includeProviderMetadata: boolean, -): IsoChannelHandlers { - return { - start: (event) => { - const config = event.arguments[configArgIndex] as - | GitHubCopilotSessionConfig - | undefined; - if (!config || typeof config !== "object") { - return; - } +): T { + const sessionStates = new WeakMap(); + const event: { arguments: unknown[]; result?: unknown; error?: Error } = { + arguments: args, + }; + const prepare = (context: typeof event) => { + const config = event.arguments[configArgIndex] as + | GitHubCopilotSessionConfig + | undefined; + if (!config || typeof config !== "object") { + return; + } - const sessionSpan = startBaseSpan( - withSpanInstrumentationName( - { - name: "Copilot Session", - spanAttributes: { type: SpanTypeAttribute.TASK }, - }, - INSTRUMENTATION_NAMES.GITHUB_COPILOT, - ), - ); + const sessionSpan = startBaseSpan( + withSpanInstrumentationName( + { + name: "Copilot Session", + spanAttributes: { type: SpanTypeAttribute.TASK }, + }, + INSTRUMENTATION_NAMES.GITHUB_COPILOT, + ), + ); - const metadata: Record = {}; - if (config.model) { - metadata["github_copilot.model"] = config.model; - } - if (includeProviderMetadata && config.provider?.type) { - metadata["github_copilot.provider_type"] = config.provider.type; - } - if (Object.keys(metadata).length > 0) { - sessionSpan.log({ metadata }); - } + const metadata: Record = {}; + if (config.model) { + metadata["github_copilot.model"] = config.model; + } + if (includeProviderMetadata && config.provider?.type) { + metadata["github_copilot.provider_type"] = config.provider.type; + } + if (Object.keys(metadata).length > 0) { + sessionSpan.log({ metadata }); + } - const state: SessionState = { - session: makeSpanWithId(sessionSpan), - activeTurns: new Map(), - pendingUserMessages: new Map(), - currentMessageContent: new Map(), - activeTools: new Map(), - subAgents: new Map(), - agentIdToToolCallId: new Map(), - processing: Promise.resolve(), - totalInputTokens: 0, - totalOutputTokens: 0, - }; + const state: SessionState = { + session: makeSpanWithId(sessionSpan), + activeTurns: new Map(), + pendingUserMessages: new Map(), + currentMessageContent: new Map(), + activeTools: new Map(), + subAgents: new Map(), + agentIdToToolCallId: new Map(), + processing: Promise.resolve(), + totalInputTokens: 0, + totalOutputTokens: 0, + }; - injectTracingHooks(config, state); - sessionStates.set(event, state); - }, + injectTracingHooks(config, state); + sessionStates.set(event, state); + }; + const complete = (context: typeof event) => { + const state = sessionStates.get(event); + if (!state) { + return; + } - asyncEnd: (event) => { - const state = sessionStates.get(event); - if (!state) { - return; - } + const session = event.result; + if (isGitHubCopilotSession(session)) { + attachSessionEventListener(session, state); + } else { + state.session.span.end(); + } + sessionStates.delete(event); + }; + const fail = (context: typeof event) => { + const state = sessionStates.get(event); + if (!state || !event.error) { + return; + } - const session = event.result; - if (isGitHubCopilotSession(session)) { - attachSessionEventListener(session, state); - } else { - state.session.span.end(); - } - sessionStates.delete(event); + state.session.span.log({ error: event.error.message }); + state.session.span.end(); + sessionStates.delete(event); + }; + runInstrumentation(() => prepare(event)); + let result: T; + try { + result = call(); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => fail(event)); + throw error; + } + return observeResult( + result, + (value) => { + event.result = value; + complete(event); }, - - error: (event) => { - const state = sessionStates.get(event); - if (!state || !event.error) { - return; - } - - state.session.span.log({ error: event.error.message }); - state.session.span.end(); - sessionStates.delete(event); + (error) => { + Object.assign(event, { error }); + fail(event); }, - }; + ); } export class GitHubCopilotPlugin extends BasePlugin { @@ -769,28 +790,23 @@ export class GitHubCopilotPlugin extends BasePlugin { } private subscribeToSessionChannels(): void { - const createChannel = gitHubCopilotChannels.createSession.tracingChannel(); - const resumeChannel = gitHubCopilotChannels.resumeSession.tracingChannel(); - - const sessionStates = new WeakMap(); - - const createHandlers = makeSessionHandlers( - sessionStates, - 0, // config is arg 0 of createSession(config) - true, // include provider metadata - ); - const resumeHandlers = makeSessionHandlers( - sessionStates, - 1, // config is arg 1 of resumeSession(sessionId, config) - false, // resumeSession config has no provider field - ); - - createChannel.subscribe(createHandlers); - resumeChannel.subscribe(resumeHandlers); - this.unsubscribers.push( - () => createChannel.unsubscribe(createHandlers), - () => resumeChannel.unsubscribe(resumeHandlers), + gitHubCopilotChannels.createSession.intercept((target, receiver, args) => + traceSessionCreation( + () => Reflect.apply(target, receiver, args), + args, + 0, + true, + ), + ), + gitHubCopilotChannels.resumeSession.intercept((target, receiver, args) => + traceSessionCreation( + () => Reflect.apply(target, receiver, args), + args, + 1, + false, + ), + ), ); } } diff --git a/js/src/instrumentation/plugins/google-adk-channels.ts b/js/src/instrumentation/plugins/google-adk-channels.ts index 9d66b91b2..1da89c466 100644 --- a/js/src/instrumentation/plugins/google-adk-channels.ts +++ b/js/src/instrumentation/plugins/google-adk-channels.ts @@ -1,8 +1,8 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { - GoogleADKRunAsyncParams, GoogleADKEvent, + GoogleADKRunAsyncParams, GoogleADKToolRunRequest, } from "../../vendor-sdk-types/google-adk"; @@ -20,37 +20,30 @@ type GoogleADKChannelContext = { * tool.runAsync is a regular async function returning Promise, * so it uses "async" kind. */ -export const googleADKChannels = defineChannels( - "@google/adk", - { - runnerRunAsync: channel< - [GoogleADKRunAsyncParams], - AsyncGenerator, - GoogleADKChannelContext, - GoogleADKEvent - >({ - channelName: "runner.runAsync", - kind: "sync-stream", - }), +export const googleADKChannels = defineInterceptor("@google/adk", { + runnerRunAsync: channel< + [GoogleADKRunAsyncParams], + AsyncGenerator, + GoogleADKChannelContext, + GoogleADKEvent + >({ + channelName: "runner.runAsync", + }), - agentRunAsync: channel< - [unknown], - AsyncGenerator, - GoogleADKChannelContext, - GoogleADKEvent - >({ - channelName: "agent.runAsync", - kind: "sync-stream", - }), + agentRunAsync: channel< + [unknown], + AsyncGenerator, + GoogleADKChannelContext, + GoogleADKEvent + >({ + channelName: "agent.runAsync", + }), - toolRunAsync: channel< - [GoogleADKToolRunRequest], - unknown, - GoogleADKChannelContext - >({ - channelName: "tool.runAsync", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.GOOGLE_ADK }, -); + toolRunAsync: channel< + [GoogleADKToolRunRequest], + PromiseLike, + GoogleADKChannelContext + >({ + channelName: "tool.runAsync", + }), +}); diff --git a/js/src/instrumentation/plugins/google-adk-plugin.test.ts b/js/src/instrumentation/plugins/google-adk-plugin.test.ts index 073af75f4..fa20490c2 100644 --- a/js/src/instrumentation/plugins/google-adk-plugin.test.ts +++ b/js/src/instrumentation/plugins/google-adk-plugin.test.ts @@ -1,24 +1,32 @@ -import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +import { invocationController } from "../test-utils/invocation"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const { mockCurrentSpanStoreSymbol: MOCK_CURRENT_SPAN_STORE_SYMBOL, mockInternalGetGlobalState, + mockWithCurrent, } = vi.hoisted(() => ({ + mockWithCurrent: vi.fn((_span: any, callback: () => unknown) => callback()), mockCurrentSpanStoreSymbol: Symbol.for("braintrust.currentSpanStore"), mockInternalGetGlobalState: vi.fn(() => undefined), })); -// Mock iso's newTracingChannel — must be before any imports that use it vi.mock("../../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(), - }, + default: {}, })); import { GoogleADKPlugin } from "./google-adk-plugin"; -import iso from "../../isomorph"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; // Mock logger const mockStartSpan = vi.fn(() => ({ @@ -32,7 +40,7 @@ vi.mock("../../logger", () => ({ _internalGetGlobalState: (...args: any[]) => (mockInternalGetGlobalState as any)(...args), BRAINTRUST_CURRENT_SPAN_STORE: MOCK_CURRENT_SPAN_STORE_SYMBOL, - withCurrent: (_span: any, callback: () => unknown) => callback(), + withCurrent: (...args: any[]) => (mockWithCurrent as any)(...args), Attachment: class MockAttachment { reference: any; constructor(params: any) { @@ -58,7 +66,10 @@ describe("GoogleADKPlugin", () => { bindStoreSpy = vi.fn(); unbindStoreSpy = vi.fn(); mockChannel = { - subscribe: subscribeSpy, + intercept: (interceptor: any) => { + subscribeSpy(invocationController(interceptor)); + return unsubscribeSpy; + }, unsubscribe: unsubscribeSpy, hasSubscribers: false, start: { @@ -67,7 +78,7 @@ describe("GoogleADKPlugin", () => { }, }; - mockNewTracingChannel.mockReturnValue(mockChannel); + mockNewInvocationHook.mockReturnValue(mockChannel); mockStartSpan.mockClear(); mockInternalGetGlobalState.mockReset(); mockInternalGetGlobalState.mockReturnValue(undefined); @@ -83,13 +94,13 @@ describe("GoogleADKPlugin", () => { plugin.enable(); // Should subscribe to 3 channels: runner.runAsync, agent.runAsync, tool.runAsync - expect(mockNewTracingChannel).toHaveBeenCalledWith( + expect(mockNewInvocationHook).toHaveBeenCalledWith( "orchestrion:@google/adk:runner.runAsync", ); - expect(mockNewTracingChannel).toHaveBeenCalledWith( + expect(mockNewInvocationHook).toHaveBeenCalledWith( "orchestrion:@google/adk:agent.runAsync", ); - expect(mockNewTracingChannel).toHaveBeenCalledWith( + expect(mockNewInvocationHook).toHaveBeenCalledWith( "orchestrion:@google/adk:tool.runAsync", ); expect(subscribeSpy).toHaveBeenCalledTimes(3); @@ -134,9 +145,9 @@ describe("GoogleADKPlugin", () => { // Find the first subscribe call (runner channel) const handlers = subscribeSpy.mock.calls[0][0]; - expect(handlers).toHaveProperty("start"); - expect(handlers).toHaveProperty("end"); - expect(handlers).toHaveProperty("error"); + expect(handlers).toHaveProperty("begin"); + expect(handlers).toHaveProperty("call"); + expect(handlers).toHaveProperty("reject"); // Simulate a start event const event = { @@ -152,7 +163,7 @@ describe("GoogleADKPlugin", () => { ], }; - handlers.start(event); + handlers.call(event); expect(mockStartSpan).toHaveBeenCalledWith( expect.objectContaining({ @@ -177,8 +188,6 @@ describe("GoogleADKPlugin", () => { result: undefined, }; - handlers.start(event); - // Simulate async iterable result const mockAsyncIterable = { [Symbol.asyncIterator]: () => ({ @@ -189,7 +198,7 @@ describe("GoogleADKPlugin", () => { }; event.result = mockAsyncIterable; - handlers.end(event); + handlers.call(event); // The stream should be patched — the span shouldn't end immediately // (it will end when the stream completes) @@ -209,34 +218,16 @@ describe("GoogleADKPlugin", () => { error: new Error("Runner failed"), }; - handlers.start(event); - + handlers.throw(event); const span = mockStartSpan.mock.results[0].value; - handlers.error(event); expect(span.log).toHaveBeenCalledWith({ error: "Runner failed" }); expect(span.end).toHaveBeenCalled(); }); - it("binds the current span store for runner events without creating duplicate spans", () => { - const currentSpanStore = {}; - const wrapSpanForStore = vi.fn(() => "wrapped-runner-store"); - mockInternalGetGlobalState.mockReturnValue({ - contextManager: { - [MOCK_CURRENT_SPAN_STORE_SYMBOL]: currentSpanStore, - wrapSpanForStore, - }, - } as any); - + it("runs runner calls with the current span without creating duplicate spans", () => { plugin.enable(); - expect(bindStoreSpy).toHaveBeenNthCalledWith( - 1, - currentSpanStore, - expect.any(Function), - ); - - const bindTransform = bindStoreSpy.mock.calls[0][1]; const handlers = subscribeSpy.mock.calls[0][0]; const event = { arguments: [ @@ -251,13 +242,12 @@ describe("GoogleADKPlugin", () => { ], }; - expect(bindTransform(event)).toBe("wrapped-runner-store"); - expect(wrapSpanForStore).toHaveBeenCalledWith( + handlers.call(event); + expect(mockWithCurrent).toHaveBeenCalledWith( mockStartSpan.mock.results[0].value, + expect.any(Function), ); - handlers.start(event); - expect(mockStartSpan).toHaveBeenCalledTimes(1); }); @@ -374,16 +364,16 @@ describe("GoogleADKPlugin", () => { arguments: [{ userId: "user-123", sessionId: "session-456" }], }; - handlers.start(event); - const span = mockStartSpan.mock.results.at(-1)?.value as { - log: ReturnType; - }; event.result = (async function* () { for (const usage of usageMetadata) { yield { usageMetadata: usage }; } })(); - handlers.end(event); + handlers.call(event); + + const span = mockStartSpan.mock.results.at(-1)?.value as { + log: ReturnType; + }; for await (const _event of event.result) { // Consume the runner stream so the span is finalized. @@ -405,10 +395,6 @@ describe("GoogleADKPlugin", () => { arguments: [{ userId: "user-123", sessionId: "session-456" }], }; - handlers.start(event); - const span = mockStartSpan.mock.results.at(-1)?.value as { - log: ReturnType; - }; event.result = (async function* () { yield { usageMetadata: { @@ -426,7 +412,11 @@ describe("GoogleADKPlugin", () => { }, }; })(); - handlers.end(event); + handlers.call(event); + + const span = mockStartSpan.mock.results.at(-1)?.value as { + log: ReturnType; + }; for await (const _event of event.result) { // Consume the runner stream so the span is finalized. @@ -463,7 +453,7 @@ describe("GoogleADKPlugin", () => { ], }; - handlers.start(event); + handlers.call(event); expect(mockStartSpan).toHaveBeenCalledWith( expect.objectContaining({ @@ -492,7 +482,7 @@ describe("GoogleADKPlugin", () => { }, }; - handlers.start(event); + handlers.call(event); const span = mockStartSpan.mock.results[0].value; expect(mockStartSpan).toHaveBeenCalledWith( @@ -518,7 +508,7 @@ describe("GoogleADKPlugin", () => { arguments: [undefined], }; - handlers.start(event); + handlers.call(event); expect(mockStartSpan).toHaveBeenCalledWith( expect.objectContaining({ @@ -527,25 +517,9 @@ describe("GoogleADKPlugin", () => { ); }); - it("binds the current span store for agent events without creating duplicate spans", () => { - const currentSpanStore = {}; - const wrapSpanForStore = vi.fn(() => "wrapped-agent-store"); - mockInternalGetGlobalState.mockReturnValue({ - contextManager: { - [MOCK_CURRENT_SPAN_STORE_SYMBOL]: currentSpanStore, - wrapSpanForStore, - }, - } as any); - + it("runs agent calls with the current span without creating duplicate spans", () => { plugin.enable(); - expect(bindStoreSpy).toHaveBeenNthCalledWith( - 2, - currentSpanStore, - expect.any(Function), - ); - - const bindTransform = bindStoreSpy.mock.calls[1][1]; const handlers = subscribeSpy.mock.calls[1][0]; const event = { arguments: [ @@ -558,13 +532,12 @@ describe("GoogleADKPlugin", () => { ], }; - expect(bindTransform(event)).toBe("wrapped-agent-store"); - expect(wrapSpanForStore).toHaveBeenCalledWith( + handlers.call(event); + expect(mockWithCurrent).toHaveBeenCalledWith( mockStartSpan.mock.results[0].value, + expect.any(Function), ); - handlers.start(event); - expect(mockStartSpan).toHaveBeenCalledTimes(1); }); @@ -593,35 +566,37 @@ describe("GoogleADKPlugin", () => { const runnerHandlers = subscribeSpy.mock.calls[0][0]; const agentHandlers = subscribeSpy.mock.calls[1][0]; - runnerHandlers.start({ - arguments: [ - { - session: { - id: "session-456", - userId: "user-123", - }, - userContent: { - role: "user", - parts: [{ text: "What is the weather?" }], - }, - }, - ], - }); - - agentHandlers.start({ - arguments: [ - { - session: { - id: "session-456", - userId: "user-123", - }, - agent: { - name: "weather_agent", - model: "gemini-2.5-flash", + runnerHandlers.call( + { + arguments: [ + { + session: { + id: "session-456", + userId: "user-123", + }, + userContent: { + role: "user", + parts: [{ text: "What is the weather?" }], + }, }, - }, - ], - }); + ], + }, + () => + agentHandlers.call({ + arguments: [ + { + session: { + id: "session-456", + userId: "user-123", + }, + agent: { + name: "weather_agent", + model: "gemini-2.5-flash", + }, + }, + ], + }), + ); expect(mockStartSpan).toHaveBeenNthCalledWith( 2, @@ -657,7 +632,7 @@ describe("GoogleADKPlugin", () => { }, }; - handlers.start(event); + handlers.begin(event); expect(mockStartSpan).toHaveBeenCalledWith( expect.objectContaining({ @@ -687,10 +662,10 @@ describe("GoogleADKPlugin", () => { result: { temperature: 72, condition: "sunny" }, }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results[0].value; - handlers.asyncEnd(event); + handlers.resolve(event); expect(span.log).toHaveBeenCalledWith( expect.objectContaining({ @@ -714,10 +689,10 @@ describe("GoogleADKPlugin", () => { error: new Error("Tool failed"), }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results[0].value; - handlers.error(event); + handlers.reject(event); expect(span.log).toHaveBeenCalledWith({ error: "Tool failed" }); expect(span.end).toHaveBeenCalled(); diff --git a/js/src/instrumentation/plugins/google-adk-plugin.ts b/js/src/instrumentation/plugins/google-adk-plugin.ts index 623264016..b516bfd9d 100644 --- a/js/src/instrumentation/plugins/google-adk-plugin.ts +++ b/js/src/instrumentation/plugins/google-adk-plugin.ts @@ -1,30 +1,26 @@ import { BasePlugin } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; -import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; -import { - BRAINTRUST_CURRENT_SPAN_STORE, - _internalGetGlobalState, - startSpan as startBaseSpan, - withCurrent, -} from "../../logger"; -import type { CurrentSpanStore, Span } from "../../logger"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute } from "../../../util/index"; +import type { Span } from "../../logger"; +import { startSpan as startBaseSpan, withCurrent } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { SpanTypeAttribute } from "../../../util/index"; import { getCurrentUnixTimestamp } from "../../util"; -import { googleADKChannels } from "./google-adk-channels"; import type { + GoogleADKBaseAgent, + GoogleADKBaseTool, GoogleADKEvent, + GoogleADKLlmAgent, GoogleADKRunAsyncParams, GoogleADKToolRunRequest, GoogleADKUsageMetadata, - GoogleADKBaseAgent, - GoogleADKLlmAgent, - GoogleADKBaseTool, } from "../../vendor-sdk-types/google-adk"; +import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; +import type { ChannelMessage } from "../core/tracing-types"; +import { googleADKChannels } from "./google-adk-channels"; type RunnerState = { span: Span; @@ -46,10 +42,6 @@ type ToolState = { startTime: number; }; -type GoogleADKStreamChannel = - | typeof googleADKChannels.runnerRunAsync - | typeof googleADKChannels.agentRunAsync; - /** * Auto-instrumentation plugin for the Google ADK. * @@ -82,10 +74,7 @@ export class GoogleADKPlugin extends BasePlugin { } private subscribeToRunnerRunAsync(): void { - const tracingChannel = - googleADKChannels.runnerRunAsync.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; + const invocationHook = googleADKChannels.runnerRunAsync; const states = new WeakMap(); const createState = ( @@ -126,80 +115,94 @@ export class GoogleADKPlugin extends BasePlugin { return { span, startTime, events: [], contextKey }; }; - const unbindCurrentSpanStore = bindCurrentSpanStoreToStart( - tracingChannel, - states, - createState, - ); - - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - ensureState(states, event, () => createState(event)); - }, - - end: (event) => { - const state = states.get(event); - if (!state) { - return; - } - - const result = event.result; - if (isAsyncIterable(result)) { - bindAsyncIterableToCurrentSpan(result, state.span); - patchStreamIfNeeded(result, { - onChunk: (adkEvent: GoogleADKEvent) => { - state.events.push(adkEvent); - }, - onComplete: () => { - finalizeRunnerSpan(state, this.activeRunnerSpans); - states.delete(event); - }, - onError: (error: Error) => { - cleanupActiveRunnerSpan(state, this.activeRunnerSpans); - state.span.log({ error: error.message }); - state.span.end(); - states.delete(event); - }, - }); - return; - } - - // Non-streaming case (unlikely for runners but handle gracefully) - try { - state.span.log({ output: result }); - } finally { + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const spanState = runInstrumentation( + () => states.get(event) ?? createState(event), + ); + if (spanState) states.set(event, spanState); + const prepare = ( + event: ChannelMessage, + ) => { + ensureState(states, event, () => createState(event)); + }; + const returned = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + + const result = event.result; + if (isAsyncIterable(result)) { + bindAsyncIterableToCurrentSpan(result, state.span); + patchStreamIfNeeded(result, { + onChunk: (adkEvent: GoogleADKEvent) => { + state.events.push(adkEvent); + }, + onComplete: () => { + finalizeRunnerSpan(state, this.activeRunnerSpans); + states.delete(event); + }, + onError: (error: Error) => { + cleanupActiveRunnerSpan(state, this.activeRunnerSpans); + state.span.log({ error: error.message }); + state.span.end(); + states.delete(event); + }, + }); + return; + } + + // Non-streaming case (unlikely for runners but handle gracefully) + try { + state.span.log({ output: result }); + } finally { + cleanupActiveRunnerSpan(state, this.activeRunnerSpans); + state.span.end(); + states.delete(event); + } + }; + const failed = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state || !event.error) { + return; + } cleanupActiveRunnerSpan(state, this.activeRunnerSpans); + state.span.log({ error: event.error.message }); state.span.end(); states.delete(event); - } - }, - - error: (event) => { - const state = states.get(event); - if (!state || !event.error) { - return; - } - cleanupActiveRunnerSpan(state, this.activeRunnerSpans); - state.span.log({ error: event.error.message }); - state.span.end(); - states.delete(event); + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } + Object.assign(event, { result }); + runInstrumentation(() => returned(event)); + return result; + }; + return spanState ? withCurrent(spanState.span, invoke) : invoke(); }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToAgentRunAsync(): void { - const tracingChannel = - googleADKChannels.agentRunAsync.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; + const invocationHook = googleADKChannels.agentRunAsync; const states = new WeakMap(); const createState = ( @@ -262,160 +265,202 @@ export class GoogleADKPlugin extends BasePlugin { return { span, startTime, events: [], contextKey, name: agentName }; }; - const unbindCurrentSpanStore = bindCurrentSpanStoreToStart( - tracingChannel, - states, - createState, - ); - - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - ensureState(states, event, () => createState(event)); - }, - - end: (event) => { - const state = states.get(event); - if (!state) { - return; - } - - const result = event.result; - if (isAsyncIterable(result)) { - bindAsyncIterableToCurrentSpan(result, state.span); - patchStreamIfNeeded(result, { - onChunk: (adkEvent: GoogleADKEvent) => { - state.events.push(adkEvent); - }, - onComplete: () => { - finalizeAgentSpan(state, this.activeAgentSpans); - states.delete(event); - }, - onError: (error: Error) => { - cleanupActiveAgentSpan(state, this.activeAgentSpans); - state.span.log({ error: error.message }); - state.span.end(); - states.delete(event); - }, - }); - return; - } - - try { - state.span.log({ output: result }); - } finally { + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const spanState = runInstrumentation( + () => states.get(event) ?? createState(event), + ); + if (spanState) states.set(event, spanState); + const prepare = ( + event: ChannelMessage, + ) => { + ensureState(states, event, () => createState(event)); + }; + const returned = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + + const result = event.result; + if (isAsyncIterable(result)) { + bindAsyncIterableToCurrentSpan(result, state.span); + patchStreamIfNeeded(result, { + onChunk: (adkEvent: GoogleADKEvent) => { + state.events.push(adkEvent); + }, + onComplete: () => { + finalizeAgentSpan(state, this.activeAgentSpans); + states.delete(event); + }, + onError: (error: Error) => { + cleanupActiveAgentSpan(state, this.activeAgentSpans); + state.span.log({ error: error.message }); + state.span.end(); + states.delete(event); + }, + }); + return; + } + + try { + state.span.log({ output: result }); + } finally { + cleanupActiveAgentSpan(state, this.activeAgentSpans); + state.span.end(); + states.delete(event); + } + }; + const failed = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state || !event.error) { + return; + } cleanupActiveAgentSpan(state, this.activeAgentSpans); + state.span.log({ error: event.error.message }); state.span.end(); states.delete(event); - } - }, - - error: (event) => { - const state = states.get(event); - if (!state || !event.error) { - return; - } - cleanupActiveAgentSpan(state, this.activeAgentSpans); - state.span.log({ error: event.error.message }); - state.span.end(); - states.delete(event); + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } + Object.assign(event, { result }); + runInstrumentation(() => returned(event)); + return result; + }; + return spanState ? withCurrent(spanState.span, invoke) : invoke(); }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToToolRunAsync(): void { - const tracingChannel = googleADKChannels.toolRunAsync.tracingChannel(); + const invocationHook = googleADKChannels.toolRunAsync; const states = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - const req = (event.arguments[0] ?? {}) as GoogleADKToolRunRequest; - const tool = event.self as GoogleADKBaseTool | undefined; - - const toolName = extractToolName(req, tool); - const parentSpan = findToolParentSpan( - req, - this.activeAgentSpans, - this.activeRunnerSpans, - ); + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + const req = (event.arguments[0] ?? {}) as GoogleADKToolRunRequest; + const tool = event.self as GoogleADKBaseTool | undefined; + + const toolName = extractToolName(req, tool); + const parentSpan = findToolParentSpan( + req, + this.activeAgentSpans, + this.activeRunnerSpans, + ); - const createSpan = () => - startBaseSpan( - withSpanInstrumentationName( - { - name: toolName ? `tool: ${toolName}` : "Google ADK Tool", - spanAttributes: { - type: SpanTypeAttribute.TOOL, - }, - event: { - input: req.args, - metadata: { - provider: "google-adk", - ...(toolName && { "google_adk.tool_name": toolName }), - ...(extractToolCallId(req) && { - "google_adk.tool_call_id": extractToolCallId(req), - }), + const createSpan = () => + startBaseSpan( + withSpanInstrumentationName( + { + name: toolName ? `tool: ${toolName}` : "Google ADK Tool", + spanAttributes: { + type: SpanTypeAttribute.TOOL, + }, + event: { + input: req.args, + metadata: { + provider: "google-adk", + ...(toolName && { "google_adk.tool_name": toolName }), + ...(extractToolCallId(req) && { + "google_adk.tool_call_id": extractToolCallId(req), + }), + }, }, }, - }, - INSTRUMENTATION_NAMES.GOOGLE_ADK, - ), - ); - const span = parentSpan - ? withCurrent(parentSpan, () => createSpan()) - : createSpan(); - const startTime = getCurrentUnixTimestamp(); - - states.set(event, { span, startTime }); - }, - - asyncEnd: (event) => { - const state = states.get(event); - if (!state) { - return; - } - - try { - const metrics: Record = {}; - const end = getCurrentUnixTimestamp(); - metrics.start = state.startTime; - metrics.end = end; - metrics.duration = end - state.startTime; - - state.span.log({ - output: event.result, - metrics: cleanMetrics(metrics), - }); - } finally { + INSTRUMENTATION_NAMES.GOOGLE_ADK, + ), + ); + const span = parentSpan + ? withCurrent(parentSpan, () => createSpan()) + : createSpan(); + const startTime = getCurrentUnixTimestamp(); + + states.set(event, { span, startTime }); + }; + const resolved = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + + try { + const metrics: Record = {}; + const end = getCurrentUnixTimestamp(); + metrics.start = state.startTime; + metrics.end = end; + metrics.duration = end - state.startTime; + + state.span.log({ + output: event.result, + metrics: cleanMetrics(metrics), + }); + } finally { + state.span.end(); + states.delete(event); + } + }; + const failed = ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state || !event.error) { + return; + } + state.span.log({ error: event.error.message }); state.span.end(); states.delete(event); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); }, - - error: (event) => { - const state = states.get(event); - if (!state || !event.error) { - return; - } - state.span.log({ error: event.error.message }); - state.span.end(); - states.delete(event); - }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } } @@ -529,49 +574,6 @@ function bindAsyncIterableToCurrentSpan(stream: unknown, span: Span): unknown { return stream; } -function bindCurrentSpanStoreToStart< - TChannel extends GoogleADKStreamChannel, - TState extends { span: Span }, ->( - tracingChannel: IsoTracingChannel>, - states: WeakMap, - create: (event: ChannelMessage) => TState, -): (() => void) | undefined { - const state = _internalGetGlobalState(); - const contextManager = state?.contextManager; - const startChannel = tracingChannel.start as - | ({ - bindStore?: ( - store: CurrentSpanStore, - callback: (event: ChannelMessage) => unknown, - ) => void; - unbindStore?: (store: CurrentSpanStore) => void; - } & object) - | undefined; - const currentSpanStore = contextManager - ? ( - contextManager as { - [BRAINTRUST_CURRENT_SPAN_STORE]?: CurrentSpanStore; - } - )[BRAINTRUST_CURRENT_SPAN_STORE] - : undefined; - - if (!startChannel?.bindStore || !currentSpanStore) { - return undefined; - } - - startChannel.bindStore(currentSpanStore, (event) => { - const span = ensureState(states, event as object, () => - create(event as ChannelMessage), - ).span; - return contextManager.wrapSpanForStore(span); - }); - - return () => { - startChannel.unbindStore?.(currentSpanStore); - }; -} - // ---- Helper functions ---- function extractRunnerContextKey( diff --git a/js/src/instrumentation/plugins/google-genai-channels.ts b/js/src/instrumentation/plugins/google-genai-channels.ts index 2003f472f..ceb8c4ce8 100644 --- a/js/src/instrumentation/plugins/google-genai-channels.ts +++ b/js/src/instrumentation/plugins/google-genai-channels.ts @@ -1,10 +1,10 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { - GoogleGenAIEmbedContentParams, - GoogleGenAIEmbedContentResponse, GoogleGenAIEditImageParams, GoogleGenAIEditImageResponse, + GoogleGenAIEmbedContentParams, + GoogleGenAIEmbedContentResponse, GoogleGenAIGenerateContentParams, GoogleGenAIGenerateContentResponse, GoogleGenAIGenerateImagesParams, @@ -23,66 +23,54 @@ type GoogleGenAIInteractionResult = | GoogleGenAIInteraction | AsyncIterable; -export const googleGenAIChannels = defineChannels( - "@google/genai", - { - generateContent: channel< - [GoogleGenAIGenerateContentParams], - GoogleGenAIGenerateContentResponse - >({ - channelName: "models.generateContent", - kind: "async", - }), - generateContentStream: channel< - [GoogleGenAIGenerateContentParams], - GoogleGenAIStreamingResult, - Record, - GoogleGenAIGenerateContentResponse - >({ - channelName: "models.generateContentStream", - kind: "async", - }), - embedContent: channel< - [GoogleGenAIEmbedContentParams], - GoogleGenAIEmbedContentResponse - >({ - channelName: "models.embedContent", - kind: "async", - }), - generateImages: channel< - [GoogleGenAIGenerateImagesParams], - GoogleGenAIGenerateImagesResponse - >({ - channelName: "models.generateImages", - kind: "async", - }), - editImage: channel< - [GoogleGenAIEditImageParams], - GoogleGenAIEditImageResponse - >({ - channelName: "models.editImage", - kind: "async", - }), - generateVideos: channel< - [GoogleGenAIGenerateVideosParams], - GoogleGenAIGenerateVideosOperation - >({ - channelName: "models.generateVideos", - kind: "async", - }), - httpResponseJson: channel<[], GoogleGenAIEmbedContentResponse>({ - channelName: "httpResponse.json", - kind: "async", - }), - interactionsCreate: channel< - [GoogleGenAIInteractionCreateParams, Record?], - GoogleGenAIInteractionResult, - Record, - GoogleGenAIInteractionSSEEvent - >({ - channelName: "interactions.create", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.GOOGLE_GENAI }, -); +export const googleGenAIChannels = defineInterceptor("@google/genai", { + generateContent: channel< + [GoogleGenAIGenerateContentParams], + PromiseLike + >({ + channelName: "models.generateContent", + }), + generateContentStream: channel< + [GoogleGenAIGenerateContentParams], + PromiseLike, + Record, + GoogleGenAIGenerateContentResponse + >({ + channelName: "models.generateContentStream", + }), + embedContent: channel< + [GoogleGenAIEmbedContentParams], + PromiseLike + >({ + channelName: "models.embedContent", + }), + generateImages: channel< + [GoogleGenAIGenerateImagesParams], + PromiseLike + >({ + channelName: "models.generateImages", + }), + editImage: channel< + [GoogleGenAIEditImageParams], + PromiseLike + >({ + channelName: "models.editImage", + }), + generateVideos: channel< + [GoogleGenAIGenerateVideosParams], + PromiseLike + >({ + channelName: "models.generateVideos", + }), + httpResponseJson: channel<[], PromiseLike>({ + channelName: "httpResponse.json", + }), + interactionsCreate: channel< + [GoogleGenAIInteractionCreateParams, Record?], + PromiseLike, + Record, + GoogleGenAIInteractionSSEEvent + >({ + channelName: "interactions.create", + }), +}); diff --git a/js/src/instrumentation/plugins/google-genai-plugin.test.ts b/js/src/instrumentation/plugins/google-genai-plugin.test.ts index 34f8ac3da..4ef608163 100644 --- a/js/src/instrumentation/plugins/google-genai-plugin.test.ts +++ b/js/src/instrumentation/plugins/google-genai-plugin.test.ts @@ -1,8 +1,17 @@ -import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +import { invocationController } from "../test-utils/invocation"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); -// Mock iso's newTracingChannel - must be before any imports that use it +// Mock platform context independently of invocation hooks. vi.mock("../../isomorph", () => ({ default: { + getEnv: () => undefined, newAsyncLocalStorage: vi.fn(() => { let current: unknown; return { @@ -18,15 +27,15 @@ vi.mock("../../isomorph", () => ({ }), }; }), - newTracingChannel: vi.fn(), }, })); -import { GoogleGenAIPlugin } from "./google-genai-plugin"; import { startSpan } from "../../logger"; -import iso from "../../isomorph"; +import { GoogleGenAIPlugin } from "./google-genai-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; const mockStartSpan = vi.mocked(startSpan); // Mock logger @@ -70,7 +79,28 @@ describe("GoogleGenAIPlugin", () => { hasSubscribers: false, }; - mockNewTracingChannel.mockReturnValue(mockChannel); + mockNewInvocationHook.mockImplementation((name: string) => { + mockChannel = { + intercept: (interceptor: any) => { + if ( + [ + "models.generateContent", + "models.generateContentStream", + "models.embedContent", + "interactions.create", + ].some( + (operation) => name === `orchestrion:@google/genai:${operation}`, + ) + ) { + subscribeSpy(invocationController(interceptor)); + } else { + interceptSpy(interceptor); + } + return unsubscribeSpy; + }, + }; + return mockChannel; + }); plugin = new GoogleGenAIPlugin(); }); @@ -116,20 +146,13 @@ describe("GoogleGenAIPlugin", () => { it("should extract input correctly", () => { plugin.enable(); - const subscribeCall = subscribeSpy.mock.calls.find( - (call: any) => - mockNewTracingChannel.mock.results[ - subscribeSpy.mock.calls.indexOf(call) - ]?.value === mockChannel, - ); - - expect(subscribeCall).toBeDefined(); + expect(subscribeSpy).toHaveBeenCalled(); // Get the handlers from the subscribe call const handlers = subscribeSpy.mock.calls[0][0]; - expect(handlers).toHaveProperty("start"); - expect(handlers).toHaveProperty("asyncEnd"); - expect(handlers).toHaveProperty("error"); + expect(handlers).toHaveProperty("begin"); + expect(handlers).toHaveProperty("resolve"); + expect(handlers).toHaveProperty("reject"); }); it.each([ @@ -216,12 +239,12 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { log: ReturnType; }; event.result = { usageMetadata }; - handlers.asyncEnd(event); + handlers.resolve(event); const metrics = span.log.mock.calls[0][0].metrics; expect(metrics).toMatchObject(expectedMetrics); @@ -244,7 +267,7 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { log: ReturnType; }; @@ -263,7 +286,7 @@ describe("GoogleGenAIPlugin", () => { totalTokenCount: 0, }, }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(span.log.mock.calls[0][0].metrics).toMatchObject({ completion_audio_tokens: 0, @@ -289,7 +312,7 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { log: ReturnType; }; @@ -310,7 +333,7 @@ describe("GoogleGenAIPlugin", () => { }, ], }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(span.log).toHaveBeenCalledWith( expect.objectContaining({ @@ -351,7 +374,7 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { log: ReturnType; }; @@ -372,7 +395,7 @@ describe("GoogleGenAIPlugin", () => { }, ], }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(span.log).toHaveBeenCalledWith( expect.objectContaining({ @@ -457,9 +480,9 @@ describe("GoogleGenAIPlugin", () => { }; } - handlers.start(event); + handlers.begin(event); event.result = stream(); - handlers.asyncEnd(event); + handlers.resolve(event); for await (const _chunk of event.result) { // Consume the provider stream so the instrumentation finalizes it. } @@ -932,7 +955,7 @@ describe("GoogleGenAIPlugin", () => { (model, contents, inputs) => { plugin.enable(); const handlers = subscribeSpy.mock.calls[2][0]; - handlers.start({ + handlers.begin({ arguments: [ { model, @@ -988,7 +1011,7 @@ describe("GoogleGenAIPlugin", () => { ], }; const original = structuredClone(params); - handlers.start({ arguments: [params] }); + handlers.begin({ arguments: [params] }); const input = mockStartSpan.mock.calls[0][0]?.event?.input; expect(input).toMatchObject({ inputs: [ @@ -1024,7 +1047,7 @@ describe("GoogleGenAIPlugin", () => { it("retains all inline media when any attachment conversion fails", () => { plugin.enable(); const handlers = subscribeSpy.mock.calls[2][0]; - handlers.start({ + handlers.begin({ arguments: [ { model: "gemini-embedding-2-preview", @@ -1095,8 +1118,8 @@ describe("GoogleGenAIPlugin", () => { ], result, }; - handlers.start(event); - handlers.asyncEnd(event); + handlers.begin(event); + handlers.resolve(event); const span = mockStartSpan.mock.results[0].value; expect(span.log).toHaveBeenCalledWith({ output: { count }, @@ -1118,8 +1141,8 @@ describe("GoogleGenAIPlugin", () => { arguments: [{ model: "gemini-embedding-2-preview", contents: "hello" }], error, }; - handlers.start(event); - handlers.error(event); + handlers.begin(event); + handlers.reject(event); const span = mockStartSpan.mock.results[0].value; expect(span.log).toHaveBeenCalledWith({ error, output: { count: 0 } }); expect(span.end).toHaveBeenCalledOnce(); @@ -1130,7 +1153,7 @@ describe("GoogleGenAIPlugin", () => { it("subscribes to the interactions.create channel", () => { plugin.enable(); - expect(mockNewTracingChannel).toHaveBeenCalledWith( + expect(mockNewInvocationHook).toHaveBeenCalledWith( "orchestrion:@google/genai:interactions.create", ); expect(subscribeSpy).toHaveBeenCalledTimes(4); @@ -1164,7 +1187,7 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { end: ReturnType; log: ReturnType; @@ -1198,7 +1221,7 @@ describe("GoogleGenAIPlugin", () => { }, }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(span.log).toHaveBeenNthCalledWith( 1, @@ -1292,7 +1315,7 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { end: ReturnType; log: ReturnType; @@ -1306,7 +1329,7 @@ describe("GoogleGenAIPlugin", () => { }, status: "completed", }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(mockStartSpan).toHaveBeenLastCalledWith( expect.objectContaining({ @@ -1378,7 +1401,7 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { log: ReturnType; }; @@ -1399,7 +1422,7 @@ describe("GoogleGenAIPlugin", () => { }, }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(span.log).toHaveBeenLastCalledWith( expect.objectContaining({ @@ -1416,7 +1439,7 @@ describe("GoogleGenAIPlugin", () => { }), ); - handlers.start(event); + handlers.begin(event); const missingUsageSpan = mockStartSpan.mock.results.at(-1)?.value as { log: ReturnType; }; @@ -1426,7 +1449,7 @@ describe("GoogleGenAIPlugin", () => { usage: {}, }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(missingUsageSpan.log).toHaveBeenLastCalledWith( expect.objectContaining({ @@ -1455,12 +1478,12 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); event.result = { id: "interaction-background", status: "in_progress", }; - handlers.asyncEnd(event); + handlers.resolve(event); expect(mockStartSpan).not.toHaveBeenCalled(); }); @@ -1513,13 +1536,13 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { end: ReturnType; log: ReturnType; }; event.result = stream(); - handlers.asyncEnd(event); + handlers.resolve(event); for await (const _chunk of event.result) { // Consume the stream so aggregation completes. @@ -1575,13 +1598,13 @@ describe("GoogleGenAIPlugin", () => { ], }; - handlers.start(event); + handlers.begin(event); const span = mockStartSpan.mock.results.at(-1)?.value as { end: ReturnType; log: ReturnType; }; event.result = stream(); - handlers.asyncEnd(event); + handlers.resolve(event); await expect(async () => { for await (const _chunk of event.result) { diff --git a/js/src/instrumentation/plugins/google-genai-plugin.ts b/js/src/instrumentation/plugins/google-genai-plugin.ts index b4ec2b7ee..b4b54ad03 100644 --- a/js/src/instrumentation/plugins/google-genai-plugin.ts +++ b/js/src/instrumentation/plugins/google-genai-plugin.ts @@ -1,25 +1,19 @@ import { uint8ArrayToBase64 } from "../../../util/bytes"; +import { debugLogger } from "../../debug-logger"; import { getExtensionFromMediaType, processInputAttachments, } from "../../wrappers/attachment-utils"; -import { debugLogger } from "../../debug-logger"; import { BasePlugin } from "../core"; -import { traceStreamingChannel, unsubscribeAll } from "../core/channel-tracing"; -import type { - ChannelMessage, - ErrorOf, - StartOf, -} from "../core/channel-definitions"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; +import { traceStreamingCall, unsubscribeAll } from "../core/channel-tracing"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute } from "../../../util/index"; import { - _internalGetGlobalState, Attachment, currentSpan, - BRAINTRUST_CURRENT_SPAN_STORE, startSpan as startBaseSpan, withCurrent, - type CurrentSpanStore, type Span, type StartSpanArgs, } from "../../logger"; @@ -27,17 +21,12 @@ import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { SpanTypeAttribute } from "../../../util/index"; import { getCurrentUnixTimestamp } from "../../util"; -import { googleGenAIChannels } from "./google-genai-channels"; -import { - isAutoInstrumentationSuppressed, - runWithAutoInstrumentationSuppressed, -} from "../auto-instrumentation-suppression"; import type { + GoogleGenAIContent, + GoogleGenAIEditImageParams, GoogleGenAIEmbedContentParams, GoogleGenAIEmbedContentResponse, - GoogleGenAIEditImageParams, GoogleGenAIGenerateContentParams, GoogleGenAIGenerateContentResponse, GoogleGenAIGenerateImagesParams, @@ -45,8 +34,6 @@ import type { GoogleGenAIGenerateVideosOperation, GoogleGenAIGenerateVideosParams, GoogleGenAIImage, - GoogleGenAIVideo, - GoogleGenAIContent, GoogleGenAIInteraction, GoogleGenAIInteractionContent, GoogleGenAIInteractionCreateParams, @@ -54,7 +41,14 @@ import type { GoogleGenAIInteractionUsage, GoogleGenAIPart, GoogleGenAIUsageMetadata, + GoogleGenAIVideo, } from "../../vendor-sdk-types/google-genai"; +import { + isAutoInstrumentationSuppressed, + runWithAutoInstrumentationSuppressed, +} from "../auto-instrumentation-suppression"; +import type { ChannelMessage, ErrorOf } from "../core/tracing-types"; +import { googleGenAIChannels } from "./google-genai-channels"; type GenerateContentChannel = typeof googleGenAIChannels.generateContent; type GenerateContentStreamChannel = @@ -131,41 +125,44 @@ export class GoogleGenAIPlugin extends BasePlugin { } private subscribeToGenerateContentChannel(): void { - const tracingChannel = - googleGenAIChannels.generateContent.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; + const invocationHook = googleGenAIChannels.generateContent; const states = new WeakMap(); - const unbindCurrentSpanStore = bindCurrentSpanStoreToStart( - tracingChannel, - states, - (event) => { - const params = event.arguments[0]; - const input = serializeGenerateContentInput(params); - const metadata = extractGenerateContentMetadata(params); - const span = startBaseSpan( - withSpanInstrumentationName( - { - name: "generate_content", - spanAttributes: { - type: SpanTypeAttribute.LLM, - }, - event: createWrapperParityEvent({ input, metadata }), - }, - INSTRUMENTATION_NAMES.GOOGLE_GENAI, - ), - ); - return { - span, - startTime: getCurrentUnixTimestamp(), - }; - }, - ); + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const spanState = runInstrumentation( + () => + states.get(event) ?? + ((event) => { + const params = event.arguments[0]; + const input = serializeGenerateContentInput(params); + const metadata = extractGenerateContentMetadata(params); + const span = startBaseSpan( + withSpanInstrumentationName( + { + name: "generate_content", + spanAttributes: { + type: SpanTypeAttribute.LLM, + }, + event: createWrapperParityEvent({ input, metadata }), + }, + INSTRUMENTATION_NAMES.GOOGLE_GENAI, + ), + ); - const handlers: IsoChannelHandlers> = - { - start: (event) => { + return { + span, + startTime: getCurrentUnixTimestamp(), + }; + })(event), + ); + if (spanState) states.set(event, spanState); + const prepare = (event: ChannelMessage) => { ensureSpanState(states, event, () => { const params = event.arguments[0]; const input = serializeGenerateContentInput(params); @@ -188,8 +185,8 @@ export class GoogleGenAIPlugin extends BasePlugin { startTime: getCurrentUnixTimestamp(), }; }); - }, - asyncEnd: (event) => { + }; + const resolved = (event: ChannelMessage) => { const spanState = states.get(event as object); if (!spanState) { return; @@ -211,59 +208,87 @@ export class GoogleGenAIPlugin extends BasePlugin { spanState.span.end(); states.delete(event as object); } - }, - error: (event) => { + }; + const failed = (event: ChannelMessage) => { logErrorAndEndSpan(states, event as ErrorOf); - }, - }; + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); + }; + return spanState ? withCurrent(spanState.span, invoke) : invoke(); + }, + ); + this.unsubscribers.push(removeHandlers); } private subscribeToGenerateContentStreamChannel(): void { - const tracingChannel = - googleGenAIChannels.generateContentStream.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; - - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - const streamEvent = event as GenerateContentStreamEvent; - const params = event.arguments[0]; - streamEvent.googleGenAIInput = serializeGenerateContentInput(params); - streamEvent.googleGenAIMetadata = - extractGenerateContentMetadata(params); - streamEvent.googleGenAIStartTime = getCurrentUnixTimestamp(); - }, - asyncEnd: (event) => { - const streamEvent = event as GenerateContentStreamEvent; - patchGoogleGenAIStreamingResult({ - input: streamEvent.googleGenAIInput, - metadata: streamEvent.googleGenAIMetadata, - startTime: streamEvent.googleGenAIStartTime, - result: streamEvent.result, - }); + const invocationHook = googleGenAIChannels.generateContentStream; + + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + const streamEvent = event as GenerateContentStreamEvent; + const params = event.arguments[0]; + streamEvent.googleGenAIInput = serializeGenerateContentInput(params); + streamEvent.googleGenAIMetadata = + extractGenerateContentMetadata(params); + streamEvent.googleGenAIStartTime = getCurrentUnixTimestamp(); + }; + const resolved = ( + event: ChannelMessage, + ) => { + const streamEvent = event as GenerateContentStreamEvent; + patchGoogleGenAIStreamingResult({ + input: streamEvent.googleGenAIInput, + metadata: streamEvent.googleGenAIMetadata, + startTime: streamEvent.googleGenAIStartTime, + result: streamEvent.result, + }); + }; + runInstrumentation(() => prepare(event)); + const result = Reflect.apply(target, receiver, args); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + () => {}, + ); }, - error: () => {}, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToEmbedContentChannel(): void { - const tracingChannel = - googleGenAIChannels.embedContent.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; + const invocationHook = googleGenAIChannels.embedContent; const states = new WeakMap(); const embeddingSpans = new WeakSet(); this.unsubscribers.push( @@ -300,266 +325,292 @@ export class GoogleGenAIPlugin extends BasePlugin { }, ), ); - const unbindCurrentSpanStore = bindCurrentSpanStoreToStart( - tracingChannel, - states, - (event) => { - const params = event.arguments[0]; - const input = serializeEmbedContentInput(params); - const metadata = { provider: "google", model: params.model }; - const span = startBaseSpan( - withSpanInstrumentationName( - { - name: "embed_content", - spanAttributes: { - type: SpanTypeAttribute.LLM, - }, - event: createWrapperParityEvent({ input, metadata }), - }, - INSTRUMENTATION_NAMES.GOOGLE_GENAI, - ), + + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const spanState = runInstrumentation( + () => + states.get(event) ?? + ((event) => { + const params = event.arguments[0]; + const input = serializeEmbedContentInput(params); + const metadata = { provider: "google", model: params.model }; + const span = startBaseSpan( + withSpanInstrumentationName( + { + name: "embed_content", + spanAttributes: { + type: SpanTypeAttribute.LLM, + }, + event: createWrapperParityEvent({ input, metadata }), + }, + INSTRUMENTATION_NAMES.GOOGLE_GENAI, + ), + ); + + embeddingSpans.add(span); + return { + span, + startTime: getCurrentUnixTimestamp(), + }; + })(event), ); + if (spanState) states.set(event, spanState); + const prepare = (event: ChannelMessage) => { + ensureSpanState(states, event, () => { + const params = event.arguments[0]; + const input = serializeEmbedContentInput(params); + const metadata = { provider: "google", model: params.model }; + const span = startBaseSpan( + withSpanInstrumentationName( + { + name: "embed_content", + spanAttributes: { + type: SpanTypeAttribute.LLM, + }, + event: createWrapperParityEvent({ input, metadata }), + }, + INSTRUMENTATION_NAMES.GOOGLE_GENAI, + ), + ); - embeddingSpans.add(span); - return { - span, - startTime: getCurrentUnixTimestamp(), + embeddingSpans.add(span); + return { + span, + startTime: getCurrentUnixTimestamp(), + }; + }); }; - }, - ); - - const handlers: IsoChannelHandlers> = { - start: (event) => { - ensureSpanState(states, event, () => { - const params = event.arguments[0]; - const input = serializeEmbedContentInput(params); - const metadata = { provider: "google", model: params.model }; - const span = startBaseSpan( - withSpanInstrumentationName( - { - name: "embed_content", - spanAttributes: { - type: SpanTypeAttribute.LLM, - }, - event: createWrapperParityEvent({ input, metadata }), - }, - INSTRUMENTATION_NAMES.GOOGLE_GENAI, - ), - ); + const resolved = (event: ChannelMessage) => { + const spanState = states.get(event as object); + if (!spanState) { + return; + } - embeddingSpans.add(span); - return { - span, - startTime: getCurrentUnixTimestamp(), - }; - }); - }, - asyncEnd: (event) => { - const spanState = states.get(event as object); - if (!spanState) { - return; - } + try { + spanState.span.log({ + output: summarizeEmbedContentOutput(event.result), + metrics: cleanMetrics( + extractEmbedContentMetrics(event.result, spanState.startTime), + ), + }); + } finally { + embeddingSpans.delete(spanState.span); + spanState.span.end(); + states.delete(event as object); + } + }; + const failed = (event: ChannelMessage) => { + const spanState = states.get(event as object); + if (!spanState) return; + try { + spanState.span.log({ error: event.error, output: { count: 0 } }); + } finally { + embeddingSpans.delete(spanState.span); + spanState.span.end(); + states.delete(event as object); + } + }; + const invoke = () => { + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; + } - try { - spanState.span.log({ - output: summarizeEmbedContentOutput(event.result), - metrics: cleanMetrics( - extractEmbedContentMetrics(event.result, spanState.startTime), - ), - }); - } finally { - embeddingSpans.delete(spanState.span); - spanState.span.end(); - states.delete(event as object); - } - }, - error: (event) => { - const spanState = states.get(event as object); - if (!spanState) return; - try { - spanState.span.log({ error: event.error, output: { count: 0 } }); - } finally { - embeddingSpans.delete(spanState.span); - spanState.span.end(); - states.delete(event as object); - } + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => { + Object.assign(event, { error }); + failed(event); + }, + ); + }; + return spanState ? withCurrent(spanState.span, invoke) : invoke(); }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - unbindCurrentSpanStore?.(); - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToInteractionsCreateChannel(): void { this.unsubscribers.push( - traceStreamingChannel( - googleGenAIChannels.interactionsCreate as InteractionsCreateChannel, - { - name: ([params]) => - isVideoInteractionCreate(params) - ? "generate_video" - : "create_interaction", - shouldTrace: ([params]) => !isBackgroundInteractionCreate(params), - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => ({ - input: isVideoInteractionCreate(params) - ? serializeVideoInteractionInput(params) - : serializeInteractionInput(params), - metadata: isVideoInteractionCreate(params) - ? { model: params.model, provider: "google" } - : extractInteractionMetadata(params), - }), - extractOutput: (result, event) => - isVideoInteractionCreate(event?.arguments?.[0]) || - getInteractionVideoOutput(result).length > 0 - ? serializeVideoInteractionOutput(result) - : serializeInteractionValue(result), - extractMetadata: (result) => - extractInteractionResponseMetadata(result), - extractMetrics: (result, startTime) => - cleanMetrics(extractInteractionMetrics(result, startTime)), - aggregateChunks: (chunks, _result, _event, startTime) => - aggregateInteractionEvents(chunks, startTime), - }, + googleGenAIChannels.interactionsCreate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GOOGLE_GENAI, + name: ([params]) => + isVideoInteractionCreate(params) + ? "generate_video" + : "create_interaction", + shouldTrace: ([params]) => !isBackgroundInteractionCreate(params), + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => ({ + input: isVideoInteractionCreate(params) + ? serializeVideoInteractionInput(params) + : serializeInteractionInput(params), + metadata: isVideoInteractionCreate(params) + ? { model: params.model, provider: "google" } + : extractInteractionMetadata(params), + }), + extractOutput: (result, event) => + isVideoInteractionCreate(event?.arguments?.[0]) || + getInteractionVideoOutput(result).length > 0 + ? serializeVideoInteractionOutput(result) + : serializeInteractionValue(result), + extractMetadata: (result) => + extractInteractionResponseMetadata(result), + extractMetrics: (result, startTime) => + cleanMetrics(extractInteractionMetrics(result, startTime)), + aggregateChunks: (chunks, _result, _event, startTime) => + aggregateInteractionEvents(chunks, startTime), + }, + ), ), ); } private subscribeToGenerateImagesChannel(): void { this.unsubscribers.push( - interceptGoogleGenAIMediaCall( - googleGenAIChannels.generateImages, - "generate_images", - serializeGenerateImagesInput, - serializeGenerateImagesOutput, + googleGenAIChannels.generateImages.intercept( + (target, receiver, args, additional) => + traceGoogleGenAIMediaCall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "generate_images", + serializeGenerateImagesInput, + serializeGenerateImagesOutput, + ), ), ); } private subscribeToEditImageChannel(): void { this.unsubscribers.push( - interceptGoogleGenAIMediaCall( - googleGenAIChannels.editImage, - "edit_image", - serializeEditImageInput, - serializeGenerateImagesOutput, + googleGenAIChannels.editImage.intercept( + (target, receiver, args, additional) => + traceGoogleGenAIMediaCall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "edit_image", + serializeEditImageInput, + serializeGenerateImagesOutput, + ), ), ); } private subscribeToGenerateVideosChannel(): void { this.unsubscribers.push( - interceptGoogleGenAIMediaCall( - googleGenAIChannels.generateVideos, - "generate_videos", - serializeGenerateVideosInput, - serializeGenerateVideosOutput, + googleGenAIChannels.generateVideos.intercept( + (target, receiver, args, additional) => + traceGoogleGenAIMediaCall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "generate_videos", + serializeGenerateVideosInput, + serializeGenerateVideosOutput, + ), ), ); } } -type GoogleGenAIMediaChannel = { - intercept( - interceptor: ( - target: (this: unknown, params: TParams) => PromiseLike, - thisArg: unknown, - args: [TParams], - ) => PromiseLike, - ): () => void; -}; - -function interceptGoogleGenAIMediaCall< - TParams extends { model: string }, - TResult, ->( - channel: GoogleGenAIMediaChannel, +function traceGoogleGenAIMediaCall( + call: () => PromiseLike, + context: { arguments: [TParams]; self: unknown; additional: unknown }, name: string, serializeInput: (params: TParams) => Record, serializeOutput: ( response: TResult, params: TParams, ) => Record, -): () => void { - return channel.intercept((target, thisArg, args) => { - const invoke = () => Reflect.apply(target, thisArg, args); - if (isAutoInstrumentationSuppressed()) { - return invoke(); - } +): PromiseLike { + const args = context.arguments; - const [params] = args; - let span: Span; + if (isAutoInstrumentationSuppressed()) { + return call(); + } + const [params] = args; + let span: Span; + try { + span = startBaseSpan( + withSpanInstrumentationName( + { + name, + spanAttributes: { type: SpanTypeAttribute.LLM }, + event: createWrapperParityEvent({ + input: serializeInput(params), + metadata: { model: params.model, provider: "google" }, + }), + }, + INSTRUMENTATION_NAMES.GOOGLE_GENAI, + ), + ); + } catch (error) { + debugLogger.error(`Error starting Google GenAI ${name} span:`, error); + return call(); + } + let ended = false; + const finish = (error?: unknown) => { + if (ended) { + return; + } + ended = true; try { - span = startBaseSpan( - withSpanInstrumentationName( - { - name, - spanAttributes: { type: SpanTypeAttribute.LLM }, - event: createWrapperParityEvent({ - input: serializeInput(params), - metadata: { model: params.model, provider: "google" }, - }), - }, - INSTRUMENTATION_NAMES.GOOGLE_GENAI, - ), + if (error !== undefined) { + span.log({ error }); + } + span.end(); + } catch (loggingError) { + debugLogger.error( + `Error ending Google GenAI ${name} span:`, + loggingError, ); - } catch (error) { - debugLogger.error(`Error starting Google GenAI ${name} span:`, error); - return invoke(); } - - let ended = false; - const finish = (error?: unknown) => { - if (ended) { - return; - } - ended = true; + }; + let result: PromiseLike; + try { + result = withCurrent(span, () => + runWithAutoInstrumentationSuppressed(call), + ); + } catch (error) { + finish(error); + throw error; + } + try { + void Promise.resolve(result).then((response) => { try { - if (error !== undefined) { - span.log({ error }); - } - span.end(); - } catch (loggingError) { + span.log({ output: serializeOutput(response, params) }); + } catch (error) { debugLogger.error( - `Error ending Google GenAI ${name} span:`, - loggingError, + `Error capturing Google GenAI ${name} output:`, + error, ); + } finally { + finish(); } - }; - - let result: PromiseLike; - try { - result = withCurrent(span, () => - runWithAutoInstrumentationSuppressed(invoke), - ); - } catch (error) { - finish(error); - throw error; - } - - try { - void Promise.resolve(result).then((response) => { - try { - span.log({ output: serializeOutput(response, params) }); - } catch (error) { - debugLogger.error( - `Error capturing Google GenAI ${name} output:`, - error, - ); - } finally { - finish(); - } - }, finish); - } catch (error) { - debugLogger.error(`Error observing Google GenAI ${name} result:`, error); - finish(); - } - - return result; - }); + }, finish); + } catch (error) { + debugLogger.error(`Error observing Google GenAI ${name} result:`, error); + finish(); + } + return result; } function isBackgroundInteractionCreate(params: unknown): boolean { @@ -581,48 +632,6 @@ function ensureSpanState( return created; } -function bindCurrentSpanStoreToStart< - TChannel extends GoogleGenAINonStreamingChannel, ->( - tracingChannel: IsoTracingChannel>, - states: WeakMap, - create: (event: StartOf) => SpanState, -): (() => void) | undefined { - const state = _internalGetGlobalState(); - const contextManager = state?.contextManager; - const startChannel = tracingChannel.start as - | ({ - bindStore?: ( - store: CurrentSpanStore, - callback: (event: ChannelMessage) => unknown, - ) => void; - unbindStore?: (store: CurrentSpanStore) => void; - } & object) - | undefined; - const currentSpanStore = contextManager - ? ( - contextManager as { - [BRAINTRUST_CURRENT_SPAN_STORE]?: CurrentSpanStore; - } - )[BRAINTRUST_CURRENT_SPAN_STORE] - : undefined; - - if (!startChannel?.bindStore || !currentSpanStore) { - return undefined; - } - - startChannel.bindStore(currentSpanStore, (event) => { - const span = ensureSpanState(states, event as object, () => - create(event as StartOf), - ).span; - return contextManager!.wrapSpanForStore(span); - }); - - return () => { - startChannel.unbindStore?.(currentSpanStore); - }; -} - function logErrorAndEndSpan( states: WeakMap, event: ErrorOf, diff --git a/js/src/instrumentation/plugins/google-generative-ai-channels.ts b/js/src/instrumentation/plugins/google-generative-ai-channels.ts index 66197b7ab..4fb1f8fba 100644 --- a/js/src/instrumentation/plugins/google-generative-ai-channels.ts +++ b/js/src/instrumentation/plugins/google-generative-ai-channels.ts @@ -1,37 +1,38 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { GenerativeAIChat, GenerativeAIModel, } from "../../vendor-sdk-types/google-generative-ai"; -export const googleGenerativeAIChannels = defineChannels( +export const googleGenerativeAIChannels = defineInterceptor( "@google/generative-ai", { generateContent: channel< Parameters, - Awaited> - >({ channelName: "GenerativeModel.generateContent", kind: "async" }), + PromiseLike>> + >({ channelName: "GenerativeModel.generateContent" }), generateContentStream: channel< Parameters, - Awaited> - >({ channelName: "GenerativeModel.generateContentStream", kind: "async" }), + PromiseLike< + Awaited> + > + >({ channelName: "GenerativeModel.generateContentStream" }), embedContent: channel< Parameters, - Awaited> - >({ channelName: "GenerativeModel.embedContent", kind: "async" }), + PromiseLike>> + >({ channelName: "GenerativeModel.embedContent" }), batchEmbedContents: channel< Parameters, - Awaited> - >({ channelName: "GenerativeModel.batchEmbedContents", kind: "async" }), + PromiseLike>> + >({ channelName: "GenerativeModel.batchEmbedContents" }), sendMessage: channel< Parameters, - Awaited> - >({ channelName: "ChatSession.sendMessage", kind: "async" }), + PromiseLike>> + >({ channelName: "ChatSession.sendMessage" }), sendMessageStream: channel< Parameters, - Awaited> - >({ channelName: "ChatSession.sendMessageStream", kind: "async" }), + PromiseLike>> + >({ channelName: "ChatSession.sendMessageStream" }), }, - { instrumentationName: INSTRUMENTATION_NAMES.GOOGLE_GENERATIVE_AI }, ); diff --git a/js/src/instrumentation/plugins/google-generative-ai-plugin.ts b/js/src/instrumentation/plugins/google-generative-ai-plugin.ts index d013bf031..86ca1665c 100644 --- a/js/src/instrumentation/plugins/google-generative-ai-plugin.ts +++ b/js/src/instrumentation/plugins/google-generative-ai-plugin.ts @@ -56,36 +56,56 @@ type Request = | { requests: GenerativeAIEmbedRequest[] }; type Operation = keyof typeof googleGenerativeAIChannels; -type GenerativeAIChannel = { - intercept( - interceptor: ( - target: (this: unknown, ...args: TArgs) => PromiseLike, - thisArg: unknown, - args: TArgs, - ) => PromiseLike, - ): () => void; -}; - export class GoogleGenerativeAIPlugin extends BasePlugin { protected onEnable(): void { this.unsubscribers.push( - interceptCall( - googleGenerativeAIChannels.generateContent, - "generateContent", + googleGenerativeAIChannels.generateContent.intercept( + (target, receiver, args, additional) => + traceGenerativeAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "generateContent", + ), + ), + googleGenerativeAIChannels.generateContentStream.intercept( + (target, receiver, args, additional) => + traceGenerativeAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "generateContentStream", + ), ), - interceptCall( - googleGenerativeAIChannels.generateContentStream, - "generateContentStream", + googleGenerativeAIChannels.sendMessage.intercept( + (target, receiver, args, additional) => + traceGenerativeAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "sendMessage", + ), ), - interceptCall(googleGenerativeAIChannels.sendMessage, "sendMessage"), - interceptCall( - googleGenerativeAIChannels.sendMessageStream, - "sendMessageStream", + googleGenerativeAIChannels.sendMessageStream.intercept( + (target, receiver, args, additional) => + traceGenerativeAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "sendMessageStream", + ), ), - interceptCall(googleGenerativeAIChannels.embedContent, "embedContent"), - interceptCall( - googleGenerativeAIChannels.batchEmbedContents, - "batchEmbedContents", + googleGenerativeAIChannels.embedContent.intercept( + (target, receiver, args, additional) => + traceGenerativeAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "embedContent", + ), + ), + googleGenerativeAIChannels.batchEmbedContents.intercept( + (target, receiver, args, additional) => + traceGenerativeAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "batchEmbedContents", + ), ), ); } @@ -95,206 +115,204 @@ export class GoogleGenerativeAIPlugin extends BasePlugin { } } -function interceptCall< +function traceGenerativeAICall< TArgs extends [Request, unknown?], TResult extends Result, >( - channel: GenerativeAIChannel, + call: () => PromiseLike, + context: { arguments: TArgs; self: unknown; additional: unknown }, operation: Operation, -): () => void { - return channel.intercept((target, thisArg, args) => { - const invokeTarget = () => Reflect.apply(target, thisArg, args); - if (isAutoInstrumentationSuppressed()) return invokeTarget(); - const self = thisArg as GenerativeAIModel | GenerativeAIChat; - const chat = operation.startsWith("sendMessage"); - const embedding = - operation === "embedContent" || operation === "batchEmbedContents"; - const start = getCurrentUnixTimestamp(); - let span: Span; +): PromiseLike { + const args = context.arguments; + const thisArg = context.self; + + if (isAutoInstrumentationSuppressed()) return call(); + const self = thisArg as GenerativeAIModel | GenerativeAIChat; + const chat = operation.startsWith("sendMessage"); + const embedding = + operation === "embedContent" || operation === "batchEmbedContents"; + const start = getCurrentUnixTimestamp(); + let span: Span; + try { + span = startSpan( + withSpanInstrumentationName( + { + name: embedding + ? operation === "embedContent" + ? "embed_content" + : "batch_embed_contents" + : "generate_content", + spanAttributes: { type: SpanTypeAttribute.LLM }, + event: extractInput(self, args[0], operation), + }, + INSTRUMENTATION_NAMES.GOOGLE_GENERATIVE_AI, + ), + ); + } catch (error) { + debugLogger.error("Error starting Google Generative AI span:", error); + return call(); + } + let ended = false; + const finish = (log: () => void) => { + if (ended) return; + ended = true; try { - span = startSpan( - withSpanInstrumentationName( - { - name: embedding - ? operation === "embedContent" - ? "embed_content" - : "batch_embed_contents" - : "generate_content", - spanAttributes: { type: SpanTypeAttribute.LLM }, - event: extractInput(self, args[0], operation), - }, - INSTRUMENTATION_NAMES.GOOGLE_GENERATIVE_AI, - ), - ); + log(); } catch (error) { - debugLogger.error("Error starting Google Generative AI span:", error); - return invokeTarget(); + debugLogger.error("Error logging Google Generative AI span:", error); } - let ended = false; - const finish = (log: () => void) => { - if (ended) return; - ended = true; - try { - log(); - } catch (error) { - debugLogger.error("Error logging Google Generative AI span:", error); - } - try { - span.end(); - } catch (error) { - debugLogger.error("Error ending Google Generative AI span:", error); - } - }; - // The SDK serializes chat sends on this promise. Capture history after the - // preceding send completes, before the current SDK call appends its turn. - if (chat) { - try { - void (self as GenerativeAIChat)._sendPromise.then( - () => { - try { - span.log(extractInput(self, args[0], operation)); - } catch (error) { - debugLogger.error("Error capturing Google chat history:", error); - } - }, - () => {}, - ); - } catch (error) { - debugLogger.error("Error observing Google chat history:", error); - } + try { + span.end(); + } catch (error) { + debugLogger.error("Error ending Google Generative AI span:", error); } - let result: PromiseLike; + }; + if (chat) { try { - result = withCurrent(span, () => - runWithAutoInstrumentationSuppressed(invokeTarget), + void (self as GenerativeAIChat)._sendPromise.then( + () => { + try { + span.log(extractInput(self, args[0], operation)); + } catch (error) { + debugLogger.error("Error capturing Google chat history:", error); + } + }, + () => {}, ); } catch (error) { - finish(() => span.log({ error })); - throw error; + debugLogger.error("Error observing Google chat history:", error); } - void Promise.resolve(result).then( - (value) => { - try { - if ("stream" in value) { - let firstToken = false; - const partial: GenerativeAIResponse = { candidates: [] }; - const candidates = new Map< - number, - NonNullable[number] - >(); - patchStreamIfNeeded(value.stream, { - aroundNext: (callback) => withCurrent(span, callback), - onChunk: (chunk) => { - for (const candidate of chunk.candidates ?? []) { - const index = candidate.index ?? 0; - const previous = candidates.get(index); - const parts = [...(previous?.content?.parts ?? [])]; - for (const part of candidate.content?.parts ?? []) { - const last = parts[parts.length - 1]; - if (part.text !== undefined && last?.text !== undefined) { - parts[parts.length - 1] = { - ...last, - text: last.text + part.text, - }; - } else parts.push(part); - } - candidates.set(index, { - ...previous, - ...candidate, - content: { - role: - candidate.content?.role ?? - previous?.content?.role ?? - "model", - parts, - }, - }); - } - partial.candidates = Array.from(candidates.values()); - if (chunk.usageMetadata) - partial.usageMetadata = chunk.usageMetadata; - if ( - !firstToken && - chunk.candidates?.some((candidate) => - candidate.content?.parts.some((part) => - part.text !== undefined - ? part.text.length > 0 - : Object.keys(part).length > 0, - ), - ) - ) { - firstToken = true; - span.log({ - metrics: { - time_to_first_token: getCurrentUnixTimestamp() - start, - }, - }); + } + let result: PromiseLike; + try { + result = withCurrent(span, () => + runWithAutoInstrumentationSuppressed(call), + ); + } catch (error) { + finish(() => span.log({ error })); + throw error; + } + void Promise.resolve(result).then( + (value) => { + try { + if ("stream" in value) { + let firstToken = false; + const partial: GenerativeAIResponse = { candidates: [] }; + const candidates = new Map< + number, + NonNullable[number] + >(); + patchStreamIfNeeded(value.stream, { + aroundNext: (callback) => withCurrent(span, callback), + onChunk: (chunk) => { + for (const candidate of chunk.candidates ?? []) { + const index = candidate.index ?? 0; + const previous = candidates.get(index); + const parts = [...(previous?.content?.parts ?? [])]; + for (const part of candidate.content?.parts ?? []) { + const last = parts[parts.length - 1]; + if (part.text !== undefined && last?.text !== undefined) { + parts[parts.length - 1] = { + ...last, + text: last.text + part.text, + }; + } else parts.push(part); } - }, - onComplete: () => {}, - onError: (error) => - finish(() => { - logResponse(span, partial); - span.log({ error }); - }), - onCancel: () => finish(() => logResponse(span, partial)), - }); - // The SDK tees its stream to build this aggregate, even when callers - // only await response. Observe it without replacing either public value. - void value.response.then( - (response) => finish(() => logResponse(span, response)), - (error) => - finish(() => { - logResponse(span, partial); - span.log({ error }); - }), - ); - } else if ("response" in value) { - finish(() => logResponse(span, value.response)); - } else { - const metrics: Record = {}; - const promptTokens = value.usageMetadata?.promptTokenCount; - if ( - typeof promptTokens === "number" && - Number.isFinite(promptTokens) && - promptTokens >= 0 - ) { - metrics.prompt_tokens = promptTokens; - metrics.tokens = promptTokens; - } - for (const detail of value.usageMetadata?.promptTokenDetails ?? - []) { + candidates.set(index, { + ...previous, + ...candidate, + content: { + role: + candidate.content?.role ?? + previous?.content?.role ?? + "model", + parts, + }, + }); + } + partial.candidates = Array.from(candidates.values()); + if (chunk.usageMetadata) + partial.usageMetadata = chunk.usageMetadata; if ( - detail.modality === "AUDIO" && - typeof detail.tokenCount === "number" && - Number.isFinite(detail.tokenCount) && - detail.tokenCount >= 0 + !firstToken && + chunk.candidates?.some((candidate) => + candidate.content?.parts.some((part) => + part.text !== undefined + ? part.text.length > 0 + : Object.keys(part).length > 0, + ), + ) ) { - metrics.prompt_audio_tokens = - (metrics.prompt_audio_tokens ?? 0) + detail.tokenCount; + firstToken = true; + span.log({ + metrics: { + time_to_first_token: getCurrentUnixTimestamp() - start, + }, + }); } - } - finish(() => - span.log({ - metrics, - output: { - count: value.embeddings?.length ?? (value.embedding ? 1 : 0), - }, + }, + onComplete: () => {}, + onError: (error) => + finish(() => { + logResponse(span, partial); + span.log({ error }); }), - ); + onCancel: () => finish(() => logResponse(span, partial)), + }); + // The SDK tees its stream to build this aggregate, even when callers + // only await response. Observe it without replacing either public value. + void value.response.then( + (response) => finish(() => logResponse(span, response)), + (error) => + finish(() => { + logResponse(span, partial); + span.log({ error }); + }), + ); + } else if ("response" in value) { + finish(() => logResponse(span, value.response)); + } else { + const metrics: Record = {}; + const promptTokens = value.usageMetadata?.promptTokenCount; + if ( + typeof promptTokens === "number" && + Number.isFinite(promptTokens) && + promptTokens >= 0 + ) { + metrics.prompt_tokens = promptTokens; + metrics.tokens = promptTokens; } - } catch (error) { - debugLogger.error( - "Error observing Google Generative AI result:", - error, + for (const detail of value.usageMetadata?.promptTokenDetails ?? []) { + if ( + detail.modality === "AUDIO" && + typeof detail.tokenCount === "number" && + Number.isFinite(detail.tokenCount) && + detail.tokenCount >= 0 + ) { + metrics.prompt_audio_tokens = + (metrics.prompt_audio_tokens ?? 0) + detail.tokenCount; + } + } + finish(() => + span.log({ + metrics, + output: { + count: value.embeddings?.length ?? (value.embedding ? 1 : 0), + }, + }), ); - finish(() => {}); } - }, - (error) => finish(() => span.log({ error })), - ); - return result; - }); + } catch (error) { + debugLogger.error( + "Error observing Google Generative AI result:", + error, + ); + finish(() => {}); + } + }, + (error) => finish(() => span.log({ error })), + ); + return result; } function normalizeContent(message: GenerativeAIMessage): GenerativeAIContent { diff --git a/js/src/instrumentation/plugins/groq-channels.ts b/js/src/instrumentation/plugins/groq-channels.ts index 5506e459e..ab3256abb 100644 --- a/js/src/instrumentation/plugins/groq-channels.ts +++ b/js/src/instrumentation/plugins/groq-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { GroqAudioSpeechCreateParams, GroqAudioTextResult, @@ -15,50 +15,41 @@ import type { type GroqChatResult = GroqChatCompletion | GroqChatStream; -export const groqChannels = defineChannels( - "groq-sdk", - { - chatCompletionsCreate: channel< - [GroqChatCreateParams, unknown?], - GroqChatResult, - Record, - GroqChatCompletionChunk - >({ - channelName: "chat.completions.create", - kind: "async", - }), +export const groqChannels = defineInterceptor("groq-sdk", { + chatCompletionsCreate: channel< + [GroqChatCreateParams, unknown?], + PromiseLike, + Record, + GroqChatCompletionChunk + >({ + channelName: "chat.completions.create", + }), - embeddingsCreate: channel< - [GroqEmbeddingCreateParams, unknown?], - GroqEmbeddingResponse - >({ - channelName: "embeddings.create", - kind: "async", - }), + embeddingsCreate: channel< + [GroqEmbeddingCreateParams, unknown?], + PromiseLike + >({ + channelName: "embeddings.create", + }), - audioSpeechCreate: channel< - [GroqAudioSpeechCreateParams, unknown?], - Response - >({ - channelName: "audio.speech.create", - kind: "async", - }), + audioSpeechCreate: channel< + [GroqAudioSpeechCreateParams, unknown?], + PromiseLike + >({ + channelName: "audio.speech.create", + }), - audioTranscriptionsCreate: channel< - [GroqAudioTranscriptionCreateParams, unknown?], - GroqAudioTextResult | string - >({ - channelName: "audio.transcriptions.create", - kind: "async", - }), + audioTranscriptionsCreate: channel< + [GroqAudioTranscriptionCreateParams, unknown?], + PromiseLike + >({ + channelName: "audio.transcriptions.create", + }), - audioTranslationsCreate: channel< - [GroqAudioTranslationCreateParams, unknown?], - GroqAudioTextResult | string - >({ - channelName: "audio.translations.create", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.GROQ }, -); + audioTranslationsCreate: channel< + [GroqAudioTranslationCreateParams, unknown?], + PromiseLike + >({ + channelName: "audio.translations.create", + }), +}); diff --git a/js/src/instrumentation/plugins/groq-plugin.ts b/js/src/instrumentation/plugins/groq-plugin.ts index b37623a46..25688cccc 100644 --- a/js/src/instrumentation/plugins/groq-plugin.ts +++ b/js/src/instrumentation/plugins/groq-plugin.ts @@ -1,28 +1,12 @@ -import { BasePlugin } from "../core"; -import { - traceAsyncChannel, - traceStreamingChannel, - unsubscribeAll, -} from "../core/channel-tracing"; import { + SpanTypeAttribute, concatUint8Arrays, isObject, - SpanTypeAttribute, } from "../../../util/index"; +import { debugLogger } from "../../debug-logger"; import { Attachment, withCurrent, type Span } from "../../logger"; -import { - convertDataToBlob, - getExtensionFromMediaType, - processInputAttachments, -} from "../../wrappers/attachment-utils"; +import { INSTRUMENTATION_NAMES } from "../../span-origin"; import { getCurrentUnixTimestamp } from "../../util"; -import { - aggregateChatCompletionChunks, - parseMetricsFromUsage, -} from "./openai-plugin"; -import { groqChannels } from "./groq-channels"; -import { isAsyncIterable, observeByteStream } from "../core/stream-patcher"; -import { debugLogger } from "../../debug-logger"; import type { GroqAudioSpeechCreateParams, GroqAudioTextResult, @@ -31,101 +15,159 @@ import type { GroqChatCompletion, GroqChatCompletionChunk, } from "../../vendor-sdk-types/groq"; +import { + convertDataToBlob, + getExtensionFromMediaType, + processInputAttachments, +} from "../../wrappers/attachment-utils"; +import { BasePlugin } from "../core"; +import { + traceAsyncCall, + traceStreamingCall, + unsubscribeAll, +} from "../core/channel-tracing"; +import { isAsyncIterable, observeByteStream } from "../core/stream-patcher"; +import { groqChannels } from "./groq-channels"; +import { + aggregateChatCompletionChunks, + parseMetricsFromUsage, +} from "./openai-plugin"; export class GroqPlugin extends BasePlugin { protected onEnable(): void { this.unsubscribers.push( - traceStreamingChannel(groqChannels.chatCompletionsCreate, { - name: "groq.chat.completions.create", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => { - const { messages, ...metadata } = params; - return { - input: processInputAttachments(messages), - metadata: { ...metadata, provider: "groq" }, - }; - }, - extractOutput: (result) => result?.choices, - extractMetrics: (result, startTime) => { - const metrics = parseGroqMetrics(result); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - aggregateChunks: aggregateGroqChatCompletionChunks, - }), + groqChannels.chatCompletionsCreate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GROQ, + name: "groq.chat.completions.create", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => { + const { messages, ...metadata } = params; + return { + input: processInputAttachments(messages), + metadata: { ...metadata, provider: "groq" }, + }; + }, + extractOutput: (result) => result?.choices, + extractMetrics: (result, startTime) => { + const metrics = parseGroqMetrics(result); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + aggregateChunks: aggregateGroqChatCompletionChunks, + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(groqChannels.embeddingsCreate, { - name: "groq.embeddings.create", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => { - const { input, ...metadata } = params; - return { - input, - metadata: { ...metadata, provider: "groq" }, - }; - }, - extractOutput: (result) => { - const embedding = result?.data?.[0]?.embedding; - return Array.isArray(embedding) - ? { embedding_length: embedding.length } - : undefined; - }, - extractMetrics: (result) => parseGroqMetrics(result), - }), + groqChannels.embeddingsCreate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GROQ, + name: "groq.embeddings.create", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => { + const { input, ...metadata } = params; + return { + input, + metadata: { ...metadata, provider: "groq" }, + }; + }, + extractOutput: (result) => { + const embedding = result?.data?.[0]?.embedding; + return Array.isArray(embedding) + ? { embedding_length: embedding.length } + : undefined; + }, + extractMetrics: (result) => parseGroqMetrics(result), + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(groqChannels.audioSpeechCreate, { - name: "groq.audio.speech.create", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => ({ - input: { - operation: "speech", - prompt: params.input, - parameters: { - voice: params.voice, - format: params.response_format, - speed: params.speed, + groqChannels.audioSpeechCreate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GROQ, + name: "groq.audio.speech.create", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => ({ + input: { + operation: "speech", + prompt: params.input, + parameters: { + voice: params.voice, + format: params.response_format, + speed: params.speed, + }, + }, + metadata: { model: params.model, provider: "groq" }, + }), + extractOutput: () => ({ content: [] }), + extractMetrics: () => ({}), + patchResult: ({ endEvent, result, span, startTime }) => + captureGroqSpeechResponse( + result, + endEvent.arguments![0], + span, + startTime, + ), }, - }, - metadata: { model: params.model, provider: "groq" }, - }), - extractOutput: () => ({ content: [] }), - extractMetrics: () => ({}), - patchResult: ({ endEvent, result, span, startTime }) => - captureGroqSpeechResponse( - result, - endEvent.arguments![0], - span, - startTime, ), - }), + ), ); this.unsubscribers.push( - traceAsyncChannel(groqChannels.audioTranscriptionsCreate, { - name: "groq.audio.transcriptions.create", - type: SpanTypeAttribute.LLM, - extractInput: ([params], _event, span) => - extractGroqAudioInput(params, "transcribe", span), - extractOutput: extractGroqAudioTextOutput, - extractMetrics: (result) => parseGroqMetricsObject(result), - }), + groqChannels.audioTranscriptionsCreate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GROQ, + name: "groq.audio.transcriptions.create", + type: SpanTypeAttribute.LLM, + extractInput: ([params], _event, span) => + extractGroqAudioInput(params, "transcribe", span), + extractOutput: extractGroqAudioTextOutput, + extractMetrics: (result) => parseGroqMetricsObject(result), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(groqChannels.audioTranslationsCreate, { - name: "groq.audio.translations.create", - type: SpanTypeAttribute.LLM, - extractInput: ([params], _event, span) => - extractGroqAudioInput(params, "translate", span), - extractOutput: extractGroqAudioTextOutput, - extractMetrics: (result) => parseGroqMetricsObject(result), - }), + groqChannels.audioTranslationsCreate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.GROQ, + name: "groq.audio.translations.create", + type: SpanTypeAttribute.LLM, + extractInput: ([params], _event, span) => + extractGroqAudioInput(params, "translate", span), + extractOutput: extractGroqAudioTextOutput, + extractMetrics: (result) => parseGroqMetricsObject(result), + }, + ), + ), ); } diff --git a/js/src/instrumentation/plugins/huggingface-channels.ts b/js/src/instrumentation/plugins/huggingface-channels.ts index 6996a9cc3..b7f986051 100644 --- a/js/src/instrumentation/plugins/huggingface-channels.ts +++ b/js/src/instrumentation/plugins/huggingface-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { HuggingFaceChatCompletion, HuggingFaceChatCompletionChunk, @@ -11,52 +11,43 @@ import type { HuggingFaceTextGenerationStreamOutput, } from "../../vendor-sdk-types/huggingface"; -export const huggingFaceChannels = defineChannels( - "@huggingface/inference", - { - chatCompletion: channel< - [HuggingFaceChatCompletionParams], - HuggingFaceChatCompletion - >({ - channelName: "chatCompletion", - kind: "async", - }), +export const huggingFaceChannels = defineInterceptor("@huggingface/inference", { + chatCompletion: channel< + [HuggingFaceChatCompletionParams], + PromiseLike + >({ + channelName: "chatCompletion", + }), - chatCompletionStream: channel< - [HuggingFaceChatCompletionParams], - AsyncIterable, - Record, - HuggingFaceChatCompletionChunk - >({ - channelName: "chatCompletionStream", - kind: "sync-stream", - }), + chatCompletionStream: channel< + [HuggingFaceChatCompletionParams], + AsyncIterable, + Record, + HuggingFaceChatCompletionChunk + >({ + channelName: "chatCompletionStream", + }), - textGeneration: channel< - [HuggingFaceTextGenerationParams], - HuggingFaceTextGenerationOutput - >({ - channelName: "textGeneration", - kind: "async", - }), + textGeneration: channel< + [HuggingFaceTextGenerationParams], + PromiseLike + >({ + channelName: "textGeneration", + }), - textGenerationStream: channel< - [HuggingFaceTextGenerationParams], - AsyncIterable, - Record, - HuggingFaceTextGenerationStreamOutput - >({ - channelName: "textGenerationStream", - kind: "sync-stream", - }), + textGenerationStream: channel< + [HuggingFaceTextGenerationParams], + AsyncIterable, + Record, + HuggingFaceTextGenerationStreamOutput + >({ + channelName: "textGenerationStream", + }), - featureExtraction: channel< - [HuggingFaceFeatureExtractionParams], - HuggingFaceFeatureExtractionOutput - >({ - channelName: "featureExtraction", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE }, -); + featureExtraction: channel< + [HuggingFaceFeatureExtractionParams], + PromiseLike + >({ + channelName: "featureExtraction", + }), +}); diff --git a/js/src/instrumentation/plugins/huggingface-plugin.ts b/js/src/instrumentation/plugins/huggingface-plugin.ts index 797c3bf82..aeb503077 100644 --- a/js/src/instrumentation/plugins/huggingface-plugin.ts +++ b/js/src/instrumentation/plugins/huggingface-plugin.ts @@ -1,14 +1,8 @@ -import { - traceAsyncChannel, - traceSyncStreamChannel, - unsubscribeAll, -} from "../core/channel-tracing"; -import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; -import { BasePlugin } from "../core"; import { SpanTypeAttribute, isObject } from "../../../util/index"; -import { getCurrentUnixTimestamp } from "../../util"; +import type { Span } from "../../logger"; import { parseMetricsFromUsage } from "../../openai-utils"; -import { huggingFaceChannels } from "./huggingface-channels"; +import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { getCurrentUnixTimestamp } from "../../util"; import type { HuggingFaceChatCompletion, HuggingFaceChatCompletionChunk, @@ -16,7 +10,14 @@ import type { HuggingFaceTextGenerationDetails, HuggingFaceTextGenerationStreamOutput, } from "../../vendor-sdk-types/huggingface"; -import type { Span } from "../../logger"; +import { BasePlugin } from "../core"; +import { + traceAsyncCall, + traceSyncStreamCall, + unsubscribeAll, +} from "../core/channel-tracing"; +import { isAsyncIterable, patchStreamIfNeeded } from "../core/stream-patcher"; +import { huggingFaceChannels } from "./huggingface-channels"; const REQUEST_METADATA_ALLOWLIST = new Set([ "dimensions", @@ -44,53 +45,135 @@ const RESPONSE_METADATA_ALLOWLIST = new Set([ export class HuggingFacePlugin extends BasePlugin { protected onEnable(): void { this.unsubscribers.push( - traceAsyncChannel(huggingFaceChannels.chatCompletion, { - name: "huggingface.chat_completion", - type: SpanTypeAttribute.LLM, - extractInput: extractChatInputWithMetadata, - extractOutput: (result) => result?.choices, - extractMetadata: (result) => extractResponseMetadata(result), - extractMetrics: (result) => parseMetricsFromUsage(result?.usage), - }), - traceSyncStreamChannel(huggingFaceChannels.chatCompletionStream, { - name: "huggingface.chat_completion_stream", - type: SpanTypeAttribute.LLM, - extractInput: extractChatInputWithMetadata, - patchResult: ({ result, span, startTime }) => - patchChatCompletionStream({ - result, - span, - startTime, - }), - }), - traceAsyncChannel(huggingFaceChannels.textGeneration, { - name: "huggingface.text_generation", - type: SpanTypeAttribute.LLM, - extractInput: extractTextGenerationInputWithMetadata, - extractOutput: (result) => - isObject(result) ? { generated_text: result.generated_text } : result, - extractMetadata: extractTextGenerationMetadata, - extractMetrics: (result) => - extractTextGenerationMetrics(result?.details ?? null), - }), - traceSyncStreamChannel(huggingFaceChannels.textGenerationStream, { - name: "huggingface.text_generation_stream", - type: SpanTypeAttribute.LLM, - extractInput: extractTextGenerationInputWithMetadata, - patchResult: ({ result, span, startTime }) => - patchTextGenerationStream({ - result, - span, - startTime, - }), - }), - traceAsyncChannel(huggingFaceChannels.featureExtraction, { - name: "huggingface.feature_extraction", - type: SpanTypeAttribute.LLM, - extractInput: extractFeatureExtractionInputWithMetadata, - extractOutput: summarizeFeatureExtractionOutput, - extractMetrics: () => ({}), - }), + huggingFaceChannels.chatCompletion.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE, + name: "huggingface.chat_completion", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => + extractChatInputWithMetadata([ + { + ...params, + ...(additional.endpointUrl && !params.endpointUrl + ? { endpointUrl: additional.endpointUrl } + : {}), + }, + ]), + extractOutput: (result) => result?.choices, + extractMetadata: (result) => extractResponseMetadata(result), + extractMetrics: (result) => parseMetricsFromUsage(result?.usage), + }, + ), + ), + huggingFaceChannels.chatCompletionStream.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE, + name: "huggingface.chat_completion_stream", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => + extractChatInputWithMetadata([ + { + ...params, + ...(additional.endpointUrl && !params.endpointUrl + ? { endpointUrl: additional.endpointUrl } + : {}), + }, + ]), + patchResult: ({ result, span, startTime }) => + patchChatCompletionStream({ + result, + span, + startTime, + }), + }, + ), + ), + huggingFaceChannels.textGeneration.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE, + name: "huggingface.text_generation", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => + extractTextGenerationInputWithMetadata([ + { + ...params, + ...(additional.endpointUrl && !params.endpointUrl + ? { endpointUrl: additional.endpointUrl } + : {}), + }, + ]), + extractOutput: (result) => + isObject(result) + ? { generated_text: result.generated_text } + : result, + extractMetadata: extractTextGenerationMetadata, + extractMetrics: (result) => + extractTextGenerationMetrics(result?.details ?? null), + }, + ), + ), + huggingFaceChannels.textGenerationStream.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE, + name: "huggingface.text_generation_stream", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => + extractTextGenerationInputWithMetadata([ + { + ...params, + ...(additional.endpointUrl && !params.endpointUrl + ? { endpointUrl: additional.endpointUrl } + : {}), + }, + ]), + patchResult: ({ result, span, startTime }) => + patchTextGenerationStream({ + result, + span, + startTime, + }), + }, + ), + ), + huggingFaceChannels.featureExtraction.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE, + name: "huggingface.feature_extraction", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => + extractFeatureExtractionInputWithMetadata([ + { + ...params, + ...(additional.endpointUrl && !params.endpointUrl + ? { endpointUrl: additional.endpointUrl } + : {}), + }, + ]), + extractOutput: summarizeFeatureExtractionOutput, + extractMetrics: () => ({}), + }, + ), + ), ); } diff --git a/js/src/instrumentation/plugins/huggingface-transformers-channels.ts b/js/src/instrumentation/plugins/huggingface-transformers-channels.ts index ac2f1fcbe..904b8fd5b 100644 --- a/js/src/instrumentation/plugins/huggingface-transformers-channels.ts +++ b/js/src/instrumentation/plugins/huggingface-transformers-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { HuggingFaceTransformersPipeline, HuggingFaceTransformersTask, @@ -8,7 +8,8 @@ import type { export type HuggingFaceTransformersEventContext = { moduleVersion?: string; - self?: HuggingFaceTransformersPipeline; + self?: unknown; + pipeline?: HuggingFaceTransformersPipeline; }; type HuggingFaceTransformersPipelineInfo = { @@ -54,26 +55,23 @@ export function getHuggingFaceTransformersPipelineInfo( return pipeline ? pipelineInfo.get(pipeline) : undefined; } -export const huggingFaceTransformersChannels = defineChannels( +export const huggingFaceTransformersChannels = defineInterceptor( "@huggingface/transformers", { pipeline: channel< [string, (string | null)?, Record?], - HuggingFaceTransformersPipeline, + PromiseLike, HuggingFaceTransformersEventContext >({ channelName: "pipeline", - kind: "async", }), pipelineCall: channel< [unknown, ...unknown[]], - unknown | HuggingFaceTransformersTensor, + PromiseLike, HuggingFaceTransformersEventContext >({ channelName: "pipeline.call", - kind: "async", }), }, - { instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE }, ); diff --git a/js/src/instrumentation/plugins/huggingface-transformers-plugin.ts b/js/src/instrumentation/plugins/huggingface-transformers-plugin.ts index 8d41b535a..973b06716 100644 --- a/js/src/instrumentation/plugins/huggingface-transformers-plugin.ts +++ b/js/src/instrumentation/plugins/huggingface-transformers-plugin.ts @@ -1,9 +1,11 @@ +import { INSTRUMENTATION_NAMES } from "../../span-origin"; import { BasePlugin } from "../core"; -import { traceAsyncChannel, unsubscribeAll } from "../core/channel-tracing"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; +import { traceAsyncCall, unsubscribeAll } from "../core/channel-tracing"; +import { observeResult } from "../core/observe-result"; + import { SpanTypeAttribute, isObject } from "../../../util"; import type { HuggingFaceTransformersPipeline } from "../../vendor-sdk-types/huggingface-transformers"; +import type { ChannelMessage } from "../core/tracing-types"; import { getHuggingFaceTransformersPipelineInfo, huggingFaceTransformersChannels, @@ -23,34 +25,44 @@ export class HuggingFaceTransformersPlugin extends BasePlugin { protected onEnable(): void { this.subscribeToPipelineFactory(); this.unsubscribers.push( - traceAsyncChannel(huggingFaceTransformersChannels.pipelineCall, { - name: (_args, event) => { - const task = getTask(event as HuggingFaceTransformersEventContext); - const operation = task?.replaceAll("-", "_") ?? "unknown"; - return `huggingface.transformers.${operation}`; - }, - type: SpanTypeAttribute.LLM, - shouldTrace: (_args, event) => - isSupportedHuggingFaceTransformersTask( - getTask(event as HuggingFaceTransformersEventContext), - ), - extractInput: (args, event) => ({ - input: extractInput( - getTask(event as HuggingFaceTransformersEventContext), - args, + huggingFaceTransformersChannels.pipelineCall.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.HUGGINGFACE, + name: (_args, event) => { + const task = getTask( + event as HuggingFaceTransformersEventContext, + ); + const operation = task?.replaceAll("-", "_") ?? "unknown"; + return `huggingface.transformers.${operation}`; + }, + type: SpanTypeAttribute.LLM, + shouldTrace: (_args, event) => + isSupportedHuggingFaceTransformersTask( + getTask(event as HuggingFaceTransformersEventContext), + ), + extractInput: (args, event) => ({ + input: extractInput( + getTask(event as HuggingFaceTransformersEventContext), + args, + ), + metadata: extractMetadata( + event as HuggingFaceTransformersEventContext, + args, + ), + }), + extractOutput: (result, event) => + extractOutput( + getTask(event as HuggingFaceTransformersEventContext), + result, + ), + extractMetrics: () => ({}), + }, ), - metadata: extractMetadata( - event as HuggingFaceTransformersEventContext, - args, - ), - }), - extractOutput: (result, event) => - extractOutput( - getTask(event as HuggingFaceTransformersEventContext), - result, - ), - extractMetrics: () => ({}), - }), + ), ); } @@ -59,34 +71,55 @@ export class HuggingFaceTransformersPlugin extends BasePlugin { } private subscribeToPipelineFactory(): void { - const channel = - huggingFaceTransformersChannels.pipeline.tracingChannel() as IsoTracingChannel< - ChannelMessage - >; - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - asyncEnd: (event) => { - if (typeof event.result !== "function") { - return; + const channel = huggingFaceTransformersChannels.pipeline; + + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const resolved = ( + event: ChannelMessage< + typeof huggingFaceTransformersChannels.pipeline + >, + ) => { + if (typeof event.result !== "function") { + return; + } + registerHuggingFaceTransformersPipeline( + event.result, + event.arguments?.[0], + event.arguments?.[1], + ); + }; + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + throw error; } - registerHuggingFaceTransformersPipeline( - event.result, - event.arguments?.[0], - event.arguments?.[1], + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + resolved(event); + }, + (error) => {}, ); }, - }; - - channel.subscribe(handlers); - this.unsubscribers.push(() => channel.unsubscribe(handlers)); + ); + this.unsubscribers.push(removeHandlers); } } function getTask( event: HuggingFaceTransformersEventContext, ): string | undefined { - const self = event.self; + const self = + event.pipeline ?? + (event.self as HuggingFaceTransformersPipeline | undefined); const registeredTask = getHuggingFaceTransformersPipelineInfo(self)?.task; if (registeredTask !== undefined) { return registeredTask; @@ -105,9 +138,15 @@ function extractMetadata( provider: "huggingface", }; const registeredModel = getHuggingFaceTransformersPipelineInfo( - event.self, + event.pipeline ?? + (event.self as HuggingFaceTransformersPipeline | undefined), )?.model; - const model = registeredModel ?? modelIdentifier(event.self); + const model = + registeredModel ?? + modelIdentifier( + event.pipeline ?? + (event.self as HuggingFaceTransformersPipeline | undefined), + ); if (model) { metadata.model = model; } diff --git a/js/src/instrumentation/plugins/instrumentation-names.test.ts b/js/src/instrumentation/plugins/instrumentation-names.test.ts index fef4d8bfb..6f0090dac 100644 --- a/js/src/instrumentation/plugins/instrumentation-names.test.ts +++ b/js/src/instrumentation/plugins/instrumentation-names.test.ts @@ -9,8 +9,8 @@ import { smithyCoreChannels, } from "./bedrock-runtime-channels"; import { claudeAgentSDKChannels } from "./claude-agent-sdk-channels"; -import { cloudflareAIChatChannels } from "./cloudflare-ai-chat-channels"; import { cloudflareAgentsChannels } from "./cloudflare-agents-channels"; +import { cloudflareAIChatChannels } from "./cloudflare-ai-chat-channels"; import { cloudflareThinkChannels } from "./cloudflare-think-channels"; import { cohereChannels } from "./cohere-channels"; import { cursorSDKChannels } from "./cursor-sdk-channels"; @@ -18,12 +18,12 @@ import { flueChannels } from "./flue-channels"; import { genkitChannels, genkitCoreChannels } from "./genkit-channels"; import { gitHubCopilotChannels } from "./github-copilot-channels"; import { googleADKChannels } from "./google-adk-channels"; -import { googleGenerativeAIChannels } from "./google-generative-ai-channels"; import { googleGenAIChannels } from "./google-genai-channels"; +import { googleGenerativeAIChannels } from "./google-generative-ai-channels"; import { groqChannels } from "./groq-channels"; import { huggingFaceChannels } from "./huggingface-channels"; -import { langGraphSDKChannels } from "./langgraph-sdk-channels"; import { langChainChannels } from "./langchain-channels"; +import { langGraphSDKChannels } from "./langgraph-sdk-channels"; import { langSmithChannels } from "./langsmith-channels"; import { mistralChannels } from "./mistral-channels"; import { ollamaChannels } from "./ollama-channels"; @@ -35,7 +35,7 @@ import { openRouterChannels } from "./openrouter-channels"; import { piCodingAgentChannels } from "./pi-coding-agent-channels"; import { strandsAgentSDKChannels } from "./strands-agent-sdk-channels"; -describe("built-in instrumentation provenance names", () => { +describe("wrapping definitions are independent of span provenance", () => { it.each([ [aiSDKChannels.generateText, INSTRUMENTATION_NAMES.AI_SDK], [harnessAgentChannels.generate, INSTRUMENTATION_NAMES.AI_SDK], @@ -88,7 +88,8 @@ describe("built-in instrumentation provenance names", () => { strandsAgentSDKChannels.agentStream, INSTRUMENTATION_NAMES.STRANDS_AGENT_SDK, ], - ])("uses %s for its canonical channel group", (channel, expected) => { - expect(channel.instrumentationName).toBe(expected); + ])("keeps span configuration out of %s", (channel, _instrumentationName) => { + expect(channel).not.toHaveProperty("instrumentationName"); + expect(channel).not.toHaveProperty("tracingChannel"); }); }); diff --git a/js/src/instrumentation/plugins/langchain-channels.ts b/js/src/instrumentation/plugins/langchain-channels.ts index 8ab3f2356..c8636b4db 100644 --- a/js/src/instrumentation/plugins/langchain-channels.ts +++ b/js/src/instrumentation/plugins/langchain-channels.ts @@ -1,27 +1,21 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { LangChainCallbackManagerConfigureArgs, LangChainCallbackManagerConfigureResult, } from "../../vendor-sdk-types/langchain"; -export const langChainChannels = defineChannels( - "@langchain/core", - { - configure: channel< - LangChainCallbackManagerConfigureArgs, - LangChainCallbackManagerConfigureResult - >({ - channelName: "CallbackManager.configure", - kind: "sync-stream", - }), - configureSync: channel< - LangChainCallbackManagerConfigureArgs, - LangChainCallbackManagerConfigureResult - >({ - channelName: "CallbackManager._configureSync", - kind: "sync-stream", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.LANGCHAIN }, -); +export const langChainChannels = defineInterceptor("@langchain/core", { + configure: channel< + LangChainCallbackManagerConfigureArgs, + LangChainCallbackManagerConfigureResult + >({ + channelName: "CallbackManager.configure", + }), + configureSync: channel< + LangChainCallbackManagerConfigureArgs, + LangChainCallbackManagerConfigureResult + >({ + channelName: "CallbackManager._configureSync", + }), +}); diff --git a/js/src/instrumentation/plugins/langchain-plugin.test.ts b/js/src/instrumentation/plugins/langchain-plugin.test.ts index 29ace20ab..89bd0ab31 100644 --- a/js/src/instrumentation/plugins/langchain-plugin.test.ts +++ b/js/src/instrumentation/plugins/langchain-plugin.test.ts @@ -1,6 +1,6 @@ import { describe, expect, it } from "vitest"; -import { LangChainPlugin } from "./langchain-plugin"; import { langChainChannels } from "./langchain-channels"; +import { LangChainPlugin } from "./langchain-plugin"; function createManager(handlers: unknown[] = []) { return { @@ -12,21 +12,30 @@ function createManager(handlers: unknown[] = []) { } function traceConfigureResult(result: unknown) { - return langChainChannels.configure.traceSync(() => result as any, { - arguments: [], - }); + return langChainChannels.configure.invoke( + () => result as any, + undefined, + [], + {}, + ); } - function traceConfigureArguments(args: unknown[]) { - return langChainChannels.configure.traceSync(() => args as any, { - arguments: args as any, - }); + return langChainChannels.configure.invoke( + (...received: unknown[]) => received as any, + undefined, + args as any, + {}, + ); } - function traceConfigureArgumentsObject(args: IArguments) { - return langChainChannels.configure.traceSync(() => args as any, { - arguments: args as any, - }); + return langChainChannels.configure.invoke( + function (..._received: unknown[]) { + return arguments as any; + }, + undefined, + args as any, + {}, + ); } function createArgumentsObject(...args: unknown[]): IArguments { @@ -41,10 +50,10 @@ describe("LangChainPlugin", () => { const args: unknown[] = []; plugin.enable(); - traceConfigureArguments(args); + const received = traceConfigureArguments(args); plugin.disable(); - expect(args[0]).toEqual([ + expect(received[0]).toEqual([ expect.objectContaining({ name: "BraintrustCallbackHandler", }), @@ -56,10 +65,10 @@ describe("LangChainPlugin", () => { const args = createArgumentsObject(); plugin.enable(); - traceConfigureArgumentsObject(args); + const received = traceConfigureArgumentsObject(args); plugin.disable(); - expect(args[0]).toEqual([ + expect(received[0]).toEqual([ expect.objectContaining({ name: "BraintrustCallbackHandler", }), diff --git a/js/src/instrumentation/plugins/langchain-plugin.ts b/js/src/instrumentation/plugins/langchain-plugin.ts index 85ef66083..b6cae5384 100644 --- a/js/src/instrumentation/plugins/langchain-plugin.ts +++ b/js/src/instrumentation/plugins/langchain-plugin.ts @@ -1,11 +1,12 @@ import { BasePlugin } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers, IsoTracingChannel } from "../../isomorph"; +import { runInstrumentation } from "../core/observe-result"; + import type { LangChainCallbackManager } from "../../vendor-sdk-types/langchain"; import { BRAINTRUST_LANGCHAIN_CALLBACK_HANDLER_NAME, BraintrustLangChainCallbackHandler, } from "../../wrappers/langchain/callback-handler"; +import type { ChannelMessage } from "../core/tracing-types"; import { langChainChannels } from "./langchain-channels"; type LangChainConfigureChannel = @@ -29,25 +30,34 @@ export class LangChainPlugin extends BasePlugin { } private subscribeToConfigure(channel: LangChainConfigureChannel): void { - const tracingChannel: IsoTracingChannel< - ChannelMessage - > = channel.tracingChannel(); - - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - injectHandlerIntoArguments(event.arguments); - }, - end: (event) => { - this.injectHandler(event.result); + const invocationHook = channel; + + const removeHandlers = invocationHook.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = (event: ChannelMessage) => { + injectHandlerIntoArguments(event.arguments); + }; + const returned = (event: ChannelMessage) => { + this.injectHandler(event.result); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + throw error; + } + Object.assign(event, { result }); + runInstrumentation(() => returned(event)); + return result; }, - }; - - tracingChannel.subscribe(handlers); - this.unsubscribers.push(() => { - tracingChannel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private injectHandler(result: unknown): void { diff --git a/js/src/instrumentation/plugins/langgraph-sdk-channels.ts b/js/src/instrumentation/plugins/langgraph-sdk-channels.ts index 1329982a5..6e43143cf 100644 --- a/js/src/instrumentation/plugins/langgraph-sdk-channels.ts +++ b/js/src/instrumentation/plugins/langgraph-sdk-channels.ts @@ -1,21 +1,17 @@ -import { INSTRUMENTATION_NAMES } from "../../span-origin"; import type { LangGraphRunArgs, LangGraphStreamEvent, } from "../../vendor-sdk-types/langgraph-sdk"; -import { channel, defineChannels } from "../core/channel-definitions"; +import { channel, defineInterceptor } from "../core/channel-definitions"; -export const langGraphSDKChannels = defineChannels( +export const langGraphSDKChannels = defineInterceptor( "@langchain/langgraph-sdk", { - wait: channel({ + wait: channel>({ channelName: "runs.wait", - kind: "async", }), stream: channel>({ channelName: "runs.stream", - kind: "sync-stream", }), }, - { instrumentationName: INSTRUMENTATION_NAMES.LANGGRAPH_SDK }, ); diff --git a/js/src/instrumentation/plugins/langsmith-channels.ts b/js/src/instrumentation/plugins/langsmith-channels.ts index ff49a7e8e..6ea403c64 100644 --- a/js/src/instrumentation/plugins/langsmith-channels.ts +++ b/js/src/instrumentation/plugins/langsmith-channels.ts @@ -1,35 +1,30 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { LangSmithBatchIngestRuns, LangSmithClient, LangSmithRun, } from "../../vendor-sdk-types/langsmith"; -export const langSmithChannels = defineChannels( - "langsmith", - { - createRun: channel< - [run: LangSmithRun, options?: unknown], - Awaited>> - >({ - channelName: "Client.createRun", - kind: "async", - }), - updateRun: channel< - [runId: string, run: LangSmithRun, options?: unknown], - Awaited>> - >({ - channelName: "Client.updateRun", - kind: "async", - }), - batchIngestRuns: channel< - [runs: LangSmithBatchIngestRuns, options?: unknown], +export const langSmithChannels = defineInterceptor("langsmith", { + createRun: channel< + [run: LangSmithRun, options?: unknown], + PromiseLike>>> + >({ + channelName: "Client.createRun", + }), + updateRun: channel< + [runId: string, run: LangSmithRun, options?: unknown], + PromiseLike>>> + >({ + channelName: "Client.updateRun", + }), + batchIngestRuns: channel< + [runs: LangSmithBatchIngestRuns, options?: unknown], + PromiseLike< Awaited>> - >({ - channelName: "Client.batchIngestRuns", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.LANGSMITH }, -); + > + >({ + channelName: "Client.batchIngestRuns", + }), +}); diff --git a/js/src/instrumentation/plugins/langsmith-plugin.test.ts b/js/src/instrumentation/plugins/langsmith-plugin.test.ts index 8e9b3aed4..bebca7f53 100644 --- a/js/src/instrumentation/plugins/langsmith-plugin.test.ts +++ b/js/src/instrumentation/plugins/langsmith-plugin.test.ts @@ -1,6 +1,6 @@ import { afterEach, beforeAll, beforeEach, describe, expect, it } from "vitest"; -import { configureNode } from "../../node/config"; import { _exportsForTestingOnly, initLogger } from "../../logger"; +import { configureNode } from "../../node/config"; import { langSmithChannels } from "./langsmith-channels"; import { LangSmithPlugin } from "./langsmith-plugin"; @@ -42,73 +42,73 @@ describe("LangSmithPlugin", () => { const firstToken = new Date(start.getTime() + 250); const end = new Date(start.getTime() + 500); - await langSmithChannels.batchIngestRuns.tracePromise( + await langSmithChannels.batchIngestRuns.invoke( async () => undefined, - { - arguments: [ - { - runCreates: [ - { - id: childId, - trace_id: rootId, - parent_run_id: rootId, - dotted_order: `20260713T000000000000Z${rootId}.20260713T000000000001Z${childId}`, - name: "answer", - run_type: "llm", - start_time: start.toISOString(), - inputs: { messages: ["hello"] }, - extra: { - metadata: { - customer: "acme", - ls_provider: "openai", - ls_model_name: "gpt-test", - ls_temperature: 0.2, - usage_metadata: { ignored: true }, - }, - runtime: { hidden: true }, + undefined, + [ + { + runCreates: [ + { + id: childId, + trace_id: rootId, + parent_run_id: rootId, + dotted_order: `20260713T000000000000Z${rootId}.20260713T000000000001Z${childId}`, + name: "answer", + run_type: "llm", + start_time: start.toISOString(), + inputs: { messages: ["hello"] }, + extra: { + metadata: { + customer: "acme", + ls_provider: "openai", + ls_model_name: "gpt-test", + ls_temperature: 0.2, + usage_metadata: { ignored: true }, }, - tags: ["unit", 1], - serialized: { secret: true }, - events: [{ name: "new_token", time: firstToken.toISOString() }], + runtime: { hidden: true }, }, - { - id: rootId, - trace_id: rootId, - name: "workflow", - run_type: "chain", - start_time: start.toISOString(), - inputs: { question: "hello" }, - }, - ], - runUpdates: [ - { - id: childId, - trace_id: rootId, - parent_run_id: rootId, - end_time: end.toISOString(), - outputs: { - generations: ["world"], - usage_metadata: { - input_tokens: 3, - output_tokens: 2, - total_tokens: 5, - input_token_details: { - cache_read: 1, - cache_creation: 2, - }, + tags: ["unit", 1], + serialized: { secret: true }, + events: [{ name: "new_token", time: firstToken.toISOString() }], + }, + { + id: rootId, + trace_id: rootId, + name: "workflow", + run_type: "chain", + start_time: start.toISOString(), + inputs: { question: "hello" }, + }, + ], + runUpdates: [ + { + id: childId, + trace_id: rootId, + parent_run_id: rootId, + end_time: end.toISOString(), + outputs: { + generations: ["world"], + usage_metadata: { + input_tokens: 3, + output_tokens: 2, + total_tokens: 5, + input_token_details: { + cache_read: 1, + cache_creation: 2, }, }, }, - { - id: rootId, - trace_id: rootId, - end_time: end.toISOString(), - outputs: { answer: "world" }, - }, - ], - }, - ], - }, + }, + { + id: rootId, + trace_id: rootId, + end_time: end.toISOString(), + outputs: { answer: "world" }, + }, + ], + }, + ], + {}, ); const spans = (await backgroundLogger.drain()) as any[]; @@ -168,15 +168,23 @@ describe("LangSmithPlugin", () => { end_time: new Date().toISOString(), }; - await langSmithChannels.updateRun.tracePromise(async () => undefined, { - arguments: [id, update], - }); - await langSmithChannels.createRun.tracePromise(async () => undefined, { - arguments: [update], - }); - await langSmithChannels.batchIngestRuns.tracePromise( + await langSmithChannels.updateRun.invoke( + async () => undefined, + undefined, + [id, update], + {}, + ); + await langSmithChannels.createRun.invoke( + async () => undefined, + undefined, + [update], + {}, + ); + await langSmithChannels.batchIngestRuns.invoke( async () => undefined, - { arguments: [{ runUpdates: [update] }] }, + undefined, + [{ runUpdates: [update] }], + {}, ); const spans = (await backgroundLogger.drain()) as any[]; @@ -192,8 +200,10 @@ describe("LangSmithPlugin", () => { const batchId = "66666666-6666-4666-8666-666666666666"; const endTime = new Date().toISOString(); - await langSmithChannels.createRun.tracePromise(async () => undefined, { - arguments: [ + await langSmithChannels.createRun.invoke( + async () => undefined, + undefined, + [ { id: directId, trace_id: directId, @@ -202,24 +212,25 @@ describe("LangSmithPlugin", () => { end_time: endTime, }, ], - }); - await langSmithChannels.batchIngestRuns.tracePromise( + {}, + ); + await langSmithChannels.batchIngestRuns.invoke( async () => undefined, - { - arguments: [ - { - runCreates: [ - { - id: batchId, - trace_id: batchId, - name: "completed batch create", - outputs: { answer: "batch" }, - end_time: endTime, - }, - ], - }, - ], - }, + undefined, + [ + { + runCreates: [ + { + id: batchId, + trace_id: batchId, + name: "completed batch create", + outputs: { answer: "batch" }, + end_time: endTime, + }, + ], + }, + ], + {}, ); const spans = (await backgroundLogger.drain()) as any[]; @@ -234,8 +245,10 @@ describe("LangSmithPlugin", () => { it("completes runs containing invalid nested dates", async () => { const id = "77777777-7777-4777-8777-777777777777"; - await langSmithChannels.createRun.tracePromise(async () => undefined, { - arguments: [ + await langSmithChannels.createRun.invoke( + async () => undefined, + undefined, + [ { id, trace_id: id, @@ -244,7 +257,8 @@ describe("LangSmithPlugin", () => { end_time: new Date().toISOString(), }, ], - }); + {}, + ); const spans = (await backgroundLogger.drain()) as any[]; expect(spans.filter((span) => span.span_id === id)).toHaveLength(1); @@ -261,9 +275,12 @@ describe("LangSmithPlugin", () => { }); await expect( - langSmithChannels.createRun.tracePromise(async () => undefined, { - arguments: [payload], - }), + langSmithChannels.createRun.invoke( + async () => undefined, + undefined, + [payload], + {}, + ), ).resolves.toBeUndefined(); expect(getterCalled).toBe(false); expect(await backgroundLogger.drain()).toEqual([]); @@ -278,17 +295,23 @@ describe("LangSmithPlugin", () => { end_time: new Date().toISOString(), }; - await langSmithChannels.updateRun.tracePromise(async () => undefined, { - arguments: [run.id, run], - }); + await langSmithChannels.updateRun.invoke( + async () => undefined, + undefined, + [run.id, run], + {}, + ); expect(await backgroundLogger.drain()).toEqual([]); plugin.disable(); plugin = new LangSmithPlugin({ skipLangChainRuns: false }); plugin.enable(); - await langSmithChannels.updateRun.tracePromise(async () => undefined, { - arguments: [run.id, run], - }); + await langSmithChannels.updateRun.invoke( + async () => undefined, + undefined, + [run.id, run], + {}, + ); expect(await backgroundLogger.drain()).toHaveLength(1); }); diff --git a/js/src/instrumentation/plugins/langsmith-plugin.ts b/js/src/instrumentation/plugins/langsmith-plugin.ts index 5027fa694..ef593eb4f 100644 --- a/js/src/instrumentation/plugins/langsmith-plugin.ts +++ b/js/src/instrumentation/plugins/langsmith-plugin.ts @@ -1,19 +1,21 @@ import { SpanTypeAttribute } from "../../../util/index"; import { debugLogger } from "../../debug-logger"; -import { startSpan as startBaseSpan } from "../../logger"; import type { Span } from "../../logger"; +import { startSpan as startBaseSpan } from "../../logger"; +import { LRUCache } from "../../lru-cache"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; -import { LRUCache } from "../../lru-cache"; import type { LangSmithBatchIngestRuns, LangSmithRun, } from "../../vendor-sdk-types/langsmith"; import { BasePlugin } from "../core"; import { unsubscribeAll } from "../core/channel-tracing"; -import type { ChannelMessage } from "../core/channel-definitions"; +import { runInstrumentation } from "../core/observe-result"; + +import type { ChannelMessage } from "../core/tracing-types"; import { langSmithChannels } from "./langsmith-channels"; type ActiveRun = { @@ -56,44 +58,64 @@ export class LangSmithPlugin extends BasePlugin { } protected onEnable(): void { - const createChannel = langSmithChannels.createRun.tracingChannel(); - const createHandlers = { - start: ( - event: ChannelMessage, - ): void => { - this.containLifecycleFailure("createRun", () => { - this.processCreate(event.arguments[0]); - }); + const createChannel = langSmithChannels.createRun; + + const removecreateHandlers = createChannel.intercept( + (target, receiver, args, additional) => { + runInstrumentation(() => + (( + event: ChannelMessage, + ): void => { + this.containLifecycleFailure("createRun", () => { + this.processCreate(additional.runTree ?? event.arguments[0]); + }); + })({ arguments: args }), + ); + return Reflect.apply(target, receiver, args); }, - }; - createChannel.subscribe(createHandlers); - this.unsubscribers.push(() => createChannel.unsubscribe(createHandlers)); - - const updateChannel = langSmithChannels.updateRun.tracingChannel(); - const updateHandlers = { - start: ( - event: ChannelMessage, - ): void => { - this.containLifecycleFailure("updateRun", () => { - this.processUpdate(event.arguments[0], event.arguments[1]); - }); + ); + this.unsubscribers.push(removecreateHandlers); + + const updateChannel = langSmithChannels.updateRun; + + const removeupdateHandlers = updateChannel.intercept( + (target, receiver, args, additional) => { + runInstrumentation(() => + (( + event: ChannelMessage, + ): void => { + this.containLifecycleFailure("updateRun", () => { + this.processUpdate( + additional.runTree + ? ((additional.runTree as LangSmithRun).id as string) + : event.arguments[0], + additional.runTree ?? event.arguments[1], + ); + }); + })({ arguments: args }), + ); + return Reflect.apply(target, receiver, args); }, - }; - updateChannel.subscribe(updateHandlers); - this.unsubscribers.push(() => updateChannel.unsubscribe(updateHandlers)); - - const batchChannel = langSmithChannels.batchIngestRuns.tracingChannel(); - const batchHandlers = { - start: ( - event: ChannelMessage, - ): void => { - this.containLifecycleFailure("batchIngestRuns", () => { - this.processBatch(event.arguments[0]); - }); + ); + this.unsubscribers.push(removeupdateHandlers); + + const batchChannel = langSmithChannels.batchIngestRuns; + + const removebatchHandlers = batchChannel.intercept( + (target, receiver, args, additional) => { + runInstrumentation(() => + (( + event: ChannelMessage, + ): void => { + this.containLifecycleFailure("batchIngestRuns", () => { + this.processBatch(event.arguments[0]); + }); + })({ arguments: args }), + ); + return Reflect.apply(target, receiver, args); }, - }; - batchChannel.subscribe(batchHandlers); - this.unsubscribers.push(() => batchChannel.unsubscribe(batchHandlers)); + ); + this.unsubscribers.push(removebatchHandlers); } protected onDisable(): void { diff --git a/js/src/instrumentation/plugins/mistral-channels.ts b/js/src/instrumentation/plugins/mistral-channels.ts index 613c542cd..9055b01af 100644 --- a/js/src/instrumentation/plugins/mistral-channels.ts +++ b/js/src/instrumentation/plugins/mistral-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { MistralAgentsCompletionEvent, MistralAgentsCompletionResponse, @@ -21,102 +21,87 @@ import type { MistralModerationResponse, } from "../../vendor-sdk-types/mistral"; -export const mistralChannels = defineChannels( - "@mistralai/mistralai", - { - chatComplete: channel< - [MistralChatCreateParams], - MistralChatCompletionResponse - >({ - channelName: "chat.complete", - kind: "async", - }), +export const mistralChannels = defineInterceptor("@mistralai/mistralai", { + chatComplete: channel< + [MistralChatCreateParams], + PromiseLike + >({ + channelName: "chat.complete", + }), - chatStream: channel< - [MistralChatCreateParams], - MistralChatResult, - Record, - MistralChatCompletionEvent - >({ - channelName: "chat.stream", - kind: "async", - }), + chatStream: channel< + [MistralChatCreateParams], + PromiseLike, + Record, + MistralChatCompletionEvent + >({ + channelName: "chat.stream", + }), - embeddingsCreate: channel< - [MistralEmbeddingCreateParams], - MistralEmbeddingResponse - >({ - channelName: "embeddings.create", - kind: "async", - }), + embeddingsCreate: channel< + [MistralEmbeddingCreateParams], + PromiseLike + >({ + channelName: "embeddings.create", + }), - classifiersModerate: channel< - [MistralClassificationCreateParams], - MistralModerationResponse - >({ - channelName: "classifiers.moderate", - kind: "async", - }), + classifiersModerate: channel< + [MistralClassificationCreateParams], + PromiseLike + >({ + channelName: "classifiers.moderate", + }), - classifiersModerateChat: channel< - [MistralChatClassificationCreateParams], - MistralModerationResponse - >({ - channelName: "classifiers.moderateChat", - kind: "async", - }), + classifiersModerateChat: channel< + [MistralChatClassificationCreateParams], + PromiseLike + >({ + channelName: "classifiers.moderateChat", + }), - classifiersClassify: channel< - [MistralClassificationCreateParams], - MistralClassificationResponse - >({ - channelName: "classifiers.classify", - kind: "async", - }), + classifiersClassify: channel< + [MistralClassificationCreateParams], + PromiseLike + >({ + channelName: "classifiers.classify", + }), - classifiersClassifyChat: channel< - [MistralChatClassificationCreateParams], - MistralClassificationResponse - >({ - channelName: "classifiers.classifyChat", - kind: "async", - }), + classifiersClassifyChat: channel< + [MistralChatClassificationCreateParams], + PromiseLike + >({ + channelName: "classifiers.classifyChat", + }), - fimComplete: channel< - [MistralFimCreateParams], - MistralFimCompletionResponse - >({ - channelName: "fim.complete", - kind: "async", - }), + fimComplete: channel< + [MistralFimCreateParams], + PromiseLike + >({ + channelName: "fim.complete", + }), - fimStream: channel< - [MistralFimCreateParams], - MistralFimResult, - Record, - MistralFimCompletionEvent - >({ - channelName: "fim.stream", - kind: "async", - }), + fimStream: channel< + [MistralFimCreateParams], + PromiseLike, + Record, + MistralFimCompletionEvent + >({ + channelName: "fim.stream", + }), - agentsComplete: channel< - [MistralAgentsCreateParams], - MistralAgentsCompletionResponse - >({ - channelName: "agents.complete", - kind: "async", - }), + agentsComplete: channel< + [MistralAgentsCreateParams], + PromiseLike + >({ + channelName: "agents.complete", + }), - agentsStream: channel< - [MistralAgentsCreateParams], - MistralAgentsResult, - Record, - MistralAgentsCompletionEvent - >({ - channelName: "agents.stream", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.MISTRAL }, -); + agentsStream: channel< + [MistralAgentsCreateParams], + PromiseLike, + Record, + MistralAgentsCompletionEvent + >({ + channelName: "agents.stream", + }), +}); diff --git a/js/src/instrumentation/plugins/mistral-plugin.ts b/js/src/instrumentation/plugins/mistral-plugin.ts index f33f82554..acdfd1fdb 100644 --- a/js/src/instrumentation/plugins/mistral-plugin.ts +++ b/js/src/instrumentation/plugins/mistral-plugin.ts @@ -1,13 +1,6 @@ -import { BasePlugin } from "../core"; -import { - traceAsyncChannel, - traceStreamingChannel, - unsubscribeAll, -} from "../core/channel-tracing"; import { SpanTypeAttribute, isObject } from "../../../util/index"; -import { processInputAttachments } from "../../wrappers/attachment-utils"; +import { INSTRUMENTATION_NAMES } from "../../span-origin"; import { getCurrentUnixTimestamp } from "../../util"; -import { mistralChannels } from "./mistral-channels"; import type { MistralChatCompletionChunk, MistralChatCompletionChunkChoice, @@ -18,6 +11,14 @@ import type { MistralThinkingContentPart, MistralToolCallDelta, } from "../../vendor-sdk-types/mistral"; +import { processInputAttachments } from "../../wrappers/attachment-utils"; +import { BasePlugin } from "../core"; +import { + traceAsyncCall, + traceStreamingCall, + unsubscribeAll, +} from "../core/channel-tracing"; +import { mistralChannels } from "./mistral-channels"; export class MistralPlugin extends BasePlugin { protected onEnable(): void { @@ -30,144 +31,248 @@ export class MistralPlugin extends BasePlugin { private subscribeToMistralChannels(): void { this.unsubscribers.push( - traceStreamingChannel(mistralChannels.chatComplete, { - name: "mistral.chat.complete", - type: SpanTypeAttribute.LLM, - extractInput: extractMessagesInputWithMetadata, - extractOutput: (result) => { - return result?.choices; - }, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result, startTime) => - extractMistralMetrics(result?.usage, startTime), - }), + mistralChannels.chatComplete.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.chat.complete", + type: SpanTypeAttribute.LLM, + extractInput: extractMessagesInputWithMetadata, + extractOutput: (result) => { + return result?.choices; + }, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result, startTime) => + extractMistralMetrics(result?.usage, startTime), + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(mistralChannels.chatStream, { - name: "mistral.chat.stream", - type: SpanTypeAttribute.LLM, - extractInput: extractMessagesInputWithMetadata, - extractOutput: extractMistralStreamOutput, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result, startTime) => - extractMistralStreamingMetrics(result, startTime), - aggregateChunks: aggregateMistralStreamChunks, - }), + mistralChannels.chatStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.chat.stream", + type: SpanTypeAttribute.LLM, + extractInput: extractMessagesInputWithMetadata, + extractOutput: extractMistralStreamOutput, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result, startTime) => + extractMistralStreamingMetrics(result, startTime), + aggregateChunks: aggregateMistralStreamChunks, + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(mistralChannels.embeddingsCreate, { - name: "mistral.embeddings.create", - type: SpanTypeAttribute.LLM, - extractInput: extractEmbeddingInputWithMetadata, - extractOutput: (result) => { - const embedding = result?.data?.[0]?.embedding; - return Array.isArray(embedding) - ? { embedding_length: embedding.length } - : undefined; - }, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result) => parseMistralMetricsFromUsage(result?.usage), - }), + mistralChannels.embeddingsCreate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.embeddings.create", + type: SpanTypeAttribute.LLM, + extractInput: extractEmbeddingInputWithMetadata, + extractOutput: (result) => { + const embedding = result?.data?.[0]?.embedding; + return Array.isArray(embedding) + ? { embedding_length: embedding.length } + : undefined; + }, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result) => + parseMistralMetricsFromUsage(result?.usage), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(mistralChannels.classifiersModerate, { - name: "mistral.classifiers.moderate", - type: SpanTypeAttribute.LLM, - extractInput: extractClassifierInputWithMetadata, - extractOutput: extractClassifierOutput, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result) => parseMistralMetricsFromUsage(result?.usage), - }), + mistralChannels.classifiersModerate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.classifiers.moderate", + type: SpanTypeAttribute.LLM, + extractInput: extractClassifierInputWithMetadata, + extractOutput: extractClassifierOutput, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result) => + parseMistralMetricsFromUsage(result?.usage), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(mistralChannels.classifiersModerateChat, { - name: "mistral.classifiers.moderateChat", - type: SpanTypeAttribute.LLM, - extractInput: extractClassifierInputWithMetadata, - extractOutput: extractClassifierOutput, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result) => parseMistralMetricsFromUsage(result?.usage), - }), + mistralChannels.classifiersModerateChat.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.classifiers.moderateChat", + type: SpanTypeAttribute.LLM, + extractInput: extractClassifierInputWithMetadata, + extractOutput: extractClassifierOutput, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result) => + parseMistralMetricsFromUsage(result?.usage), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(mistralChannels.classifiersClassify, { - name: "mistral.classifiers.classify", - type: SpanTypeAttribute.LLM, - extractInput: extractClassifierInputWithMetadata, - extractOutput: extractClassifierOutput, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result) => parseMistralMetricsFromUsage(result?.usage), - }), + mistralChannels.classifiersClassify.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.classifiers.classify", + type: SpanTypeAttribute.LLM, + extractInput: extractClassifierInputWithMetadata, + extractOutput: extractClassifierOutput, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result) => + parseMistralMetricsFromUsage(result?.usage), + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(mistralChannels.classifiersClassifyChat, { - name: "mistral.classifiers.classifyChat", - type: SpanTypeAttribute.LLM, - extractInput: extractClassifierInputWithMetadata, - extractOutput: extractClassifierOutput, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result) => parseMistralMetricsFromUsage(result?.usage), - }), + mistralChannels.classifiersClassifyChat.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.classifiers.classifyChat", + type: SpanTypeAttribute.LLM, + extractInput: extractClassifierInputWithMetadata, + extractOutput: extractClassifierOutput, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result) => + parseMistralMetricsFromUsage(result?.usage), + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(mistralChannels.fimComplete, { - name: "mistral.fim.complete", - type: SpanTypeAttribute.LLM, - extractInput: extractPromptInputWithMetadata, - extractOutput: (result) => { - return result?.choices; - }, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result, startTime) => - extractMistralMetrics(result?.usage, startTime), - }), + mistralChannels.fimComplete.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.fim.complete", + type: SpanTypeAttribute.LLM, + extractInput: extractPromptInputWithMetadata, + extractOutput: (result) => { + return result?.choices; + }, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result, startTime) => + extractMistralMetrics(result?.usage, startTime), + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(mistralChannels.fimStream, { - name: "mistral.fim.stream", - type: SpanTypeAttribute.LLM, - extractInput: extractPromptInputWithMetadata, - extractOutput: extractMistralStreamOutput, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result, startTime) => - extractMistralStreamingMetrics(result, startTime), - aggregateChunks: aggregateMistralStreamChunks, - }), + mistralChannels.fimStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.fim.stream", + type: SpanTypeAttribute.LLM, + extractInput: extractPromptInputWithMetadata, + extractOutput: extractMistralStreamOutput, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result, startTime) => + extractMistralStreamingMetrics(result, startTime), + aggregateChunks: aggregateMistralStreamChunks, + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(mistralChannels.agentsComplete, { - name: "mistral.agents.complete", - type: SpanTypeAttribute.LLM, - extractInput: extractMessagesInputWithMetadata, - extractOutput: (result) => { - return result?.choices; - }, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result, startTime) => - extractMistralMetrics(result?.usage, startTime), - }), + mistralChannels.agentsComplete.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.agents.complete", + type: SpanTypeAttribute.LLM, + extractInput: extractMessagesInputWithMetadata, + extractOutput: (result) => { + return result?.choices; + }, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result, startTime) => + extractMistralMetrics(result?.usage, startTime), + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(mistralChannels.agentsStream, { - name: "mistral.agents.stream", - type: SpanTypeAttribute.LLM, - extractInput: extractMessagesInputWithMetadata, - extractOutput: extractMistralStreamOutput, - extractMetadata: (result) => extractMistralResponseMetadata(result), - extractMetrics: (result, startTime) => - extractMistralStreamingMetrics(result, startTime), - aggregateChunks: aggregateMistralStreamChunks, - }), + mistralChannels.agentsStream.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.MISTRAL, + name: "mistral.agents.stream", + type: SpanTypeAttribute.LLM, + extractInput: extractMessagesInputWithMetadata, + extractOutput: extractMistralStreamOutput, + extractMetadata: (result) => + extractMistralResponseMetadata(result), + extractMetrics: (result, startTime) => + extractMistralStreamingMetrics(result, startTime), + aggregateChunks: aggregateMistralStreamChunks, + }, + ), + ), ); } } diff --git a/js/src/instrumentation/plugins/ollama-channels.ts b/js/src/instrumentation/plugins/ollama-channels.ts index c6e4f2044..cf2472cd2 100644 --- a/js/src/instrumentation/plugins/ollama-channels.ts +++ b/js/src/instrumentation/plugins/ollama-channels.ts @@ -1,4 +1,3 @@ -import { INSTRUMENTATION_NAMES } from "../../span-origin"; import type { OllamaChatRequest, OllamaChatResponse, @@ -9,33 +8,26 @@ import type { OllamaGenerateResponse, OllamaGenerateResult, } from "../../vendor-sdk-types/ollama"; -import { channel, defineChannels } from "../core/channel-definitions"; +import { channel, defineInterceptor } from "../core/channel-definitions"; -export const ollamaChannels = defineChannels( - "ollama", - { - chat: channel< - [OllamaChatRequest], - OllamaChatResult, - Record, - OllamaChatResponse - >({ - channelName: "chat", - kind: "async", - }), - generate: channel< - [OllamaGenerateRequest], - OllamaGenerateResult, - Record, - OllamaGenerateResponse - >({ - channelName: "generate", - kind: "async", - }), - embed: channel<[OllamaEmbedRequest], OllamaEmbedResponse>({ - channelName: "embed", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.OLLAMA }, -); +export const ollamaChannels = defineInterceptor("ollama", { + chat: channel< + [OllamaChatRequest], + PromiseLike, + Record, + OllamaChatResponse + >({ + channelName: "chat", + }), + generate: channel< + [OllamaGenerateRequest], + PromiseLike, + Record, + OllamaGenerateResponse + >({ + channelName: "generate", + }), + embed: channel<[OllamaEmbedRequest], PromiseLike>({ + channelName: "embed", + }), +}); diff --git a/js/src/instrumentation/plugins/ollama-plugin.ts b/js/src/instrumentation/plugins/ollama-plugin.ts index f48040348..3413f51d3 100644 --- a/js/src/instrumentation/plugins/ollama-plugin.ts +++ b/js/src/instrumentation/plugins/ollama-plugin.ts @@ -1,7 +1,7 @@ import { SpanTypeAttribute, isObject } from "../../../util/index"; import iso from "../../isomorph"; import { Attachment } from "../../logger"; -import { processInputAttachments } from "../../wrappers/attachment-utils"; +import { INSTRUMENTATION_NAMES } from "../../span-origin"; import type { OllamaChatRequest, OllamaChatResponse, @@ -14,48 +14,71 @@ import type { OllamaToolCall, OllamaUsageResponse, } from "../../vendor-sdk-types/ollama"; +import { processInputAttachments } from "../../wrappers/attachment-utils"; import { BasePlugin } from "../core"; -import type { AsyncEndOf } from "../core/channel-definitions"; + import { - traceAsyncChannel, - traceStreamingChannel, + traceAsyncCall, + traceStreamingCall, unsubscribeAll, } from "../core/channel-tracing"; +import type { AsyncEndOf } from "../core/tracing-types"; import { ollamaChannels } from "./ollama-channels"; export class OllamaPlugin extends BasePlugin { protected onEnable(): void { this.unsubscribers.push( - traceStreamingChannel(ollamaChannels.chat, { - name: "ollama.chat", - type: SpanTypeAttribute.LLM, - extractInput: extractOllamaChatInput, - extractOutput: (result, event) => - extractOllamaChatOutput( - result, - countOllamaToolCalls(event?.arguments?.[0]?.messages), - ), - extractMetadata: extractOllamaResponseMetadata, - extractMetrics: extractOllamaMetrics, - aggregateChunks: aggregateOllamaChatChunks, - }), - traceStreamingChannel(ollamaChannels.generate, { - name: "ollama.generate", - type: SpanTypeAttribute.LLM, - extractInput: extractOllamaGenerateInput, - extractOutput: extractOllamaGenerateOutput, - extractMetadata: extractOllamaResponseMetadata, - extractMetrics: extractOllamaMetrics, - aggregateChunks: aggregateOllamaGenerateChunks, - }), - traceAsyncChannel(ollamaChannels.embed, { - name: "ollama.embed", - type: SpanTypeAttribute.LLM, - extractInput: extractOllamaEmbedInput, - extractOutput: extractOllamaEmbedOutput, - extractMetadata: extractOllamaResponseMetadata, - extractMetrics: extractOllamaMetrics, - }), + ollamaChannels.chat.intercept((target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OLLAMA, + name: "ollama.chat", + type: SpanTypeAttribute.LLM, + extractInput: extractOllamaChatInput, + extractOutput: (result, event) => + extractOllamaChatOutput( + result, + countOllamaToolCalls(event?.arguments?.[0]?.messages), + ), + extractMetadata: extractOllamaResponseMetadata, + extractMetrics: extractOllamaMetrics, + aggregateChunks: aggregateOllamaChatChunks, + }, + ), + ), + ollamaChannels.generate.intercept((target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OLLAMA, + name: "ollama.generate", + type: SpanTypeAttribute.LLM, + extractInput: extractOllamaGenerateInput, + extractOutput: extractOllamaGenerateOutput, + extractMetadata: extractOllamaResponseMetadata, + extractMetrics: extractOllamaMetrics, + aggregateChunks: aggregateOllamaGenerateChunks, + }, + ), + ), + ollamaChannels.embed.intercept((target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OLLAMA, + name: "ollama.embed", + type: SpanTypeAttribute.LLM, + extractInput: extractOllamaEmbedInput, + extractOutput: extractOllamaEmbedOutput, + extractMetadata: extractOllamaResponseMetadata, + extractMetrics: extractOllamaMetrics, + }, + ), + ), ); } diff --git a/js/src/instrumentation/plugins/openai-agents-channels.ts b/js/src/instrumentation/plugins/openai-agents-channels.ts index fb4760869..302e86f30 100644 --- a/js/src/instrumentation/plugins/openai-agents-channels.ts +++ b/js/src/instrumentation/plugins/openai-agents-channels.ts @@ -1,29 +1,24 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { OpenAIAgentsSpan, OpenAIAgentsTrace, } from "../../vendor-sdk-types/openai-agents"; -export const openAIAgentsCoreChannels = defineChannels( +export const openAIAgentsCoreChannels = defineInterceptor( "@openai/agents-core", { - onTraceStart: channel<[OpenAIAgentsTrace], void>({ + onTraceStart: channel<[OpenAIAgentsTrace], PromiseLike>({ channelName: "tracing.processor.onTraceStart", - kind: "async", }), - onTraceEnd: channel<[OpenAIAgentsTrace], void>({ + onTraceEnd: channel<[OpenAIAgentsTrace], PromiseLike>({ channelName: "tracing.processor.onTraceEnd", - kind: "async", }), - onSpanStart: channel<[OpenAIAgentsSpan], void>({ + onSpanStart: channel<[OpenAIAgentsSpan], PromiseLike>({ channelName: "tracing.processor.onSpanStart", - kind: "async", }), - onSpanEnd: channel<[OpenAIAgentsSpan], void>({ + onSpanEnd: channel<[OpenAIAgentsSpan], PromiseLike>({ channelName: "tracing.processor.onSpanEnd", - kind: "async", }), }, - { instrumentationName: INSTRUMENTATION_NAMES.OPENAI_AGENTS }, ); diff --git a/js/src/instrumentation/plugins/openai-agents-plugin.test.ts b/js/src/instrumentation/plugins/openai-agents-plugin.test.ts index 9f2a604bd..4cf44e613 100644 --- a/js/src/instrumentation/plugins/openai-agents-plugin.test.ts +++ b/js/src/instrumentation/plugins/openai-agents-plugin.test.ts @@ -1,6 +1,6 @@ import { afterEach, beforeAll, beforeEach, describe, expect, it } from "vitest"; -import { configureNode } from "../../node/config"; import { _exportsForTestingOnly, initLogger } from "../../logger"; +import { configureNode } from "../../node/config"; import { openAIAgentsCoreChannels } from "./openai-agents-channels"; import { OpenAIAgentsPlugin } from "./openai-agents-plugin"; @@ -66,21 +66,29 @@ describe("OpenAIAgentsPlugin", () => { }, }; - await openAIAgentsCoreChannels.onTraceStart.tracePromise( + await openAIAgentsCoreChannels.onTraceStart.invoke( async () => undefined, - { arguments: [trace as any] }, + undefined, + [trace as any], + {}, ); - await openAIAgentsCoreChannels.onSpanStart.tracePromise( + await openAIAgentsCoreChannels.onSpanStart.invoke( async () => undefined, - { arguments: [span as any] }, + undefined, + [span as any], + {}, ); - await openAIAgentsCoreChannels.onSpanEnd.tracePromise( + await openAIAgentsCoreChannels.onSpanEnd.invoke( async () => undefined, - { arguments: [span as any] }, + undefined, + [span as any], + {}, ); - await openAIAgentsCoreChannels.onTraceEnd.tracePromise( + await openAIAgentsCoreChannels.onTraceEnd.invoke( async () => undefined, - { arguments: [trace as any] }, + undefined, + [trace as any], + {}, ); const spans = (await backgroundLogger.drain()) as any[]; diff --git a/js/src/instrumentation/plugins/openai-agents-plugin.ts b/js/src/instrumentation/plugins/openai-agents-plugin.ts index 4cf8af1a0..387c1ceac 100644 --- a/js/src/instrumentation/plugins/openai-agents-plugin.ts +++ b/js/src/instrumentation/plugins/openai-agents-plugin.ts @@ -1,12 +1,13 @@ -import { BasePlugin } from "../core"; -import { unsubscribeAll } from "../core/channel-tracing"; import { isObject } from "../../../util/index"; -import { openAIAgentsCoreChannels } from "./openai-agents-channels"; -import { OpenAIAgentsTraceProcessor } from "./openai-agents-trace-processor"; import type { OpenAIAgentsSpan, OpenAIAgentsTrace, } from "../../vendor-sdk-types/openai-agents"; +import { BasePlugin } from "../core"; +import { unsubscribeAll } from "../core/channel-tracing"; +import { runInstrumentation } from "../core/observe-result"; +import { openAIAgentsCoreChannels } from "./openai-agents-channels"; +import { OpenAIAgentsTraceProcessor } from "./openai-agents-trace-processor"; function firstArgument(args: unknown): unknown { if (Array.isArray(args)) { @@ -54,61 +55,72 @@ export class OpenAIAgentsPlugin extends BasePlugin { } private subscribeToTraceLifecycle(): void { - const traceStartChannel = - openAIAgentsCoreChannels.onTraceStart.tracingChannel(); - const traceStartHandlers = { - start: (event: { arguments: unknown }) => { - const trace = firstArgument(event.arguments); - if (isOpenAIAgentsTrace(trace)) { - void this.processor.onTraceStart(trace); - } + const traceStartChannel = openAIAgentsCoreChannels.onTraceStart; + + const removetraceStartHandlers = traceStartChannel.intercept( + (target, receiver, args) => { + runInstrumentation(() => + ((event: { arguments: unknown }) => { + const trace = firstArgument(event.arguments); + if (isOpenAIAgentsTrace(trace)) { + void this.processor.onTraceStart(trace); + } + })({ arguments: args }), + ); + return Reflect.apply(target, receiver, args); }, - }; - traceStartChannel.subscribe(traceStartHandlers); - this.unsubscribers.push(() => - traceStartChannel.unsubscribe(traceStartHandlers), ); + this.unsubscribers.push(removetraceStartHandlers); + + const traceEndChannel = openAIAgentsCoreChannels.onTraceEnd; - const traceEndChannel = - openAIAgentsCoreChannels.onTraceEnd.tracingChannel(); - const traceEndHandlers = { - start: (event: { arguments: unknown }) => { - const trace = firstArgument(event.arguments); - if (isOpenAIAgentsTrace(trace)) { - void this.processor.onTraceEnd(trace); - } + const removetraceEndHandlers = traceEndChannel.intercept( + (target, receiver, args) => { + runInstrumentation(() => + ((event: { arguments: unknown }) => { + const trace = firstArgument(event.arguments); + if (isOpenAIAgentsTrace(trace)) { + void this.processor.onTraceEnd(trace); + } + })({ arguments: args }), + ); + return Reflect.apply(target, receiver, args); }, - }; - traceEndChannel.subscribe(traceEndHandlers); - this.unsubscribers.push(() => - traceEndChannel.unsubscribe(traceEndHandlers), ); + this.unsubscribers.push(removetraceEndHandlers); - const spanStartChannel = - openAIAgentsCoreChannels.onSpanStart.tracingChannel(); - const spanStartHandlers = { - start: (event: { arguments: unknown }) => { - const span = firstArgument(event.arguments); - if (isOpenAIAgentsSpan(span)) { - void this.processor.onSpanStart(span); - } + const spanStartChannel = openAIAgentsCoreChannels.onSpanStart; + + const removespanStartHandlers = spanStartChannel.intercept( + (target, receiver, args) => { + runInstrumentation(() => + ((event: { arguments: unknown }) => { + const span = firstArgument(event.arguments); + if (isOpenAIAgentsSpan(span)) { + void this.processor.onSpanStart(span); + } + })({ arguments: args }), + ); + return Reflect.apply(target, receiver, args); }, - }; - spanStartChannel.subscribe(spanStartHandlers); - this.unsubscribers.push(() => - spanStartChannel.unsubscribe(spanStartHandlers), ); + this.unsubscribers.push(removespanStartHandlers); + + const spanEndChannel = openAIAgentsCoreChannels.onSpanEnd; - const spanEndChannel = openAIAgentsCoreChannels.onSpanEnd.tracingChannel(); - const spanEndHandlers = { - start: (event: { arguments: unknown }) => { - const span = firstArgument(event.arguments); - if (isOpenAIAgentsSpan(span)) { - void this.processor.onSpanEnd(span); - } + const removespanEndHandlers = spanEndChannel.intercept( + (target, receiver, args) => { + runInstrumentation(() => + ((event: { arguments: unknown }) => { + const span = firstArgument(event.arguments); + if (isOpenAIAgentsSpan(span)) { + void this.processor.onSpanEnd(span); + } + })({ arguments: args }), + ); + return Reflect.apply(target, receiver, args); }, - }; - spanEndChannel.subscribe(spanEndHandlers); - this.unsubscribers.push(() => spanEndChannel.unsubscribe(spanEndHandlers)); + ); + this.unsubscribers.push(removespanEndHandlers); } } diff --git a/js/src/instrumentation/plugins/openai-channels.ts b/js/src/instrumentation/plugins/openai-channels.ts index 6f6bac9a0..71a7d8496 100644 --- a/js/src/instrumentation/plugins/openai-channels.ts +++ b/js/src/instrumentation/plugins/openai-channels.ts @@ -1,12 +1,21 @@ +import type { CompiledPrompt } from "../../logger"; import type { OpenAIMediaParams, OpenAIMediaResponse, } from "../../vendor-sdk-types/openai-media"; -import type { CompiledPrompt } from "../../logger"; -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; -import type { StartOf } from "../core/channel-definitions"; -import type { ChannelSpanInfo, SpanInfoCarrier } from "../core/types"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + +import type { + OpenAIAgentsTraceStartChannelArgs, + OpenAIAgentsTraceState, +} from "../../openai-agents-api-types"; +import type { + CompleteOpenAIBatchTraceArgs, + OpenAIBatchLike, + OpenAIBatchesRetrieveTraceArgs, + OpenAIFileLike, + OpenAIFilesCreateTraceArgs, +} from "../../openai-batch-types"; import type { OpenAIChatCompletion, OpenAIChatCompletionChunk, @@ -21,209 +30,181 @@ import type { OpenAIResponseCreateParams, OpenAIResponseStreamEvent, } from "../../vendor-sdk-types/openai"; -import type { - CompleteOpenAIBatchTraceArgs, - OpenAIBatchLike, - OpenAIBatchesRetrieveTraceArgs, - OpenAIFileLike, - OpenAIFilesCreateTraceArgs, -} from "../../openai-batch-types"; -import type { - OpenAIAgentsTraceStartChannelArgs, - OpenAIAgentsTraceState, -} from "../../openai-agents-api-types"; +import type { ChannelSpanInfo, SpanInfoCarrier } from "../core/types"; type OpenAIChatSpanInfo = NonNullable["span_info"]>; type OpenAIChannelExtras = SpanInfoCarrier & { response?: Response; + responseInfo?: { response?: Response }; }; type OpenAIChatChannelExtras = OpenAIChannelExtras; type OpenAIResponsesChannelExtras = OpenAIChannelExtras; -export const openAIChannels = defineChannels( - "openai", - { - imagesGenerate: channel< - [OpenAIMediaParams, unknown?], - OpenAIMediaResponse, - OpenAIChannelExtras - >({ channelName: "images.generate", kind: "async" }), - imagesEdit: channel< - [OpenAIMediaParams, unknown?], - OpenAIMediaResponse, - OpenAIChannelExtras - >({ channelName: "images.edit", kind: "async" }), - imagesCreateVariation: channel< - [OpenAIMediaParams, unknown?], - OpenAIMediaResponse, - OpenAIChannelExtras - >({ channelName: "images.createVariation", kind: "async" }), - audioSpeechCreate: channel< - [OpenAIMediaParams, unknown?], - OpenAIMediaResponse, - OpenAIChannelExtras - >({ channelName: "audio.speech.create", kind: "async" }), - audioTranscriptionsCreate: channel< - [OpenAIMediaParams, unknown?], - OpenAIMediaResponse, - OpenAIChannelExtras - >({ channelName: "audio.transcriptions.create", kind: "async" }), - audioTranslationsCreate: channel< - [OpenAIMediaParams, unknown?], - OpenAIMediaResponse, - OpenAIChannelExtras - >({ channelName: "audio.translations.create", kind: "async" }), - - agentsTraceStart: channel< - [OpenAIAgentsTraceStartChannelArgs], - OpenAIAgentsTraceState | null, - OpenAIChannelExtras - >({ - channelName: "agents.trace.start", - kind: "async", - }), - - agentsTraceCapture: channel< - [{ event: unknown; state: OpenAIAgentsTraceState | null }], - OpenAIAgentsTraceState | null, - OpenAIChannelExtras - >({ - channelName: "agents.trace.capture", - kind: "async", - }), - - agentsTraceFail: channel< - [{ error: unknown; state: OpenAIAgentsTraceState | null }], - OpenAIAgentsTraceState | null, - OpenAIChannelExtras - >({ - channelName: "agents.trace.fail", - kind: "async", - }), - - filesCreateTraced: channel< - [OpenAIFilesCreateTraceArgs], - OpenAIFileLike, - OpenAIChannelExtras - >({ - channelName: "files.create-traced", - kind: "async", - }), - - batchesRetrieveTraced: channel< - [OpenAIBatchesRetrieveTraceArgs], - OpenAIBatchLike, - OpenAIChannelExtras - >({ - channelName: "batches.retrieve-traced", - kind: "async", - }), - - batchesCompleteTrace: channel< - [CompleteOpenAIBatchTraceArgs], - void, - OpenAIChannelExtras - >({ - channelName: "batches.complete-trace", - kind: "async", - }), - - chatCompletionsCreate: channel< - [OpenAIChatCreateParams], - OpenAIChatCompletion | OpenAIChatStream, - OpenAIChatChannelExtras, - OpenAIChatCompletionChunk - >({ - channelName: "chat.completions.create", - kind: "async", - }), - - embeddingsCreate: channel< - [OpenAIEmbeddingCreateParams], - OpenAIEmbeddingResponse, - OpenAIChatChannelExtras - >({ - channelName: "embeddings.create", - kind: "async", - }), - - betaChatCompletionsParse: channel< - [OpenAIChatCreateParams], - OpenAIChatCompletion, - OpenAIChatChannelExtras, - OpenAIChatCompletionChunk - >({ - channelName: "beta.chat.completions.parse", - kind: "async", - }), - - betaChatCompletionsStream: channel< - [OpenAIChatCreateParams], - unknown, - OpenAIChatChannelExtras - >({ - channelName: "beta.chat.completions.stream", - kind: "sync-stream", - }), - - moderationsCreate: channel< - [OpenAIModerationCreateParams], - OpenAIModerationResponse, - OpenAIChatChannelExtras - >({ - channelName: "moderations.create", - kind: "async", - }), - - responsesCreate: channel< - [OpenAIResponseCreateParams], - OpenAIResponse | AsyncIterable, - OpenAIResponsesChannelExtras, - OpenAIResponseStreamEvent - >({ - channelName: "responses.create", - kind: "async", - }), - - responsesStream: channel< - [OpenAIResponseCreateParams], - unknown, - OpenAIResponsesChannelExtras, - OpenAIResponseStreamEvent - >({ - channelName: "responses.stream", - kind: "sync-stream", - }), - - responsesParse: channel< - [OpenAIResponseCreateParams], - OpenAIResponse, - OpenAIResponsesChannelExtras, - OpenAIResponseStreamEvent - >({ - channelName: "responses.parse", - kind: "async", - }), - - responsesCompact: channel< - [OpenAIResponseCompactParams], - OpenAIResponse, - OpenAIResponsesChannelExtras - >({ - channelName: "responses.compact", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.OPENAI }, -); +export const openAIChannels = defineInterceptor("openai", { + imagesGenerate: channel< + [OpenAIMediaParams, unknown?], + PromiseLike, + OpenAIChannelExtras + >({ channelName: "images.generate" }), + imagesEdit: channel< + [OpenAIMediaParams, unknown?], + PromiseLike, + OpenAIChannelExtras + >({ channelName: "images.edit" }), + imagesCreateVariation: channel< + [OpenAIMediaParams, unknown?], + PromiseLike, + OpenAIChannelExtras + >({ channelName: "images.createVariation" }), + audioSpeechCreate: channel< + [OpenAIMediaParams, unknown?], + PromiseLike, + OpenAIChannelExtras + >({ channelName: "audio.speech.create" }), + audioTranscriptionsCreate: channel< + [OpenAIMediaParams, unknown?], + PromiseLike, + OpenAIChannelExtras + >({ channelName: "audio.transcriptions.create" }), + audioTranslationsCreate: channel< + [OpenAIMediaParams, unknown?], + PromiseLike, + OpenAIChannelExtras + >({ channelName: "audio.translations.create" }), + + agentsTraceStart: channel< + [OpenAIAgentsTraceStartChannelArgs], + PromiseLike, + OpenAIChannelExtras + >({ + channelName: "agents.trace.start", + }), + + agentsTraceCapture: channel< + [{ event: unknown; state: OpenAIAgentsTraceState | null }], + PromiseLike, + OpenAIChannelExtras + >({ + channelName: "agents.trace.capture", + }), + + agentsTraceFail: channel< + [{ error: unknown; state: OpenAIAgentsTraceState | null }], + PromiseLike, + OpenAIChannelExtras + >({ + channelName: "agents.trace.fail", + }), + + filesCreateTraced: channel< + [OpenAIFilesCreateTraceArgs], + PromiseLike, + OpenAIChannelExtras + >({ + channelName: "files.create-traced", + }), + + batchesRetrieveTraced: channel< + [OpenAIBatchesRetrieveTraceArgs], + PromiseLike, + OpenAIChannelExtras + >({ + channelName: "batches.retrieve-traced", + }), + + batchesCompleteTrace: channel< + [CompleteOpenAIBatchTraceArgs], + PromiseLike, + OpenAIChannelExtras + >({ + channelName: "batches.complete-trace", + }), + + chatCompletionsCreate: channel< + [OpenAIChatCreateParams], + PromiseLike, + OpenAIChatChannelExtras, + OpenAIChatCompletionChunk + >({ + channelName: "chat.completions.create", + }), + + embeddingsCreate: channel< + [OpenAIEmbeddingCreateParams], + PromiseLike, + OpenAIChatChannelExtras + >({ + channelName: "embeddings.create", + }), + + betaChatCompletionsParse: channel< + [OpenAIChatCreateParams], + PromiseLike, + OpenAIChatChannelExtras, + OpenAIChatCompletionChunk + >({ + channelName: "beta.chat.completions.parse", + }), + + betaChatCompletionsStream: channel< + [OpenAIChatCreateParams], + unknown, + OpenAIChatChannelExtras + >({ + channelName: "beta.chat.completions.stream", + }), + + moderationsCreate: channel< + [OpenAIModerationCreateParams], + PromiseLike, + OpenAIChatChannelExtras + >({ + channelName: "moderations.create", + }), + + responsesCreate: channel< + [OpenAIResponseCreateParams], + PromiseLike>, + OpenAIResponsesChannelExtras, + OpenAIResponseStreamEvent + >({ + channelName: "responses.create", + }), + + responsesStream: channel< + [OpenAIResponseCreateParams], + unknown, + OpenAIResponsesChannelExtras, + OpenAIResponseStreamEvent + >({ + channelName: "responses.stream", + }), + + responsesParse: channel< + [OpenAIResponseCreateParams], + PromiseLike, + OpenAIResponsesChannelExtras, + OpenAIResponseStreamEvent + >({ + channelName: "responses.parse", + }), + + responsesCompact: channel< + [OpenAIResponseCompactParams], + PromiseLike, + OpenAIResponsesChannelExtras + >({ + channelName: "responses.compact", + }), +}); export type OpenAIChannel = (typeof openAIChannels)[keyof typeof openAIChannels]; -export type OpenAIAsyncChannel = Extract; - -export type OpenAIStartContext = - StartOf; +export type OpenAIAsyncChannel = Extract< + OpenAIChannel, + { __result?: PromiseLike } +>; diff --git a/js/src/instrumentation/plugins/openai-codex-channels.ts b/js/src/instrumentation/plugins/openai-codex-channels.ts index bd622fe5c..e0ba37d6f 100644 --- a/js/src/instrumentation/plugins/openai-codex-channels.ts +++ b/js/src/instrumentation/plugins/openai-codex-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { OpenAICodexInput, OpenAICodexStreamedTurn, @@ -9,26 +9,20 @@ import type { OpenAICodexTurnOptions, } from "../../vendor-sdk-types/openai-codex"; -export const openAICodexChannels = defineChannels( - "@openai/codex-sdk", - { - run: channel< - [OpenAICodexInput, OpenAICodexTurnOptions | undefined], - OpenAICodexTurn, - { operation?: "run"; thread?: OpenAICodexThread } - >({ - channelName: "Thread.run", - kind: "async", - }), - runStreamed: channel< - [OpenAICodexInput, OpenAICodexTurnOptions | undefined], - OpenAICodexStreamedTurn, - { operation?: "runStreamed"; thread?: OpenAICodexThread }, - OpenAICodexThreadEvent - >({ - channelName: "Thread.runStreamed", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.OPENAI_CODEX }, -); +export const openAICodexChannels = defineInterceptor("@openai/codex-sdk", { + run: channel< + [OpenAICodexInput, OpenAICodexTurnOptions | undefined], + PromiseLike, + { operation?: "run"; thread?: OpenAICodexThread } + >({ + channelName: "Thread.run", + }), + runStreamed: channel< + [OpenAICodexInput, OpenAICodexTurnOptions | undefined], + PromiseLike, + { operation?: "runStreamed"; thread?: OpenAICodexThread }, + OpenAICodexThreadEvent + >({ + channelName: "Thread.runStreamed", + }), +}); diff --git a/js/src/instrumentation/plugins/openai-codex-plugin.test.ts b/js/src/instrumentation/plugins/openai-codex-plugin.test.ts index 932abe0ba..e25e99f62 100644 --- a/js/src/instrumentation/plugins/openai-codex-plugin.test.ts +++ b/js/src/instrumentation/plugins/openai-codex-plugin.test.ts @@ -1,23 +1,30 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +import { invocationController } from "../test-utils/invocation"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const { mockStartSpan } = vi.hoisted(() => ({ mockStartSpan: vi.fn(), })); vi.mock("../../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(), - }, + default: {}, })); vi.mock("../../logger", () => ({ startSpan: (...args: unknown[]) => mockStartSpan(...args), })); -import iso from "../../isomorph"; import { OpenAICodexPlugin } from "./openai-codex-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; describe("OpenAICodexPlugin", () => { let handlersByName: Map; @@ -31,9 +38,11 @@ describe("OpenAICodexPlugin", () => { beforeEach(() => { handlersByName = new Map(); spans = []; - mockNewTracingChannel.mockImplementation((name: string) => ({ - subscribe: vi.fn((handlers) => handlersByName.set(name, handlers)), - unsubscribe: vi.fn(), + mockNewInvocationHook.mockImplementation((name: string) => ({ + intercept: vi.fn((interceptor) => { + handlersByName.set(name, invocationController(interceptor)); + return vi.fn(); + }), })); mockStartSpan.mockImplementation((args: any) => { const span = { @@ -79,8 +88,8 @@ describe("OpenAICodexPlugin", () => { thread: { id: "thread-1" }, }; - runHandlers.start(event); - await runHandlers.asyncEnd(event); + runHandlers.begin(event); + await runHandlers.resolve(event); const rootSpan = spans.find((span) => span.name === "OpenAI Codex"); expect(rootSpan?.log).toHaveBeenCalledWith( diff --git a/js/src/instrumentation/plugins/openai-codex-plugin.ts b/js/src/instrumentation/plugins/openai-codex-plugin.ts index 4005ecb86..fd4109d8e 100644 --- a/js/src/instrumentation/plugins/openai-codex-plugin.ts +++ b/js/src/instrumentation/plugins/openai-codex-plugin.ts @@ -1,16 +1,15 @@ import { BasePlugin, toLoggedError } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; -import type { IsoChannelHandlers } from "../../isomorph"; +import { observeResult, runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute } from "../../../util/index"; import { debugLogger } from "../../debug-logger"; -import { startSpan as startBaseSpan } from "../../logger"; import type { Span, StartSpanArgs } from "../../logger"; +import { startSpan as startBaseSpan } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; import { getCurrentUnixTimestamp } from "../../util"; -import { SpanTypeAttribute } from "../../../util/index"; -import { openAICodexChannels } from "./openai-codex-channels"; import type { OpenAICodexCommandExecutionItem, OpenAICodexFileChangeItem, @@ -26,6 +25,8 @@ import type { OpenAICodexUsage, OpenAICodexWebSearchItem, } from "../../vendor-sdk-types/openai-codex"; +import type { ChannelMessage } from "../core/tracing-types"; +import { openAICodexChannels } from "./openai-codex-channels"; type CodexRunState = { activeLlmSpan?: CodexLlmSpanState; @@ -69,71 +70,125 @@ export class OpenAICodexPlugin extends BasePlugin { } private subscribeToRun(): void { - const channel = openAICodexChannels.run.tracingChannel(); + const channel = openAICodexChannels.run; const states = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - states.set(event, startCodexRun(event, "Thread.run")); - }, - asyncEnd: async (event) => { - const state = states.get(event); - if (!state) { - return; + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + states.set(event, startCodexRun(event, "Thread.run")); + }; + const resolved = async ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + states.delete(event); + await finalizeCompletedRun(state, event.result); + }; + const failed = async ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + states.delete(event); + await finalizeCodexRun(state, { error: event.error }); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } - states.delete(event); - await finalizeCompletedRun(state, event.result); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + return resolved(event); + }, + (error) => { + Object.assign(event, { error }); + return failed(event); + }, + ); }, - error: async (event) => { - const state = states.get(event); - if (!state) { - return; - } - states.delete(event); - await finalizeCodexRun(state, { error: event.error }); - }, - }; - - channel.subscribe(handlers); - this.unsubscribers.push(() => { - channel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } private subscribeToRunStreamed(): void { - const channel = openAICodexChannels.runStreamed.tracingChannel(); + const channel = openAICodexChannels.runStreamed; const states = new WeakMap(); - const handlers: IsoChannelHandlers< - ChannelMessage - > = { - start: (event) => { - states.set(event, startCodexRun(event, "Thread.runStreamed")); - }, - asyncEnd: async (event) => { - const state = states.get(event); - if (!state) { - return; + const removeHandlers = channel.intercept( + (target, receiver, args, additional) => { + const event = { + ...additional, + arguments: args, + self: receiver, + } as ChannelMessage; + const prepare = ( + event: ChannelMessage, + ) => { + states.set(event, startCodexRun(event, "Thread.runStreamed")); + }; + const resolved = async ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + states.delete(event); + await patchStreamedTurn(event.result, state); + }; + const failed = async ( + event: ChannelMessage, + ) => { + const state = states.get(event); + if (!state) { + return; + } + states.delete(event); + await finalizeCodexRun(state, { error: event.error }); + }; + runInstrumentation(() => prepare(event)); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + Object.assign(event, { error }); + runInstrumentation(() => failed(event)); + throw error; } - states.delete(event); - await patchStreamedTurn(event.result, state); + return observeResult( + result, + (value) => { + Object.assign(event, { result: value }); + return resolved(event); + }, + (error) => { + Object.assign(event, { error }); + return failed(event); + }, + ); }, - error: async (event) => { - const state = states.get(event); - if (!state) { - return; - } - states.delete(event); - await finalizeCodexRun(state, { error: event.error }); - }, - }; - - channel.subscribe(handlers); - this.unsubscribers.push(() => { - channel.unsubscribe(handlers); - }); + ); + this.unsubscribers.push(removeHandlers); } } diff --git a/js/src/instrumentation/plugins/openai-media.ts b/js/src/instrumentation/plugins/openai-media.ts index 3e20aec1f..4da80fc59 100644 --- a/js/src/instrumentation/plugins/openai-media.ts +++ b/js/src/instrumentation/plugins/openai-media.ts @@ -1,3 +1,7 @@ +import type { + ArgsOf, + InvocationAdditionalOf, +} from "../core/channel-definitions"; import { Attachment, startSpan, withCurrent, type Span } from "../../logger"; import { debugLogger } from "../../debug-logger"; import { getCurrentUnixTimestamp } from "../../util"; @@ -654,193 +658,196 @@ function observeSpeech( } type MediaChannel = typeof openAIChannels.imagesGenerate; -export function interceptOpenAIMedia( - channel: MediaChannel, +export function traceOpenAIMedia( + call: () => PromiseLike, + context: { + arguments: ArgsOf; + self: unknown; + additional: InvocationAdditionalOf; + }, + channelName: string, operation: string, -): () => void { - return channel.intercept((target, thisArg, args, additional) => { - if (isAutoInstrumentationSuppressed()) return target.apply(thisArg, args); - const params = args[0]; - const captureAttachments = isAutoCaptureAttachmentsEnabled(); - let span: Span; - try { - const { name, spanAttributes, spanInfoMetadata } = buildStartSpanArgs( - { name: `openai.${channel.channelName}`, type: "llm" }, - { arguments: args, span_info: additional.span_info }, - ); - span = startSpan( - withSpanInstrumentationName( - { - name, - spanAttributes, - event: { - metadata: mergeInputMetadata( - { - // Provider defaults can change independently of the SDK version. - model: - params.model ?? - (operation === "variation" ? "dall-e-2" : undefined), - provider: "openai", - }, - spanInfoMetadata, - ), - }, +): PromiseLike { + const args = context.arguments; + const additional = context.additional; + if (isAutoInstrumentationSuppressed()) return call(); + const params = args[0]; + const captureAttachments = isAutoCaptureAttachmentsEnabled(); + let span: Span; + try { + const { name, spanAttributes, spanInfoMetadata } = buildStartSpanArgs( + { name: `openai.${channelName}`, type: "llm" }, + { arguments: args, span_info: additional.span_info }, + ); + span = startSpan( + withSpanInstrumentationName( + { + name, + spanAttributes, + event: { + metadata: mergeInputMetadata( + { + // Provider defaults can change independently of the SDK version. + model: + params.model ?? + (operation === "variation" ? "dall-e-2" : undefined), + provider: "openai", + }, + spanInfoMetadata, + ), }, - INSTRUMENTATION_NAMES.OPENAI, - ), - ); - } catch (error) { - debugLogger.debug("OpenAI media span failed", error); - return target.apply(thisArg, args); + }, + INSTRUMENTATION_NAMES.OPENAI, + ), + ); + } catch (error) { + debugLogger.debug("OpenAI media span failed", error); + return call(); + } + void mediaInput(params, operation, captureAttachments) + .then((input) => span.log({ input })) + .catch((error) => debugLogger.debug("OpenAI media input failed", error)); + const start = getCurrentUnixTimestamp(); + let seen = false; + let ended = false; + const first = () => { + if (!seen) { + seen = true; + span.log({ + metrics: { time_to_first_token: getCurrentUnixTimestamp() - start }, + }); } - void mediaInput(params, operation, captureAttachments) - .then((input) => span.log({ input })) - .catch((error) => debugLogger.debug("OpenAI media input failed", error)); - const start = getCurrentUnixTimestamp(); - let seen = false; - let ended = false; - const first = () => { - if (!seen) { - seen = true; + }; + const finish = (result?: OpenAIMediaResult | string, error?: unknown) => { + if (ended) return; + ended = true; + try { + if (result !== undefined) span.log({ - metrics: { time_to_first_token: getCurrentUnixTimestamp() - start }, + output: mediaOutput(result, params, captureAttachments), + ...(typeof result === "object" + ? { + metrics: mediaUsage(result.usage), + ...(result.model ? { metadata: { model: result.model } } : {}), + } + : {}), }); - } - }; - const finish = (result?: OpenAIMediaResult | string, error?: unknown) => { - if (ended) return; - ended = true; - try { - if (result !== undefined) - span.log({ - output: mediaOutput(result, params, captureAttachments), - ...(typeof result === "object" - ? { - metrics: mediaUsage(result.usage), - ...(result.model - ? { metadata: { model: result.model } } - : {}), - } - : {}), - }); - if (error !== undefined) span.log({ error }); - } catch (error) { - debugLogger.debug("OpenAI media output failed", error); - } finally { - try { - span.end(); - } catch (error) { - debugLogger.debug("OpenAI media finalization failed", error); - } - } - }; - let observed = false; - const onValue = (value: OpenAIMediaResponse) => { - if (observed) return; - observed = true; + if (error !== undefined) span.log({ error }); + } catch (error) { + debugLogger.debug("OpenAI media output failed", error); + } finally { try { - if ( - value instanceof Response || - (isObject(value) && - typeof Reflect.get(value, "arrayBuffer") === "function" && - Reflect.get(value, "headers")) - ) { - span.log({ output: { content: [] } }); - if (captureAttachments) - observeSpeech(value as Response, params, span, first); - finish(); - } else if (isAsyncIterable(value)) { - const accumulated: OpenAIMediaResult = {}; - const audio: Blob[] = []; - patchStreamIfNeeded(value, { - onChunk: (event) => { - if (event.b64_json || event.audio || event.delta || event.text) - first(); - if (event.usage) accumulated.usage = event.usage; - if (event.model) accumulated.model = event.model; - if (event.type.endsWith(".completed") && event.b64_json) { - accumulated.data = [ - ...(accumulated.data ?? []), - { - b64_json: captureAttachments ? event.b64_json : "", - }, - ]; - accumulated.output_format = event.output_format; - } - if (event.type === "transcript.text.delta") - accumulated.text = - (accumulated.text ?? "") + (event.delta ?? ""); - if (event.type === "transcript.text.done") - accumulated.text = event.text ?? accumulated.text; - if (event.type === "transcript.text.segment") - accumulated.segments = [...(accumulated.segments ?? []), event]; - if (captureAttachments && event.audio) { - const blob = convertDataToBlob( - event.audio, - AUDIO_TYPES.get(params.response_format ?? "mp3") ?? - "application/octet-stream", - ); - if (blob) audio.push(blob); - } - }, - onComplete: () => { - finish(accumulated); - if (captureAttachments && audio.length) { - const contentType = - AUDIO_TYPES.get(params.response_format ?? "mp3") ?? - "application/octet-stream"; - const filename = `speech.${params.response_format ?? "mp3"}`; - span.log({ - output: { - content: [ - { - type: "file", - file: { - filename, - file_data: new Attachment({ - data: new Blob(audio, { type: contentType }), - contentType, - filename, - }), - }, - }, - ], - }, - }); - } - }, - onCancel: () => finish(accumulated), - onError: (error) => finish(accumulated, error), - }); - } else finish(value as OpenAIMediaResult | string); + span.end(); } catch (error) { - debugLogger.debug("OpenAI media observation failed", error); - finish(); + debugLogger.debug("OpenAI media finalization failed", error); } - }; - let result; - try { - result = withCurrent(span, () => - runWithAutoInstrumentationSuppressed(() => target.apply(thisArg, args)), - ); - } catch (error) { - finish(undefined, error); - throw error; } + }; + let observed = false; + const onValue = (value: OpenAIMediaResponse) => { + if (observed) return; + observed = true; try { - observeMediaPromise( - result, - onValue, - (error) => finish(undefined, error), - (response) => { - if (operation === "speech") onValue(response); - else finish(); - }, - ); + if ( + value instanceof Response || + (isObject(value) && + typeof Reflect.get(value, "arrayBuffer") === "function" && + Reflect.get(value, "headers")) + ) { + span.log({ output: { content: [] } }); + if (captureAttachments) + observeSpeech(value as Response, params, span, first); + finish(); + } else if (isAsyncIterable(value)) { + const accumulated: OpenAIMediaResult = {}; + const audio: Blob[] = []; + patchStreamIfNeeded(value, { + onChunk: (event) => { + if (event.b64_json || event.audio || event.delta || event.text) + first(); + if (event.usage) accumulated.usage = event.usage; + if (event.model) accumulated.model = event.model; + if (event.type.endsWith(".completed") && event.b64_json) { + accumulated.data = [ + ...(accumulated.data ?? []), + { + b64_json: captureAttachments ? event.b64_json : "", + }, + ]; + accumulated.output_format = event.output_format; + } + if (event.type === "transcript.text.delta") + accumulated.text = (accumulated.text ?? "") + (event.delta ?? ""); + if (event.type === "transcript.text.done") + accumulated.text = event.text ?? accumulated.text; + if (event.type === "transcript.text.segment") + accumulated.segments = [...(accumulated.segments ?? []), event]; + if (captureAttachments && event.audio) { + const blob = convertDataToBlob( + event.audio, + AUDIO_TYPES.get(params.response_format ?? "mp3") ?? + "application/octet-stream", + ); + if (blob) audio.push(blob); + } + }, + onComplete: () => { + finish(accumulated); + if (captureAttachments && audio.length) { + const contentType = + AUDIO_TYPES.get(params.response_format ?? "mp3") ?? + "application/octet-stream"; + const filename = `speech.${params.response_format ?? "mp3"}`; + span.log({ + output: { + content: [ + { + type: "file", + file: { + filename, + file_data: new Attachment({ + data: new Blob(audio, { type: contentType }), + contentType, + filename, + }), + }, + }, + ], + }, + }); + } + }, + onCancel: () => finish(accumulated), + onError: (error) => finish(accumulated, error), + }); + } else finish(value as OpenAIMediaResult | string); } catch (error) { - debugLogger.debug("OpenAI media promise observation failed", error); + debugLogger.debug("OpenAI media observation failed", error); finish(); } - return result; - }); + }; + let result; + try { + result = withCurrent(span, () => + runWithAutoInstrumentationSuppressed(() => call()), + ); + } catch (error) { + finish(undefined, error); + throw error; + } + try { + observeMediaPromise( + result, + onValue, + (error) => finish(undefined, error), + (response) => { + if (operation === "speech") onValue(response); + else finish(); + }, + ); + } catch (error) { + debugLogger.debug("OpenAI media promise observation failed", error); + finish(); + } + return result; } diff --git a/js/src/instrumentation/plugins/openai-plugin.ts b/js/src/instrumentation/plugins/openai-plugin.ts index f36df9ed3..1b747076e 100644 --- a/js/src/instrumentation/plugins/openai-plugin.ts +++ b/js/src/instrumentation/plugins/openai-plugin.ts @@ -1,41 +1,42 @@ -import { interceptOpenAIMedia } from "./openai-media"; -import { BasePlugin } from "../core"; -import { - traceAsyncChannel, - traceStreamingChannel, - traceSyncStreamChannel, - unsubscribeAll, -} from "../core/channel-tracing"; import { SpanTypeAttribute, isObject } from "../../../util/index"; -import { getCurrentUnixTimestamp } from "../../util"; -import { openAIChannels } from "./openai-channels"; -import { - extractOpenAIChatInput, - extractOpenAIResponsesInput, - extractOpenAIResponsesMetadata, - processImagesInOutput, -} from "./openai-span-data"; -import { - interceptOpenAIBatchesRetrieveTraced, - interceptOpenAIBatchTraceComplete, - interceptOpenAIFilesCreateTraced, -} from "./openai-batch-instrumentation"; -import { - interceptOpenAIAgentsTraceCapture, - interceptOpenAIAgentsTraceFail, - interceptOpenAIAgentsTraceStart, -} from "./openai-agents-api-instrumentation"; import { BRAINTRUST_CACHED_STREAM_METRIC, getCachedMetricFromHeaders, parseMetricsFromUsage, } from "../../openai-utils"; +import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { getCurrentUnixTimestamp } from "../../util"; import type { OpenAIChatChoice, OpenAIChatCompletionChunk, OpenAIChatLogprobs, OpenAIResponseStreamEvent, } from "../../vendor-sdk-types/openai"; +import { BasePlugin } from "../core"; +import { + traceAsyncCall, + traceStreamingCall, + traceSyncStreamCall, + unsubscribeAll, +} from "../core/channel-tracing"; +import { + interceptOpenAIAgentsTraceCapture, + interceptOpenAIAgentsTraceFail, + interceptOpenAIAgentsTraceStart, +} from "./openai-agents-api-instrumentation"; +import { + interceptOpenAIBatchTraceComplete, + interceptOpenAIBatchesRetrieveTraced, + interceptOpenAIFilesCreateTraced, +} from "./openai-batch-instrumentation"; +import { openAIChannels } from "./openai-channels"; +import { traceOpenAIMedia } from "./openai-media"; +import { + extractOpenAIChatInput, + extractOpenAIResponsesInput, + extractOpenAIResponsesMetadata, + processImagesInOutput, +} from "./openai-span-data"; /** * Plugin for OpenAI SDK instrumentation. @@ -73,228 +74,420 @@ export class OpenAIPlugin extends BasePlugin { ); this.unsubscribers.push( - interceptOpenAIMedia(openAIChannels.imagesGenerate, "generate"), - interceptOpenAIMedia(openAIChannels.imagesEdit, "edit"), - interceptOpenAIMedia(openAIChannels.imagesCreateVariation, "variation"), - interceptOpenAIMedia(openAIChannels.audioSpeechCreate, "speech"), - interceptOpenAIMedia( - openAIChannels.audioTranscriptionsCreate, - "transcribe", + openAIChannels.imagesGenerate.intercept( + (target, receiver, args, additional) => + traceOpenAIMedia( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + openAIChannels.imagesGenerate.channelName, + "generate", + ), + ), + openAIChannels.imagesEdit.intercept( + (target, receiver, args, additional) => + traceOpenAIMedia( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + openAIChannels.imagesEdit.channelName, + "edit", + ), + ), + openAIChannels.imagesCreateVariation.intercept( + (target, receiver, args, additional) => + traceOpenAIMedia( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + openAIChannels.imagesCreateVariation.channelName, + "variation", + ), + ), + openAIChannels.audioSpeechCreate.intercept( + (target, receiver, args, additional) => + traceOpenAIMedia( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + openAIChannels.audioSpeechCreate.channelName, + "speech", + ), + ), + openAIChannels.audioTranscriptionsCreate.intercept( + (target, receiver, args, additional) => + traceOpenAIMedia( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + openAIChannels.audioTranscriptionsCreate.channelName, + "transcribe", + ), + ), + openAIChannels.audioTranslationsCreate.intercept( + (target, receiver, args, additional) => + traceOpenAIMedia( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + openAIChannels.audioTranslationsCreate.channelName, + "translate", + ), ), - interceptOpenAIMedia(openAIChannels.audioTranslationsCreate, "translate"), ); // Chat Completions - supports streaming this.unsubscribers.push( - traceStreamingChannel(openAIChannels.chatCompletionsCreate, { - name: "Chat Completion", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => extractOpenAIChatInput(params), - extractOutput: (result) => { - return result?.choices; - }, - extractMetrics: (result, startTime, endEvent) => { - const metrics = withCachedMetric( - parseMetricsFromUsage(result?.usage), - result, - endEvent, - ); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - aggregateChunks: aggregateChatCompletionChunks, - }), + openAIChannels.chatCompletionsCreate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "Chat Completion", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => extractOpenAIChatInput(params), + extractOutput: (result) => { + return result?.choices; + }, + extractMetrics: (result, startTime, endEvent) => { + const metrics = withCachedMetric( + parseMetricsFromUsage(result?.usage), + result, + endEvent, + ); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + aggregateChunks: aggregateChatCompletionChunks, + }, + ), + ), ); // Embeddings this.unsubscribers.push( - traceAsyncChannel(openAIChannels.embeddingsCreate, { - name: "Embedding", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => { - const { input, ...metadata } = params; - return { - input, - metadata: { ...metadata, provider: "openai" }, - }; - }, - extractOutput: (result) => { - const embedding = result?.data?.[0]?.embedding; - return Array.isArray(embedding) - ? { embedding_length: embedding.length } - : undefined; - }, - extractMetrics: (result, _startTime, endEvent) => { - return withCachedMetric( - parseMetricsFromUsage(result?.usage), - result, - endEvent, - ); - }, - }), + openAIChannels.embeddingsCreate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "Embedding", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => { + const { input, ...metadata } = params; + return { + input, + metadata: { ...metadata, provider: "openai" }, + }; + }, + extractOutput: (result) => { + const embedding = result?.data?.[0]?.embedding; + return Array.isArray(embedding) + ? { embedding_length: embedding.length } + : undefined; + }, + extractMetrics: (result, _startTime, endEvent) => { + return withCachedMetric( + parseMetricsFromUsage(result?.usage), + result, + endEvent, + ); + }, + }, + ), + ), ); // Beta Chat Completions Parse this.unsubscribers.push( - traceStreamingChannel(openAIChannels.betaChatCompletionsParse, { - name: "Chat Completion", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => extractOpenAIChatInput(params), - extractOutput: (result) => { - return result?.choices; - }, - extractMetrics: (result, startTime, endEvent) => { - const metrics = withCachedMetric( - parseMetricsFromUsage(result?.usage), - result, - endEvent, - ); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - aggregateChunks: aggregateChatCompletionChunks, - }), + openAIChannels.betaChatCompletionsParse.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "Chat Completion", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => extractOpenAIChatInput(params), + extractOutput: (result) => { + return result?.choices; + }, + extractMetrics: (result, startTime, endEvent) => { + const metrics = withCachedMetric( + parseMetricsFromUsage(result?.usage), + result, + endEvent, + ); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + aggregateChunks: aggregateChatCompletionChunks, + }, + ), + ), ); // Beta Chat Completions Stream (sync method returning event-based stream) this.unsubscribers.push( - traceSyncStreamChannel(openAIChannels.betaChatCompletionsStream, { - name: "Chat Completion", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => extractOpenAIChatInput(params), - }), + openAIChannels.betaChatCompletionsStream.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "Chat Completion", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => extractOpenAIChatInput(params), + }, + ), + ), ); // Moderations this.unsubscribers.push( - traceAsyncChannel(openAIChannels.moderationsCreate, { - name: "Moderation", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => { - const { input, ...metadata } = params; - return { - input, - metadata: { ...metadata, provider: "openai" }, - }; - }, - extractOutput: (result) => { - return result?.results; - }, - extractMetrics: (result, _startTime, endEvent) => { - return withCachedMetric( - parseMetricsFromUsage(result?.usage), - result, - endEvent, - ); - }, - }), + openAIChannels.moderationsCreate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "Moderation", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => { + const { input, ...metadata } = params; + return { + input, + metadata: { ...metadata, provider: "openai" }, + }; + }, + extractOutput: (result) => { + return result?.results; + }, + extractMetrics: (result, _startTime, endEvent) => { + return withCachedMetric( + parseMetricsFromUsage(result?.usage), + result, + endEvent, + ); + }, + }, + ), + ), ); // Responses API - create (supports streaming via stream=true param) this.unsubscribers.push( - traceStreamingChannel(openAIChannels.responsesCreate, { - name: "openai.responses.create", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => extractOpenAIResponsesInput(params), - extractOutput: (result) => { - return processImagesInOutput(result?.output); - }, - extractMetadata: (result) => extractOpenAIResponsesMetadata(result), - extractMetrics: (result, startTime, endEvent) => { - const metrics = withCachedMetric( - parseMetricsFromUsage(result?.usage), - result, - endEvent, - ); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - aggregateChunks: aggregateResponseStreamEvents, - }), + openAIChannels.responsesCreate.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "openai.responses.create", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => extractOpenAIResponsesInput(params), + extractOutput: (result) => { + return processImagesInOutput(result?.output); + }, + extractMetadata: (result) => + extractOpenAIResponsesMetadata(result), + extractMetrics: (result, startTime, endEvent) => { + const metrics = withCachedMetric( + parseMetricsFromUsage(result?.usage), + result, + endEvent, + ); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + aggregateChunks: aggregateResponseStreamEvents, + }, + ), + ), ); // Responses API - stream (sync method returning event-based stream) this.unsubscribers.push( - traceSyncStreamChannel(openAIChannels.responsesStream, { - name: "openai.responses.create", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => extractOpenAIResponsesInput(params), - extractFromEvent: (event) => { - if (event.type !== "response.completed" || !event.response) { - return {}; - } - - const response = event.response; - const data: Record = {}; - - if (response.output !== undefined) { - data.output = processImagesInOutput(response.output); - } - - const { usage: _usage, output: _output, ...metadata } = response; - if (Object.keys(metadata).length > 0) { - data.metadata = metadata; - } - - data.metrics = parseMetricsFromUsage(response.usage); - return data; - }, - }), + openAIChannels.responsesStream.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "openai.responses.create", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => extractOpenAIResponsesInput(params), + extractFromEvent: (event) => { + if (event.type !== "response.completed" || !event.response) { + return {}; + } + + const response = event.response; + const data: Record = {}; + + if (response.output !== undefined) { + data.output = processImagesInOutput(response.output); + } + + const { + usage: _usage, + output: _output, + ...metadata + } = response; + if (Object.keys(metadata).length > 0) { + data.metadata = metadata; + } + + data.metrics = parseMetricsFromUsage(response.usage); + return data; + }, + }, + ), + ), ); // Responses API - parse this.unsubscribers.push( - traceStreamingChannel(openAIChannels.responsesParse, { - name: "openai.responses.parse", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => extractOpenAIResponsesInput(params), - extractOutput: (result) => { - return processImagesInOutput(result?.output); - }, - extractMetadata: (result) => extractOpenAIResponsesMetadata(result), - extractMetrics: (result, startTime, endEvent) => { - const metrics = withCachedMetric( - parseMetricsFromUsage(result?.usage), - result, - endEvent, - ); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - aggregateChunks: aggregateResponseStreamEvents, - }), + openAIChannels.responsesParse.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "openai.responses.parse", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => extractOpenAIResponsesInput(params), + extractOutput: (result) => { + return processImagesInOutput(result?.output); + }, + extractMetadata: (result) => + extractOpenAIResponsesMetadata(result), + extractMetrics: (result, startTime, endEvent) => { + const metrics = withCachedMetric( + parseMetricsFromUsage(result?.usage), + result, + endEvent, + ); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + aggregateChunks: aggregateResponseStreamEvents, + }, + ), + ), ); // Responses API - compact this.unsubscribers.push( - traceAsyncChannel(openAIChannels.responsesCompact, { - name: "openai.responses.compact", - type: SpanTypeAttribute.LLM, - extractInput: ([params]) => extractOpenAIResponsesInput(params), - extractOutput: (result) => { - return processImagesInOutput(result?.output); - }, - extractMetadata: (result) => extractOpenAIResponsesMetadata(result), - extractMetrics: (result, startTime, endEvent) => { - const metrics = withCachedMetric( - parseMetricsFromUsage(result?.usage), - result, - endEvent, - ); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - }), + openAIChannels.responsesCompact.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { + ...additional, + arguments: args, + self: receiver, + get response() { + return additional.responseInfo?.response ?? additional.response; + }, + }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENAI, + name: "openai.responses.compact", + type: SpanTypeAttribute.LLM, + extractInput: ([params]) => extractOpenAIResponsesInput(params), + extractOutput: (result) => { + return processImagesInOutput(result?.output); + }, + extractMetadata: (result) => + extractOpenAIResponsesMetadata(result), + extractMetrics: (result, startTime, endEvent) => { + const metrics = withCachedMetric( + parseMetricsFromUsage(result?.usage), + result, + endEvent, + ); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + }, + ), + ), ); } diff --git a/js/src/instrumentation/plugins/openrouter-agent-channels.ts b/js/src/instrumentation/plugins/openrouter-agent-channels.ts index a8fa85015..979545e6a 100644 --- a/js/src/instrumentation/plugins/openrouter-agent-channels.ts +++ b/js/src/instrumentation/plugins/openrouter-agent-channels.ts @@ -1,45 +1,38 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { OpenRouterAgentCallModelArgs, OpenRouterAgentCallModelRequest, } from "../../vendor-sdk-types/openrouter-agent"; -export const openRouterAgentChannels = defineChannels( - "@openrouter/agent", - { - callModel: channel({ - channelName: "callModel", - kind: "sync-stream", - }), +export const openRouterAgentChannels = defineInterceptor("@openrouter/agent", { + callModel: channel({ + channelName: "callModel", + }), - callModelTurn: channel< - [OpenRouterAgentCallModelRequest | undefined], - unknown, - { - step: number; - stepType: "initial" | "continue"; - } - >({ - channelName: "callModel.turn", - kind: "async", - }), + callModelTurn: channel< + [OpenRouterAgentCallModelRequest | undefined], + PromiseLike, + { + step: number; + stepType: "initial" | "continue"; + } + >({ + channelName: "callModel.turn", + }), - toolExecute: channel< - [unknown], - unknown | AsyncIterable, - { - span_info?: { - name?: string; - }; - toolCallId?: string; - toolName: string; - }, - unknown - >({ - channelName: "tool.execute", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER_AGENT }, -); + toolExecute: channel< + [unknown], + unknown | AsyncIterable, + { + span_info?: { + name?: string; + }; + toolCallId?: string; + toolName: string; + }, + unknown + >({ + channelName: "tool.execute", + }), +}); diff --git a/js/src/instrumentation/plugins/openrouter-agent-plugin.test.ts b/js/src/instrumentation/plugins/openrouter-agent-plugin.test.ts index 85d66d51c..8f50f9dab 100644 --- a/js/src/instrumentation/plugins/openrouter-agent-plugin.test.ts +++ b/js/src/instrumentation/plugins/openrouter-agent-plugin.test.ts @@ -7,8 +7,8 @@ import { it, vi, } from "vitest"; -import { configureNode } from "../../node/config"; import { _exportsForTestingOnly, initLogger } from "../../logger"; +import { configureNode } from "../../node/config"; import { openRouterAgentChannels } from "./openrouter-agent-channels"; import { aggregateOpenRouterChatChunks, @@ -459,7 +459,7 @@ describe("OpenRouter Agent Plugin", () => { tools: [tool], }; - const result = openRouterAgentChannels.callModel.traceSync( + const result = openRouterAgentChannels.callModel.invoke( () => { const modelResult = { allToolExecutionRounds: [] as any[], @@ -515,9 +515,9 @@ describe("OpenRouter Agent Plugin", () => { return modelResult; }, - { - arguments: [request as any], - }, + undefined, + [request as any], + {}, ); expect(request.tools[0]).not.toBe(tool); @@ -635,7 +635,7 @@ describe("OpenRouter Agent Plugin", () => { model: "openai/gpt-4.1-mini", }; - const result = openRouterAgentChannels.callModel.traceSync( + const result = openRouterAgentChannels.callModel.invoke( () => ({ async getResponse() { return finalResponse; @@ -644,9 +644,9 @@ describe("OpenRouter Agent Plugin", () => { return "ok"; }, }), - { - arguments: [request], - }, + undefined, + [request], + {}, ); await expect(result.getText()).resolves.toBe("ok"); diff --git a/js/src/instrumentation/plugins/openrouter-agent-plugin.ts b/js/src/instrumentation/plugins/openrouter-agent-plugin.ts index facda63c0..7a02b581e 100644 --- a/js/src/instrumentation/plugins/openrouter-agent-plugin.ts +++ b/js/src/instrumentation/plugins/openrouter-agent-plugin.ts @@ -1,30 +1,29 @@ +import { INSTRUMENTATION_NAMES } from "../../span-origin"; import { BasePlugin, toLoggedError } from "../core"; import { - traceAsyncChannel, - traceStreamingChannel, - traceSyncStreamChannel, + traceAsyncCall, + traceStreamingCall, + traceSyncStreamCall, unsubscribeAll, } from "../core/channel-tracing"; -import type { ChannelMessage } from "../core/channel-definitions"; -import { - SpanTypeAttribute, - isObject, - isPromiseLike, -} from "../../../util/index"; -import { withCurrent } from "../../logger"; +import { runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute, isObject } from "../../../util/index"; import type { Span } from "../../logger"; -import { zodToJsonSchema } from "../../zod/utils"; -import { openRouterAgentChannels } from "./openrouter-agent-channels"; +import { withCurrent } from "../../logger"; import type { + OpenRouterAgentCallModelRequest, OpenRouterAgentChatChoice, OpenRouterAgentChatCompletionChunk, - OpenRouterAgentCallModelRequest, OpenRouterAgentEmbeddingResponse, OpenRouterAgentResponse, OpenRouterAgentResponseStreamEvent, OpenRouterAgentTool, OpenRouterAgentToolTurnContext, } from "../../vendor-sdk-types/openrouter-agent"; +import { zodToJsonSchema } from "../../zod/utils"; +import type { ChannelMessage } from "../core/tracing-types"; +import { openRouterAgentChannels } from "./openrouter-agent-channels"; export class OpenRouterAgentPlugin extends BasePlugin { protected onEnable(): void { @@ -37,113 +36,137 @@ export class OpenRouterAgentPlugin extends BasePlugin { private subscribeToOpenRouterAgentChannels(): void { this.unsubscribers.push( - traceSyncStreamChannel(openRouterAgentChannels.callModel, { - name: "openrouter.callModel", - type: SpanTypeAttribute.TASK, - extractInput: (args) => { - const request = getOpenRouterCallModelRequestArg(args); - return { - input: request - ? extractOpenRouterCallModelInput(request) - : undefined, - metadata: request - ? extractOpenRouterCallModelMetadata(request) - : { provider: "openrouter" }, - }; - }, - patchResult: ({ endEvent, result, span }) => { - return patchOpenRouterCallModelResult({ - request: getOpenRouterCallModelRequestArg(endEvent.arguments), - result, - span, - }); - }, - }), + openRouterAgentChannels.callModel.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER_AGENT, + name: "openrouter.callModel", + type: SpanTypeAttribute.TASK, + extractInput: (args) => { + const request = getOpenRouterCallModelRequestArg(args); + return { + input: request + ? extractOpenRouterCallModelInput(request) + : undefined, + metadata: request + ? extractOpenRouterCallModelMetadata(request) + : { provider: "openrouter" }, + }; + }, + patchResult: ({ endEvent, result, span }) => { + return patchOpenRouterCallModelResult({ + request: getOpenRouterCallModelRequestArg(endEvent.arguments), + result, + span, + }); + }, + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(openRouterAgentChannels.callModelTurn, { - name: "openrouter.beta.responses.send", - type: SpanTypeAttribute.LLM, - extractInput: (args, event) => { - const request = getOpenRouterCallModelRequestArg(args); - const metadata = request - ? extractOpenRouterCallModelMetadata(request) - : { provider: "openrouter" }; - - if (isObject(metadata) && "tools" in metadata) { - delete (metadata as Record).tools; - } - - return { - input: request - ? extractOpenRouterCallModelInput(request) - : undefined, - metadata: { - ...metadata, - step: event.step, - step_type: event.stepType, + openRouterAgentChannels.callModelTurn.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER_AGENT, + name: "openrouter.beta.responses.send", + type: SpanTypeAttribute.LLM, + extractInput: (args, event) => { + const request = getOpenRouterCallModelRequestArg(args); + const metadata = request + ? extractOpenRouterCallModelMetadata(request) + : { provider: "openrouter" }; + + if (isObject(metadata) && "tools" in metadata) { + delete (metadata as Record).tools; + } + + return { + input: request + ? extractOpenRouterCallModelInput(request) + : undefined, + metadata: { + ...metadata, + step: event.step, + step_type: event.stepType, + }, + }; + }, + extractOutput: (result) => + extractOpenRouterResponseOutput( + result as Record, + ), + extractMetadata: (result, event) => { + if (!isObject(result)) { + return { + step: event?.step, + step_type: event?.stepType, + }; + } + + return { + ...(extractOpenRouterResponseMetadata(result) || {}), + ...(event?.step !== undefined ? { step: event.step } : {}), + ...(event?.stepType ? { step_type: event.stepType } : {}), + }; + }, + extractMetrics: (result) => + isObject(result) + ? parseOpenRouterMetricsFromUsage(result.usage) + : {}, }, - }; - }, - extractOutput: (result) => - extractOpenRouterResponseOutput(result as Record), - extractMetadata: (result, event) => { - if (!isObject(result)) { - return { - step: event?.step, - step_type: event?.stepType, - }; - } + ), + ), + ); - return { - ...(extractOpenRouterResponseMetadata(result) || {}), - ...(event?.step !== undefined ? { step: event.step } : {}), - ...(event?.stepType ? { step_type: event.stepType } : {}), - }; - }, - extractMetrics: (result) => - isObject(result) ? parseOpenRouterMetricsFromUsage(result.usage) : {}, - }), + this.unsubscribers.push( + openRouterAgentChannels.toolExecute.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER_AGENT, + name: "openrouter.tool", + type: SpanTypeAttribute.TOOL, + extractInput: (args, event) => ({ + input: args[0], + metadata: { + provider: "openrouter", + tool_name: event.toolName, + ...(event.toolCallId + ? { tool_call_id: event.toolCallId } + : {}), + }, + }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + aggregateChunks: (chunks) => ({ + output: + chunks.length > 0 ? chunks[chunks.length - 1] : undefined, + metrics: {}, + }), + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(openRouterAgentChannels.toolExecute, { - name: "openrouter.tool", - type: SpanTypeAttribute.TOOL, - extractInput: (args, event) => ({ - input: args[0], - metadata: { - provider: "openrouter", - tool_name: event.toolName, - ...(event.toolCallId ? { tool_call_id: event.toolCallId } : {}), - }, - }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - aggregateChunks: (chunks) => ({ - output: chunks.length > 0 ? chunks[chunks.length - 1] : undefined, - metrics: {}, - }), + openRouterAgentChannels.callModel.intercept((target, receiver, args) => { + runInstrumentation(() => { + const request = getOpenRouterCallModelRequestArg(args); + if (request) patchOpenRouterCallModelRequestTools(request); + }); + return Reflect.apply(target, receiver, args); }), ); - - const callModelChannel = openRouterAgentChannels.callModel.tracingChannel(); - const callModelHandlers = { - start: (event: { arguments: unknown[] }) => { - const request = getOpenRouterCallModelRequestArg(event.arguments); - if (!request) { - return; - } - - patchOpenRouterCallModelRequestTools(request); - }, - }; - - callModelChannel.subscribe(callModelHandlers); - this.unsubscribers.push(() => { - callModelChannel.unsubscribe(callModelHandlers); - }); } } @@ -568,7 +591,6 @@ function traceToolExecution(args: { toolCallId?: string; toolName: string; }): unknown { - const tracingChannel = openRouterAgentChannels.toolExecute.tracingChannel(); const input = args.args.length > 0 ? args.args[0] : undefined; const event: OpenRouterToolTraceContext = { arguments: [input], @@ -579,43 +601,12 @@ function traceToolExecution(args: { toolName: args.toolName, }; - tracingChannel.start!.publish(event); - - try { - const result = args.execute(); - return publishToolResult(tracingChannel, event, result); - } catch (error) { - event.error = normalizeError(error); - tracingChannel.error!.publish(event); - throw error; - } -} - -function publishToolResult( - tracingChannel: ReturnType< - typeof openRouterAgentChannels.toolExecute.tracingChannel - >, - event: OpenRouterToolTraceContext, - result: unknown, -): unknown { - if (isPromiseLike(result)) { - return result.then( - (resolved) => { - event.result = resolved; - tracingChannel.asyncEnd!.publish(event); - return resolved; - }, - (error) => { - event.error = normalizeError(error); - tracingChannel.error!.publish(event); - throw error; - }, - ); - } - - event.result = result; - tracingChannel.asyncEnd!.publish(event); - return result; + return openRouterAgentChannels.toolExecute.invoke( + args.execute, + undefined, + [input], + event, + ); } function getToolCallId(context: unknown): string | undefined { @@ -1109,7 +1100,12 @@ async function traceOpenRouterCallModelTurn(args: { }; return await withCurrent(args.parentSpan, () => - openRouterAgentChannels.callModelTurn.tracePromise(args.fn, context), + openRouterAgentChannels.callModelTurn.invoke( + args.fn, + undefined, + [args.request], + context, + ), ); } @@ -1321,8 +1317,4 @@ function isAsyncIterable(value: unknown): value is AsyncIterable { ); } -function normalizeError(error: unknown): Error { - return error instanceof Error ? error : new Error(String(error)); -} - export { parseOpenRouterMetricsFromUsage }; diff --git a/js/src/instrumentation/plugins/openrouter-channels.ts b/js/src/instrumentation/plugins/openrouter-channels.ts index bb0bdc4ad..b69060256 100644 --- a/js/src/instrumentation/plugins/openrouter-channels.ts +++ b/js/src/instrumentation/plugins/openrouter-channels.ts @@ -1,9 +1,9 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { + OpenRouterCallModelRequest, OpenRouterChatCompletion, OpenRouterChatCompletionChunk, - OpenRouterCallModelRequest, OpenRouterChatCreateParams, OpenRouterEmbeddingCreateParams, OpenRouterEmbeddingResponse, @@ -22,77 +22,66 @@ type OpenRouterResponsesResult = | OpenRouterResponse | AsyncIterable; -export const openRouterChannels = defineChannels( - "@openrouter/sdk", - { - chatSend: channel< - [OpenRouterChatCreateParams], - OpenRouterChatResult, - Record, - OpenRouterChatCompletionChunk - >({ - channelName: "chat.send", - kind: "async", - }), +export const openRouterChannels = defineInterceptor("@openrouter/sdk", { + chatSend: channel< + [OpenRouterChatCreateParams], + PromiseLike, + Record, + OpenRouterChatCompletionChunk + >({ + channelName: "chat.send", + }), - embeddingsGenerate: channel< - [OpenRouterEmbeddingCreateParams], - OpenRouterEmbeddingResponse - >({ - channelName: "embeddings.generate", - kind: "async", - }), + embeddingsGenerate: channel< + [OpenRouterEmbeddingCreateParams], + PromiseLike + >({ + channelName: "embeddings.generate", + }), - rerankRerank: channel< - [OpenRouterRerankCreateParams], - OpenRouterRerankResult - >({ - channelName: "rerank.rerank", - kind: "async", - }), + rerankRerank: channel< + [OpenRouterRerankCreateParams], + PromiseLike + >({ + channelName: "rerank.rerank", + }), - betaResponsesSend: channel< - [OpenRouterResponsesCreateParams], - OpenRouterResponsesResult, - Record, - OpenRouterResponseStreamEvent - >({ - channelName: "beta.responses.send", - kind: "async", - }), + betaResponsesSend: channel< + [OpenRouterResponsesCreateParams], + PromiseLike, + Record, + OpenRouterResponseStreamEvent + >({ + channelName: "beta.responses.send", + }), - callModel: channel<[OpenRouterCallModelRequest], unknown>({ - channelName: "callModel", - kind: "sync-stream", - }), + callModel: channel<[OpenRouterCallModelRequest], unknown>({ + channelName: "callModel", + }), - callModelTurn: channel< - [OpenRouterCallModelRequest | undefined], - unknown, - { - step: number; - stepType: "initial" | "continue"; - } - >({ - channelName: "callModel.turn", - kind: "async", - }), + callModelTurn: channel< + [OpenRouterCallModelRequest | undefined], + PromiseLike, + { + step: number; + stepType: "initial" | "continue"; + } + >({ + channelName: "callModel.turn", + }), - toolExecute: channel< - [unknown], - unknown | AsyncIterable, - { - span_info?: { - name?: string; - }; - toolCallId?: string; - toolName: string; - }, - unknown - >({ - channelName: "tool.execute", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER }, -); + toolExecute: channel< + [unknown], + unknown | AsyncIterable, + { + span_info?: { + name?: string; + }; + toolCallId?: string; + toolName: string; + }, + unknown + >({ + channelName: "tool.execute", + }), +}); diff --git a/js/src/instrumentation/plugins/openrouter-plugin.test.ts b/js/src/instrumentation/plugins/openrouter-plugin.test.ts index 93627512e..615fad481 100644 --- a/js/src/instrumentation/plugins/openrouter-plugin.test.ts +++ b/js/src/instrumentation/plugins/openrouter-plugin.test.ts @@ -7,8 +7,8 @@ import { it, vi, } from "vitest"; -import { configureNode } from "../../node/config"; import { _exportsForTestingOnly, initLogger } from "../../logger"; +import { configureNode } from "../../node/config"; import { openRouterChannels } from "./openrouter-channels"; import { aggregateOpenRouterChatChunks, @@ -460,7 +460,7 @@ describe("OpenRouter Plugin", () => { model: "openai/gpt-4.1-mini", tools: [tool], }; - const result = openRouterChannels.callModel.traceSync( + const result = openRouterChannels.callModel.invoke( () => { const modelResult = { allToolExecutionRounds: [] as any[], @@ -516,7 +516,9 @@ describe("OpenRouter Plugin", () => { return modelResult; }, - { arguments: [request as any] }, + undefined, + [request as any], + {}, ); expect(request.tools[0]).not.toBe(tool); @@ -623,7 +625,7 @@ describe("OpenRouter Plugin", () => { }, }; - await openRouterChannels.rerankRerank.tracePromise( + await openRouterChannels.rerankRerank.invoke( async () => ({ id: "rerank_123", model: `${TEST_RERANK_PROVIDER}/${TEST_RERANK_MODEL}`, @@ -644,7 +646,9 @@ describe("OpenRouter Plugin", () => { totalTokens: 5, }, }), - { arguments: [request] }, + undefined, + [request], + {}, ); const spans = await backgroundLogger.drain(); diff --git a/js/src/instrumentation/plugins/openrouter-plugin.ts b/js/src/instrumentation/plugins/openrouter-plugin.ts index c89b8f9a7..be4bbd859 100644 --- a/js/src/instrumentation/plugins/openrouter-plugin.ts +++ b/js/src/instrumentation/plugins/openrouter-plugin.ts @@ -1,25 +1,21 @@ +import { INSTRUMENTATION_NAMES } from "../../span-origin"; import { BasePlugin, toLoggedError } from "../core"; import { - traceAsyncChannel, - traceStreamingChannel, - traceSyncStreamChannel, + traceAsyncCall, + traceStreamingCall, + traceSyncStreamCall, unsubscribeAll, } from "../core/channel-tracing"; -import type { ChannelMessage } from "../core/channel-definitions"; -import { - SpanTypeAttribute, - isObject, - isPromiseLike, -} from "../../../util/index"; -import { withCurrent } from "../../logger"; +import { runInstrumentation } from "../core/observe-result"; + +import { SpanTypeAttribute, isObject } from "../../../util/index"; import type { Span } from "../../logger"; +import { withCurrent } from "../../logger"; import { getCurrentUnixTimestamp } from "../../util"; -import { zodToJsonSchema } from "../../zod/utils"; -import { openRouterChannels } from "./openrouter-channels"; import type { + OpenRouterCallModelRequest, OpenRouterChatChoice, OpenRouterChatCompletionChunk, - OpenRouterCallModelRequest, OpenRouterEmbeddingResponse, OpenRouterRerankResult, OpenRouterResponse, @@ -27,6 +23,9 @@ import type { OpenRouterTool, OpenRouterToolTurnContext, } from "../../vendor-sdk-types/openrouter"; +import { zodToJsonSchema } from "../../zod/utils"; +import type { ChannelMessage } from "../core/tracing-types"; +import { openRouterChannels } from "./openrouter-channels"; export class OpenRouterPlugin extends BasePlugin { protected onEnable(): void { @@ -39,253 +38,329 @@ export class OpenRouterPlugin extends BasePlugin { private subscribeToOpenRouterChannels(): void { this.unsubscribers.push( - traceStreamingChannel(openRouterChannels.chatSend, { - name: "openrouter.chat.send", - type: SpanTypeAttribute.LLM, - extractInput: (args) => { - const request = getOpenRouterRequestArg(args); - const chatGenerationParams = isObject(request?.chatGenerationParams) - ? request.chatGenerationParams - : {}; - const httpReferer = request?.httpReferer; - const xTitle = request?.xTitle; - const { messages, ...metadata } = chatGenerationParams; - return { - input: messages, - metadata: buildOpenRouterMetadata(metadata, httpReferer, xTitle), - }; - }, - extractOutput: (result) => { - return isObject(result) ? result.choices : undefined; - }, - extractMetrics: (result, startTime) => { - const metrics = parseOpenRouterMetricsFromUsage(result?.usage); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - aggregateChunks: aggregateOpenRouterChatChunks, - }), + openRouterChannels.chatSend.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER, + name: "openrouter.chat.send", + type: SpanTypeAttribute.LLM, + extractInput: (args) => { + const request = getOpenRouterRequestArg(args); + const chatGenerationParams = isObject( + request?.chatGenerationParams, + ) + ? request.chatGenerationParams + : {}; + const httpReferer = request?.httpReferer; + const xTitle = request?.xTitle; + const { messages, ...metadata } = chatGenerationParams; + return { + input: messages, + metadata: buildOpenRouterMetadata( + metadata, + httpReferer, + xTitle, + ), + }; + }, + extractOutput: (result) => { + return isObject(result) ? result.choices : undefined; + }, + extractMetrics: (result, startTime) => { + const metrics = parseOpenRouterMetricsFromUsage(result?.usage); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + aggregateChunks: aggregateOpenRouterChatChunks, + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(openRouterChannels.embeddingsGenerate, { - name: "openrouter.embeddings.generate", - type: SpanTypeAttribute.LLM, - extractInput: (args) => { - const request = getOpenRouterRequestArg(args); - const requestBody = isObject(request?.requestBody) - ? request.requestBody - : {}; - const httpReferer = request?.httpReferer; - const xTitle = request?.xTitle; - const { input, ...metadata } = requestBody; - return { - input, - metadata: buildOpenRouterEmbeddingMetadata( - metadata, - httpReferer, - xTitle, - ), - }; - }, - extractOutput: (result) => { - if (!isObject(result)) { - return undefined; - } - - const embedding = result.data?.[0]?.embedding; - return Array.isArray(embedding) - ? { embedding_length: embedding.length } - : undefined; - }, - extractMetadata: (result) => { - if (!isObject(result)) { - return undefined; - } - - return extractOpenRouterResponseMetadata(result); - }, - extractMetrics: (result) => { - return isObject(result) - ? parseOpenRouterMetricsFromUsage(result.usage) - : {}; - }, - }), + openRouterChannels.embeddingsGenerate.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER, + name: "openrouter.embeddings.generate", + type: SpanTypeAttribute.LLM, + extractInput: (args) => { + const request = getOpenRouterRequestArg(args); + const requestBody = isObject(request?.requestBody) + ? request.requestBody + : {}; + const httpReferer = request?.httpReferer; + const xTitle = request?.xTitle; + const { input, ...metadata } = requestBody; + return { + input, + metadata: buildOpenRouterEmbeddingMetadata( + metadata, + httpReferer, + xTitle, + ), + }; + }, + extractOutput: (result) => { + if (!isObject(result)) { + return undefined; + } + + const embedding = result.data?.[0]?.embedding; + return Array.isArray(embedding) + ? { embedding_length: embedding.length } + : undefined; + }, + extractMetadata: (result) => { + if (!isObject(result)) { + return undefined; + } + + return extractOpenRouterResponseMetadata(result); + }, + extractMetrics: (result) => { + return isObject(result) + ? parseOpenRouterMetricsFromUsage(result.usage) + : {}; + }, + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(openRouterChannels.rerankRerank, { - name: "openrouter.rerank.rerank", - type: SpanTypeAttribute.LLM, - extractInput: (args) => { - const request = getOpenRouterRequestArg(args); - const requestBody = isObject(request?.requestBody) - ? request.requestBody - : {}; - const httpReferer = request?.httpReferer; - const xTitle = request?.xTitle ?? request?.appTitle; - const { documents, query, ...metadata } = requestBody; - return { - input: { - documents, - query, + openRouterChannels.rerankRerank.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER, + name: "openrouter.rerank.rerank", + type: SpanTypeAttribute.LLM, + extractInput: (args) => { + const request = getOpenRouterRequestArg(args); + const requestBody = isObject(request?.requestBody) + ? request.requestBody + : {}; + const httpReferer = request?.httpReferer; + const xTitle = request?.xTitle ?? request?.appTitle; + const { documents, query, ...metadata } = requestBody; + return { + input: { + documents, + query, + }, + metadata: buildOpenRouterRerankMetadata( + metadata, + documents, + httpReferer, + xTitle, + ), + }; + }, + extractOutput: (result) => extractOpenRouterRerankOutput(result), + extractMetadata: (result) => + extractOpenRouterResponseMetadata(result), + extractMetrics: (result) => + isObject(result) + ? parseOpenRouterMetricsFromUsage(result.usage) + : {}, }, - metadata: buildOpenRouterRerankMetadata( - metadata, - documents, - httpReferer, - xTitle, - ), - }; - }, - extractOutput: (result) => extractOpenRouterRerankOutput(result), - extractMetadata: (result) => extractOpenRouterResponseMetadata(result), - extractMetrics: (result) => - isObject(result) ? parseOpenRouterMetricsFromUsage(result.usage) : {}, - }), + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(openRouterChannels.betaResponsesSend, { - name: "openrouter.beta.responses.send", - type: SpanTypeAttribute.LLM, - extractInput: (args) => { - const request = getOpenRouterRequestArg(args); - const openResponsesRequest = isObject(request?.openResponsesRequest) - ? request.openResponsesRequest - : {}; - const httpReferer = request?.httpReferer; - const xTitle = request?.xTitle; - const { input, ...metadata } = openResponsesRequest; - return { - input, - metadata: buildOpenRouterMetadata(metadata, httpReferer, xTitle), - }; - }, - extractOutput: (result) => - extractOpenRouterResponseOutput(result as Record), - extractMetadata: (result) => extractOpenRouterResponseMetadata(result), - extractMetrics: (result, startTime) => { - const metrics = parseOpenRouterMetricsFromUsage(result?.usage); - if (startTime) { - metrics.time_to_first_token = getCurrentUnixTimestamp() - startTime; - } - return metrics; - }, - aggregateChunks: aggregateOpenRouterResponseStreamEvents, - }), + openRouterChannels.betaResponsesSend.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER, + name: "openrouter.beta.responses.send", + type: SpanTypeAttribute.LLM, + extractInput: (args) => { + const request = getOpenRouterRequestArg(args); + const openResponsesRequest = isObject( + request?.openResponsesRequest, + ) + ? request.openResponsesRequest + : {}; + const httpReferer = request?.httpReferer; + const xTitle = request?.xTitle; + const { input, ...metadata } = openResponsesRequest; + return { + input, + metadata: buildOpenRouterMetadata( + metadata, + httpReferer, + xTitle, + ), + }; + }, + extractOutput: (result) => + extractOpenRouterResponseOutput( + result as Record, + ), + extractMetadata: (result) => + extractOpenRouterResponseMetadata(result), + extractMetrics: (result, startTime) => { + const metrics = parseOpenRouterMetricsFromUsage(result?.usage); + if (startTime) { + metrics.time_to_first_token = + getCurrentUnixTimestamp() - startTime; + } + return metrics; + }, + aggregateChunks: aggregateOpenRouterResponseStreamEvents, + }, + ), + ), ); this.unsubscribers.push( - traceSyncStreamChannel(openRouterChannels.callModel, { - name: "openrouter.callModel", - type: SpanTypeAttribute.TASK, - extractInput: (args) => { - const request = getOpenRouterCallModelRequestArg(args); - return { - input: request - ? extractOpenRouterCallModelInput(request) - : undefined, - metadata: request - ? extractOpenRouterCallModelMetadata(request) - : { provider: "openrouter" }, - }; - }, - patchResult: ({ endEvent, result, span }) => { - return patchOpenRouterCallModelResult({ - request: getOpenRouterCallModelRequestArg(endEvent.arguments), - result, - span, - }); - }, - }), + openRouterChannels.callModel.intercept( + (target, receiver, args, additional) => + traceSyncStreamCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER, + name: "openrouter.callModel", + type: SpanTypeAttribute.TASK, + extractInput: (args) => { + const request = getOpenRouterCallModelRequestArg(args); + return { + input: request + ? extractOpenRouterCallModelInput(request) + : undefined, + metadata: request + ? extractOpenRouterCallModelMetadata(request) + : { provider: "openrouter" }, + }; + }, + patchResult: ({ endEvent, result, span }) => { + return patchOpenRouterCallModelResult({ + request: getOpenRouterCallModelRequestArg(endEvent.arguments), + result, + span, + }); + }, + }, + ), + ), ); this.unsubscribers.push( - traceAsyncChannel(openRouterChannels.callModelTurn, { - name: "openrouter.beta.responses.send", - type: SpanTypeAttribute.LLM, - extractInput: (args, event) => { - const request = getOpenRouterCallModelRequestArg(args); - const metadata = request - ? extractOpenRouterCallModelMetadata(request) - : { provider: "openrouter" }; - - if (isObject(metadata) && "tools" in metadata) { - delete (metadata as Record).tools; - } - - return { - input: request - ? extractOpenRouterCallModelInput(request) - : undefined, - metadata: { - ...metadata, - step: event.step, - step_type: event.stepType, + openRouterChannels.callModelTurn.intercept( + (target, receiver, args, additional) => + traceAsyncCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER, + name: "openrouter.beta.responses.send", + type: SpanTypeAttribute.LLM, + extractInput: (args, event) => { + const request = getOpenRouterCallModelRequestArg(args); + const metadata = request + ? extractOpenRouterCallModelMetadata(request) + : { provider: "openrouter" }; + + if (isObject(metadata) && "tools" in metadata) { + delete (metadata as Record).tools; + } + + return { + input: request + ? extractOpenRouterCallModelInput(request) + : undefined, + metadata: { + ...metadata, + step: event.step, + step_type: event.stepType, + }, + }; + }, + extractOutput: (result) => + extractOpenRouterResponseOutput( + result as Record, + ), + extractMetadata: (result, event) => { + if (!isObject(result)) { + return { + step: event?.step, + step_type: event?.stepType, + }; + } + + return { + ...(extractOpenRouterResponseMetadata(result) || {}), + ...(event?.step !== undefined ? { step: event.step } : {}), + ...(event?.stepType ? { step_type: event.stepType } : {}), + }; + }, + extractMetrics: (result) => + isObject(result) + ? parseOpenRouterMetricsFromUsage(result.usage) + : {}, }, - }; - }, - extractOutput: (result) => - extractOpenRouterResponseOutput(result as Record), - extractMetadata: (result, event) => { - if (!isObject(result)) { - return { - step: event?.step, - step_type: event?.stepType, - }; - } + ), + ), + ); - return { - ...(extractOpenRouterResponseMetadata(result) || {}), - ...(event?.step !== undefined ? { step: event.step } : {}), - ...(event?.stepType ? { step_type: event.stepType } : {}), - }; - }, - extractMetrics: (result) => - isObject(result) ? parseOpenRouterMetricsFromUsage(result.usage) : {}, - }), + this.unsubscribers.push( + openRouterChannels.toolExecute.intercept( + (target, receiver, args, additional) => + traceStreamingCall( + () => Reflect.apply(target, receiver, args), + { ...additional, arguments: args, self: receiver }, + { + instrumentationName: INSTRUMENTATION_NAMES.OPENROUTER, + name: "openrouter.tool", + type: SpanTypeAttribute.TOOL, + extractInput: (args, event) => ({ + input: args[0], + metadata: { + provider: "openrouter", + tool_name: event.toolName, + ...(event.toolCallId + ? { tool_call_id: event.toolCallId } + : {}), + }, + }), + extractOutput: (result) => result, + extractMetrics: () => ({}), + aggregateChunks: (chunks) => ({ + output: + chunks.length > 0 ? chunks[chunks.length - 1] : undefined, + metrics: {}, + }), + }, + ), + ), ); this.unsubscribers.push( - traceStreamingChannel(openRouterChannels.toolExecute, { - name: "openrouter.tool", - type: SpanTypeAttribute.TOOL, - extractInput: (args, event) => ({ - input: args[0], - metadata: { - provider: "openrouter", - tool_name: event.toolName, - ...(event.toolCallId ? { tool_call_id: event.toolCallId } : {}), - }, - }), - extractOutput: (result) => result, - extractMetrics: () => ({}), - aggregateChunks: (chunks) => ({ - output: chunks.length > 0 ? chunks[chunks.length - 1] : undefined, - metrics: {}, - }), + openRouterChannels.callModel.intercept((target, receiver, args) => { + runInstrumentation(() => { + const request = getOpenRouterCallModelRequestArg(args); + if (request) patchOpenRouterCallModelRequestTools(request); + }); + return Reflect.apply(target, receiver, args); }), ); - - const callModelChannel = openRouterChannels.callModel.tracingChannel(); - const callModelHandlers = { - start: (event: { arguments: unknown[] }) => { - const request = getOpenRouterCallModelRequestArg(event.arguments); - if (!request) { - return; - } - - patchOpenRouterCallModelRequestTools(request); - }, - }; - - callModelChannel.subscribe(callModelHandlers); - this.unsubscribers.push(() => { - callModelChannel.unsubscribe(callModelHandlers); - }); } } @@ -750,7 +825,6 @@ function traceToolExecution(args: { toolCallId?: string; toolName: string; }): unknown { - const tracingChannel = openRouterChannels.toolExecute.tracingChannel(); const input = args.args.length > 0 ? args.args[0] : undefined; const event: OpenRouterToolTraceContext = { arguments: [input], @@ -761,43 +835,12 @@ function traceToolExecution(args: { toolName: args.toolName, }; - tracingChannel.start!.publish(event); - - try { - const result = args.execute(); - return publishToolResult(tracingChannel, event, result); - } catch (error) { - event.error = normalizeError(error); - tracingChannel.error!.publish(event); - throw error; - } -} - -function publishToolResult( - tracingChannel: ReturnType< - typeof openRouterChannels.toolExecute.tracingChannel - >, - event: OpenRouterToolTraceContext, - result: unknown, -): unknown { - if (isPromiseLike(result)) { - return result.then( - (resolved) => { - event.result = resolved; - tracingChannel.asyncEnd!.publish(event); - return resolved; - }, - (error) => { - event.error = normalizeError(error); - tracingChannel.error!.publish(event); - throw error; - }, - ); - } - - event.result = result; - tracingChannel.asyncEnd!.publish(event); - return result; + return openRouterChannels.toolExecute.invoke( + args.execute, + undefined, + [input], + event, + ); } function getToolCallId(context: unknown): string | undefined { @@ -1291,7 +1334,12 @@ async function traceOpenRouterCallModelTurn(args: { }; return await withCurrent(args.parentSpan, () => - openRouterChannels.callModelTurn.tracePromise(args.fn, context), + openRouterChannels.callModelTurn.invoke( + args.fn, + undefined, + [args.request], + context, + ), ); } @@ -1503,8 +1551,4 @@ function isAsyncIterable(value: unknown): value is AsyncIterable { ); } -function normalizeError(error: unknown): Error { - return error instanceof Error ? error : new Error(String(error)); -} - export { parseOpenRouterMetricsFromUsage }; diff --git a/js/src/instrumentation/plugins/pi-coding-agent-channels.ts b/js/src/instrumentation/plugins/pi-coding-agent-channels.ts index c537618bb..cd7f1da7a 100644 --- a/js/src/instrumentation/plugins/pi-coding-agent-channels.ts +++ b/js/src/instrumentation/plugins/pi-coding-agent-channels.ts @@ -1,21 +1,19 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { PiAgentSession, PiPromptOptions, } from "../../vendor-sdk-types/pi-coding-agent"; -export const piCodingAgentChannels = defineChannels( +export const piCodingAgentChannels = defineInterceptor( "@earendil-works/pi-coding-agent", { prompt: channel< [string, PiPromptOptions | undefined], - void, + PromiseLike, { session?: PiAgentSession } >({ channelName: "AgentSession.prompt", - kind: "async", }), }, - { instrumentationName: INSTRUMENTATION_NAMES.PI_CODING_AGENT }, ); diff --git a/js/src/instrumentation/plugins/pi-coding-agent-plugin.test.ts b/js/src/instrumentation/plugins/pi-coding-agent-plugin.test.ts index 73b5afd24..3b8ea79b8 100644 --- a/js/src/instrumentation/plugins/pi-coding-agent-plugin.test.ts +++ b/js/src/instrumentation/plugins/pi-coding-agent-plugin.test.ts @@ -20,11 +20,17 @@ vi.mock("../../isomorph", async (importOriginal) => { default: { ...actual.default, newAsyncLocalStorage: () => new AsyncLocalStorage(), - newTracingChannel: mockNewTracingChannel, }, }; }); +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: mockNewTracingChannel, +})); + vi.mock("../../logger", () => ({ startSpan: (...args: unknown[]) => mockStartSpan(...args), withCurrent: (_span: unknown, callback: () => unknown) => callback(), diff --git a/js/src/instrumentation/plugins/pi-coding-agent-plugin.ts b/js/src/instrumentation/plugins/pi-coding-agent-plugin.ts index 174ded6ce..28c374976 100644 --- a/js/src/instrumentation/plugins/pi-coding-agent-plugin.ts +++ b/js/src/instrumentation/plugins/pi-coding-agent-plugin.ts @@ -1,21 +1,15 @@ import { BasePlugin, toLoggedError } from "../core"; -import type { ChannelMessage } from "../core/channel-definitions"; -import iso, { type IsoAsyncLocalStorage } from "../../isomorph"; + +import { SpanTypeAttribute, isObject } from "../../../util/index"; import { debugLogger } from "../../debug-logger"; -import { startSpan as startBaseSpan, withCurrent } from "../../logger"; +import iso, { type IsoAsyncLocalStorage } from "../../isomorph"; import type { Span } from "../../logger"; +import { startSpan as startBaseSpan, withCurrent } from "../../logger"; import { INSTRUMENTATION_NAMES, withSpanInstrumentationName, } from "../../span-origin"; import { getCurrentUnixTimestamp } from "../../util"; -import { SpanTypeAttribute, isObject } from "../../../util/index"; -import { processInputAttachments } from "../../wrappers/attachment-utils"; -import { - runWithAutoInstrumentationAllowed, - runWithAutoInstrumentationSuppressed, -} from "../auto-instrumentation-suppression"; -import { piCodingAgentChannels } from "./pi-coding-agent-channels"; import type { PiAgent, PiAgentEvent, @@ -35,6 +29,13 @@ import type { PiToolCall, PiToolResultMessage, } from "../../vendor-sdk-types/pi-coding-agent"; +import { processInputAttachments } from "../../wrappers/attachment-utils"; +import { + runWithAutoInstrumentationAllowed, + runWithAutoInstrumentationSuppressed, +} from "../auto-instrumentation-suppression"; +import type { ChannelMessage } from "../core/tracing-types"; +import { piCodingAgentChannels } from "./pi-coding-agent-channels"; type PiPromptState = { activeLlmSpans: Set; diff --git a/js/src/instrumentation/plugins/strands-agent-sdk-channels.ts b/js/src/instrumentation/plugins/strands-agent-sdk-channels.ts index f385f308e..f55bf6d2a 100644 --- a/js/src/instrumentation/plugins/strands-agent-sdk-channels.ts +++ b/js/src/instrumentation/plugins/strands-agent-sdk-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { StrandsAgent, StrandsAgentResult, @@ -17,7 +17,7 @@ type StrandsChannelContext = { self?: unknown; }; -export const strandsAgentSDKChannels = defineChannels( +export const strandsAgentSDKChannels = defineInterceptor( "@strands-agents/sdk", { agentStream: channel< @@ -27,7 +27,6 @@ export const strandsAgentSDKChannels = defineChannels( StrandsAgentStreamEvent >({ channelName: "Agent.stream", - kind: "sync-stream", }), graphStream: channel< @@ -41,7 +40,6 @@ export const strandsAgentSDKChannels = defineChannels( StrandsMultiAgentStreamEvent >({ channelName: "Graph.stream", - kind: "sync-stream", }), swarmStream: channel< @@ -55,8 +53,6 @@ export const strandsAgentSDKChannels = defineChannels( StrandsMultiAgentStreamEvent >({ channelName: "Swarm.stream", - kind: "sync-stream", }), }, - { instrumentationName: INSTRUMENTATION_NAMES.STRANDS_AGENT_SDK }, ); diff --git a/js/src/instrumentation/plugins/strands-agent-sdk-plugin.test.ts b/js/src/instrumentation/plugins/strands-agent-sdk-plugin.test.ts index a2b3f0527..7397d107f 100644 --- a/js/src/instrumentation/plugins/strands-agent-sdk-plugin.test.ts +++ b/js/src/instrumentation/plugins/strands-agent-sdk-plugin.test.ts @@ -1,4 +1,11 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../../global-instrumentation-hooks"; +vi.mock("../../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal< + typeof import("../../global-instrumentation-hooks") + >()), + newGlobalInvocationHook: vi.fn(), +})); const { mockWithCurrent, mockNewAsyncLocalStorage, mockStartSpan } = vi.hoisted( () => ({ @@ -26,7 +33,6 @@ vi.mock("../../isomorph", () => ({ default: { getEnv: vi.fn(), newAsyncLocalStorage: mockNewAsyncLocalStorage, - newTracingChannel: vi.fn(), }, })); @@ -39,12 +45,13 @@ vi.mock("../../logger", async (importOriginal) => { }; }); -import iso from "../../isomorph"; import { Attachment } from "../../logger"; import { isAutoInstrumentationSuppressed } from "../auto-instrumentation-suppression"; import { StrandsAgentSDKPlugin } from "./strands-agent-sdk-plugin"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; describe("StrandsAgentSDKPlugin", () => { let handlersByName: Map; @@ -65,7 +72,7 @@ describe("StrandsAgentSDKPlugin", () => { beforeEach(() => { handlersByName = new Map(); spans = []; - mockNewTracingChannel.mockImplementation((name: string) => ({ + mockNewInvocationHook.mockImplementation((name: string) => ({ intercept: vi.fn((interceptor) => { const handlers = { end: (event: any) => diff --git a/js/src/instrumentation/plugins/typesafe-channels.ts b/js/src/instrumentation/plugins/typesafe-channels.ts index 5c54fa41b..f7cae57bd 100644 --- a/js/src/instrumentation/plugins/typesafe-channels.ts +++ b/js/src/instrumentation/plugins/typesafe-channels.ts @@ -1,20 +1,14 @@ -import { INSTRUMENTATION_NAMES } from "../../span-origin"; import type { TypeSafeSystemOneRequest, TypeSafeSystemOneResult, } from "../../vendor-sdk-types/typesafe"; -import { channel, defineChannels } from "../core/channel-definitions"; +import { channel, defineInterceptor } from "../core/channel-definitions"; -export const typeSafeChannels = defineChannels( - "@typesafe-ai/sdk", - { - systemOne: channel< - [TypeSafeSystemOneRequest, options?: unknown], - TypeSafeSystemOneResult - >({ - channelName: "systemOne", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.TYPESAFE }, -); +export const typeSafeChannels = defineInterceptor("@typesafe-ai/sdk", { + systemOne: channel< + [TypeSafeSystemOneRequest, options?: unknown], + PromiseLike + >({ + channelName: "systemOne", + }), +}); diff --git a/js/src/instrumentation/plugins/voyageai-channels.ts b/js/src/instrumentation/plugins/voyageai-channels.ts index 757337007..1f31199e5 100644 --- a/js/src/instrumentation/plugins/voyageai-channels.ts +++ b/js/src/instrumentation/plugins/voyageai-channels.ts @@ -1,5 +1,5 @@ -import { channel, defineChannels } from "../core/channel-definitions"; -import { INSTRUMENTATION_NAMES } from "../../span-origin"; +import { channel, defineInterceptor } from "../core/channel-definitions"; + import type { VoyageAIContextualizedEmbedRequest, VoyageAIContextualizedResult, @@ -10,40 +10,32 @@ import type { VoyageAIRerankResponse, } from "../../vendor-sdk-types/voyageai"; -export const voyageAIChannels = defineChannels( - "voyageai", - { - embed: channel< - [VoyageAIEmbedRequest, options?: unknown], - VoyageAIEmbeddingResponse - >({ - channelName: "embed", - kind: "async", - }), +export const voyageAIChannels = defineInterceptor("voyageai", { + embed: channel< + [VoyageAIEmbedRequest, options?: unknown], + PromiseLike + >({ + channelName: "embed", + }), - multimodalEmbed: channel< - [VoyageAIMultimodalEmbedRequest, options?: unknown], - VoyageAIEmbeddingResponse - >({ - channelName: "multimodalEmbed", - kind: "async", - }), + multimodalEmbed: channel< + [VoyageAIMultimodalEmbedRequest, options?: unknown], + PromiseLike + >({ + channelName: "multimodalEmbed", + }), - rerank: channel< - [VoyageAIRerankRequest, options?: unknown], - VoyageAIRerankResponse - >({ - channelName: "rerank", - kind: "async", - }), + rerank: channel< + [VoyageAIRerankRequest, options?: unknown], + PromiseLike + >({ + channelName: "rerank", + }), - contextualizedEmbed: channel< - [VoyageAIContextualizedEmbedRequest, options?: unknown], - VoyageAIContextualizedResult - >({ - channelName: "contextualizedEmbed", - kind: "async", - }), - }, - { instrumentationName: INSTRUMENTATION_NAMES.VOYAGEAI }, -); + contextualizedEmbed: channel< + [VoyageAIContextualizedEmbedRequest, options?: unknown], + PromiseLike + >({ + channelName: "contextualizedEmbed", + }), +}); diff --git a/js/src/instrumentation/plugins/voyageai-plugin.ts b/js/src/instrumentation/plugins/voyageai-plugin.ts index 37f864b26..25547b4ef 100644 --- a/js/src/instrumentation/plugins/voyageai-plugin.ts +++ b/js/src/instrumentation/plugins/voyageai-plugin.ts @@ -31,32 +31,46 @@ const RERANK_METADATA_ALLOWLIST = new Set([ export class VoyageAIPlugin extends BasePlugin { protected onEnable(): void { this.unsubscribers.push( - interceptVoyageAICall( - voyageAIChannels.embed, - "voyageai.embed", - extractTextEmbeddingInput, - summarizeEmbeddingOutput, - extractEmbeddingUsageMetrics, + voyageAIChannels.embed.intercept((target, receiver, args, additional) => + traceVoyageAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "voyageai.embed", + extractTextEmbeddingInput, + summarizeEmbeddingOutput, + extractEmbeddingUsageMetrics, + ), ), - interceptVoyageAICall( - voyageAIChannels.multimodalEmbed, - "voyageai.multimodalEmbed", - extractMultimodalEmbeddingInput, - summarizeEmbeddingOutput, - extractEmbeddingUsageMetrics, + voyageAIChannels.multimodalEmbed.intercept( + (target, receiver, args, additional) => + traceVoyageAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "voyageai.multimodalEmbed", + extractMultimodalEmbeddingInput, + summarizeEmbeddingOutput, + extractEmbeddingUsageMetrics, + ), ), - interceptVoyageAICall( - voyageAIChannels.rerank, - "voyageai.rerank", - extractRerankInput, - summarizeRerankOutput, + voyageAIChannels.rerank.intercept((target, receiver, args, additional) => + traceVoyageAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "voyageai.rerank", + extractRerankInput, + summarizeRerankOutput, + ), ), - interceptVoyageAICall( - voyageAIChannels.contextualizedEmbed, - "voyageai.contextualizedEmbed", - extractContextualizedEmbeddingInput, - summarizeContextualizedEmbeddingOutput, - extractEmbeddingUsageMetrics, + voyageAIChannels.contextualizedEmbed.intercept( + (target, receiver, args, additional) => + traceVoyageAICall( + () => Reflect.apply(target, receiver, args), + { arguments: args, self: receiver, additional }, + "voyageai.contextualizedEmbed", + extractContextualizedEmbeddingInput, + summarizeContextualizedEmbeddingOutput, + extractEmbeddingUsageMetrics, + ), ), ); } @@ -71,21 +85,12 @@ type VoyageAIResult = | VoyageAIRerankResponse | VoyageAIContextualizedResult; -type VoyageAIChannel = { - intercept( - interceptor: ( - target: (this: unknown, ...args: TArgs) => PromiseLike, - thisArg: unknown, - args: TArgs, - ) => PromiseLike, - ): () => void; -}; - -function interceptVoyageAICall< +function traceVoyageAICall< TArgs extends unknown[], TResult extends VoyageAIResult, >( - channel: VoyageAIChannel, + call: () => PromiseLike, + context: { arguments: TArgs; self: unknown; additional: unknown }, name: string, extractInput: (args: TArgs) => { input: unknown; @@ -95,55 +100,51 @@ function interceptVoyageAICall< extractMetrics: ( result: TResult, ) => Record = extractUsageMetrics, -): () => void { - return channel.intercept((target, thisArg, args) => { - const invokeTarget = () => Reflect.apply(target, thisArg, args); - if (isAutoInstrumentationSuppressed()) { - return invokeTarget(); - } - - let span: Span; - try { - const { input, metadata } = extractInput(args); - span = startSpan( - withSpanInstrumentationName( - { - event: { input, metadata }, - name, - spanAttributes: { type: SpanTypeAttribute.LLM }, - }, - INSTRUMENTATION_NAMES.VOYAGEAI, - ), - ); - } catch (error) { - debugLogger.error(`Error starting span for ${name}:`, error); - return invokeTarget(); - } - - let result: PromiseLike; - try { - result = withCurrent(span, () => - runWithAutoInstrumentationSuppressed(invokeTarget), - ); - } catch (error) { - finishVoyageAISpan(span, name, () => span.log({ error })); - throw error; - } +): PromiseLike { + const args = context.arguments; - void Promise.resolve(result).then( - (value) => - finishVoyageAISpan(span, name, () => { - const metadata = extractResponseMetadata(value); - span.log({ - output: extractOutput(value), - ...(metadata ? { metadata } : {}), - metrics: extractMetrics(value), - }); - }), - (error) => finishVoyageAISpan(span, name, () => span.log({ error })), + if (isAutoInstrumentationSuppressed()) { + return call(); + } + let span: Span; + try { + const { input, metadata } = extractInput(args); + span = startSpan( + withSpanInstrumentationName( + { + event: { input, metadata }, + name, + spanAttributes: { type: SpanTypeAttribute.LLM }, + }, + INSTRUMENTATION_NAMES.VOYAGEAI, + ), + ); + } catch (error) { + debugLogger.error(`Error starting span for ${name}:`, error); + return call(); + } + let result: PromiseLike; + try { + result = withCurrent(span, () => + runWithAutoInstrumentationSuppressed(call), ); - return result; - }); + } catch (error) { + finishVoyageAISpan(span, name, () => span.log({ error })); + throw error; + } + void Promise.resolve(result).then( + (value) => + finishVoyageAISpan(span, name, () => { + const metadata = extractResponseMetadata(value); + span.log({ + output: extractOutput(value), + ...(metadata ? { metadata } : {}), + metrics: extractMetrics(value), + }); + }), + (error) => finishVoyageAISpan(span, name, () => span.log({ error })), + ); + return result; } function finishVoyageAISpan(span: Span, name: string, log: () => void): void { diff --git a/js/src/instrumentation/registry.test.ts b/js/src/instrumentation/registry.test.ts index 9d01c94ef..fb8707de4 100644 --- a/js/src/instrumentation/registry.test.ts +++ b/js/src/instrumentation/registry.test.ts @@ -1,9 +1,13 @@ -import { describe, it, expect, beforeEach, afterEach, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { newGlobalInvocationHook } from "../global-instrumentation-hooks"; +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(), +})); -// Mock iso's newTracingChannel - must be before any imports that use it +// Mock platform context independently of invocation hooks. vi.mock("../isomorph", () => ({ default: { - newTracingChannel: vi.fn(), newAsyncLocalStorage: vi.fn(() => ({ getStore: vi.fn(() => undefined), run: vi.fn((_store: unknown, callback: () => unknown) => callback()), @@ -13,20 +17,21 @@ vi.mock("../isomorph", () => ({ }, })); -import { registry, configureInstrumentation } from "./registry"; -import iso from "../isomorph"; +import { configureInstrumentation, registry } from "./registry"; -const mockNewTracingChannel = iso.newTracingChannel as ReturnType; +const mockNewInvocationHook = newGlobalInvocationHook as ReturnType< + typeof vi.fn +>; describe("Plugin Registry", () => { beforeEach(() => { // Setup mock channel const mockChannel = { - subscribe: vi.fn(), - unsubscribe: vi.fn(), + intercept: vi.fn(() => vi.fn()), + unintercept: vi.fn(() => vi.fn()), hasSubscribers: false, }; - mockNewTracingChannel.mockReturnValue(mockChannel); + mockNewInvocationHook.mockReturnValue(mockChannel); }); // Clean up after each test diff --git a/js/src/instrumentation/test-utils/invocation.ts b/js/src/instrumentation/test-utils/invocation.ts new file mode 100644 index 000000000..04167d6e9 --- /dev/null +++ b/js/src/instrumentation/test-utils/invocation.ts @@ -0,0 +1,61 @@ +import type { GlobalInvocationInterceptor } from "../../global-instrumentation-hooks"; + +type Call = { + arguments?: unknown[]; + self?: unknown; + result?: unknown; + error?: unknown; + [key: string]: unknown; +}; + +/** Drive a real interceptor with controllable provider settlement, without tracing events. */ +export function invocationController(interceptor: GlobalInvocationInterceptor) { + const pending = new WeakMap< + Call, + { resolve: (value: unknown) => void; reject: (error: unknown) => void } + >(); + return { + begin(call: Call) { + const result = { + then( + resolve: (value: unknown) => void, + reject: (error: unknown) => void, + ) { + pending.set(call, { resolve, reject }); + }, + }; + return interceptor(() => result, call.self, call.arguments ?? [], call); + }, + resolve(call: Call) { + const operation = pending.get(call); + if (!operation) return; + pending.delete(call); + return operation.resolve(call.result); + }, + reject(call: Call) { + const operation = pending.get(call); + if (!operation) return; + pending.delete(call); + return operation.reject(call.error); + }, + call(call: Call, target = () => call.result) { + return interceptor(target, call.self, call.arguments ?? [], call); + }, + throw(call: Call) { + try { + interceptor( + () => { + throw call.error; + }, + call.self, + call.arguments ?? [], + call, + ); + } catch (error) { + if (error !== call.error) throw error; + return; + } + throw new Error("Interceptor swallowed the provider error"); + }, + }; +} diff --git a/js/src/isomorph.ts b/js/src/isomorph.ts index a6adf3b48..354adf2cb 100644 --- a/js/src/isomorph.ts +++ b/js/src/isomorph.ts @@ -2,13 +2,6 @@ import { type GitMetadataSettingsType as GitMetadataSettings, type RepoInfoType as RepoInfo, } from "./generated_types"; -import { - newGlobalTracingChannel, - type GlobalHookAsyncLocalStorage, - type GlobalHookHandlers, - type GlobalTracingChannel, - type GlobalTracingChannelCollection, -} from "./global-instrumentation-hooks"; export interface CallerLocation { caller_functionname: string; @@ -16,7 +9,10 @@ export interface CallerLocation { caller_lineno: number; } -export type IsoAsyncLocalStorage = GlobalHookAsyncLocalStorage; +export interface IsoAsyncLocalStorage { + run(store: T | undefined, callback: () => R): R; + getStore(): T | undefined; +} class DefaultAsyncLocalStorage implements IsoAsyncLocalStorage { constructor() {} @@ -29,10 +25,6 @@ class DefaultAsyncLocalStorage implements IsoAsyncLocalStorage { } } -type IsoTracingChannelCollection = GlobalTracingChannelCollection; -export type IsoTracingChannel = GlobalTracingChannel; -export type IsoChannelHandlers = GlobalHookHandlers; - interface Common { buildType: | "browser" // deprecated, use /workerd or /edge-light entrypoints for edge environments @@ -51,10 +43,6 @@ interface Common { getCallerLocation: () => CallerLocation | undefined; newAsyncLocalStorage: () => IsoAsyncLocalStorage; // eslint-disable-next-line @typescript-eslint/no-explicit-any - newTracingChannel: ( - nameOrChannels: string | IsoTracingChannelCollection, - ) => IsoTracingChannel; - // eslint-disable-next-line @typescript-eslint/no-explicit-any processOn: (event: string, handler: (code: any) => void) => void; // hash a string. not guaranteed to be crypto safe. @@ -111,9 +99,7 @@ const iso: Common = { getCallerLocation: () => undefined, newAsyncLocalStorage: () => new DefaultAsyncLocalStorage(), // eslint-disable-next-line @typescript-eslint/no-explicit-any - newTracingChannel: ( - nameOrChannels: string | IsoTracingChannelCollection, - ) => newGlobalTracingChannel(nameOrChannels), + processOn: (_0, _1) => {}, basename: (filepath: string) => filepath.split(/[\\/]/).pop() || filepath, // eslint-disable-next-line no-restricted-properties -- preserving intentional console usage. diff --git a/js/src/openai-promise-utils.test.ts b/js/src/openai-promise-utils.test.ts index dc474c316..4a6393636 100644 --- a/js/src/openai-promise-utils.test.ts +++ b/js/src/openai-promise-utils.test.ts @@ -71,20 +71,25 @@ describe("OpenAI API promise wrappers", () => { expect(tracedAsResponseCalls).toBe(1); }); - it("dispatches the OpenAI channel lifecycle for asResponse-only calls", async () => { + it("invokes the OpenAI hook for asResponse-only calls", async () => { const phases: string[] = []; let result: unknown; - const tracingChannel = - openAIChannels.chatCompletionsCreate.tracingChannel(); - const handlers = { - asyncEnd: (event: { result?: unknown }) => { - phases.push("asyncEnd"); - result = event.result; + const remove = openAIChannels.chatCompletionsCreate.intercept( + (target, receiver, args) => { + phases.push("called"); + const promise = Reflect.apply(target, receiver, args); + promise.then( + (value) => { + phases.push("resolved"); + result = value; + }, + () => { + phases.push("error"); + }, + ); + return promise; }, - error: () => phases.push("error"), - start: () => phases.push("start"), - }; - tracingChannel.subscribe(handlers); + ); try { const client = wrapOpenAI({ @@ -106,10 +111,10 @@ describe("OpenAI API promise wrappers", () => { expect(response.bodyUsed).toBe(false); await new Promise((resolve) => setTimeout(resolve, 0)); } finally { - tracingChannel.unsubscribe(handlers); + remove(); } - expect(phases).toEqual(["start", "asyncEnd"]); + expect(phases).toEqual(["called", "resolved"]); expect(result).toBeUndefined(); }); @@ -141,14 +146,21 @@ describe("OpenAI API promise wrappers", () => { }, ); const phases: string[] = []; - const tracingChannel = - openAIChannels.chatCompletionsCreate.tracingChannel(); - const handlers = { - asyncEnd: () => phases.push("asyncEnd"), - error: () => phases.push("error"), - start: () => phases.push("start"), - }; - tracingChannel.subscribe(handlers); + const remove = openAIChannels.chatCompletionsCreate.intercept( + (target, receiver, args) => { + phases.push("called"); + const promise = Reflect.apply(target, receiver, args); + promise.then( + (value) => { + phases.push("resolved"); + }, + () => { + phases.push("error"); + }, + ); + return promise; + }, + ); try { const client = wrapOpenAI({ @@ -169,14 +181,14 @@ describe("OpenAI API promise wrappers", () => { expect(raw).toBe(response); expect(raw.bodyUsed).toBe(false); expect(pulls).toBe(0); - expect(phases).toEqual(["start", "asyncEnd"]); + expect(phases).toEqual(["called", "resolved"]); const reader = raw.body?.getReader(); await reader?.read(); await reader?.cancel(); expect(cancelled).toBe(true); } finally { - tracingChannel.unsubscribe(handlers); + remove(); } }); diff --git a/js/src/wrappers/ai-sdk/ai-sdk.ts b/js/src/wrappers/ai-sdk/ai-sdk.ts index 7c9859b1c..c0ff3a754 100644 --- a/js/src/wrappers/ai-sdk/ai-sdk.ts +++ b/js/src/wrappers/ai-sdk/ai-sdk.ts @@ -14,9 +14,9 @@ import type { AISDKEmbedFunction, AISDKEmbedParams, AISDKEvaluateParams, + AISDKGenerateFunction, AISDKGenerateImageFunction, AISDKGenerateImageParams, - AISDKGenerateFunction, AISDKHarnessAgentCallParams, AISDKHarnessAgentCreateSessionFunction, AISDKHarnessAgentGenerateFunction, @@ -370,12 +370,11 @@ const wrapHarnessAgentCreateSession = ( const wrapper = function ( params?: Parameters[0], ) { - return harnessAgentChannels.createSession.tracePromise( - () => - params === undefined - ? createSession.call(instance) - : createSession.call(instance, params), - createAISDKChannelContext(params ?? {}, { self: instance }), + return harnessAgentChannels.createSession.invoke( + createSession, + instance, + params === undefined ? [] : [params], + {}, ); }; Object.defineProperty(wrapper, "name", { @@ -517,8 +516,10 @@ const makeGenerateTextWrapper = ( const { span_info, ...params } = allParams; const tracedParams = { ...params }; - return channel.tracePromise( - () => generateText(tracedParams), + return channel.invoke( + generateText, + contextOptions.self, + [tracedParams], createAISDKChannelContext(tracedParams, { aiSDK: contextOptions.aiSDK, denyOutputPaths: options.denyOutputPaths, @@ -585,8 +586,10 @@ const makeGenerateImageWrapper = ( const { span_info, ...params } = allParams; const tracedParams = { ...params }; - return aiSDKChannels.generateImage.tracePromise( - () => generateImage(tracedParams), + return aiSDKChannels.generateImage.invoke( + generateImage, + contextOptions.self, + [tracedParams], createAISDKChannelContext(tracedParams, { aiSDK: contextOptions.aiSDK, denyOutputPaths: options.denyOutputPaths, @@ -620,8 +623,10 @@ const makeEmbedWrapper = ( const { span_info, ...params } = allParams; const tracedParams = { ...params }; - return channel.tracePromise( - () => embed(tracedParams), + return channel.invoke( + embed, + contextOptions.self, + [tracedParams], createAISDKChannelContext(tracedParams, { aiSDK: contextOptions.aiSDK, denyOutputPaths: options.denyOutputPaths, @@ -678,8 +683,10 @@ const makeRerankWrapper = ( const { span_info, ...params } = allParams; const tracedParams = { ...params }; - return aiSDKChannels.rerank.tracePromise( - () => rerank(tracedParams), + return aiSDKChannels.rerank.invoke( + rerank, + contextOptions.self, + [tracedParams], createAISDKChannelContext(tracedParams, { aiSDK: contextOptions.aiSDK, denyOutputPaths: options.denyOutputPaths, @@ -736,7 +743,12 @@ const makeStreamWrapper = ( }), }); - return channel.tracePromise(() => streamText(tracedParams) as any, context); + return channel.invoke( + streamText, + contextOptions.self, + [tracedParams], + context, + ); }; Object.defineProperty(wrapper, "name", { value: name, writable: false }); return wrapper; diff --git a/js/src/wrappers/ai-sdk/harness-agent-context.ts b/js/src/wrappers/ai-sdk/harness-agent-context.ts index 78d60e912..2768f9ab7 100644 --- a/js/src/wrappers/ai-sdk/harness-agent-context.ts +++ b/js/src/wrappers/ai-sdk/harness-agent-context.ts @@ -1,8 +1,9 @@ +import { SpanComponentsV4 } from "../../../util/span_identifier_v4"; +import type { IsoAsyncLocalStorage } from "../../isomorph"; import iso from "../../isomorph"; -import type { IsoAsyncLocalStorage, IsoTracingChannel } from "../../isomorph"; import { - _internalGetGlobalState, _internalExportParentSynchronously, + _internalGetGlobalState, currentSpan, startSpan, updateSpan, @@ -18,7 +19,6 @@ import type { AISDKHarnessAgentCreateSessionParams, AISDKHarnessAgentSession, } from "../../vendor-sdk-types/ai-sdk"; -import { SpanComponentsV4 } from "../../../util/span_identifier_v4"; const BRAINTRUST_TURN_CONTEXT_KEY = "__braintrust_trace_context"; @@ -311,26 +311,17 @@ export function currentHarnessTurnParent(): HarnessTurnParent | undefined { ); } -export function bindHarnessTurnParentToStart( - tracingChannel: IsoTracingChannel, - parentFromEvent: (event: T) => HarnessTurnParent | undefined, -): () => void { - const startChannel = tracingChannel.start; - if (!startChannel) { - return () => {}; - } - +export function runWithHarnessTurnParent( + parent: HarnessTurnParent | undefined, + call: () => T, +): T { harnessTurnParentStore ??= iso.newAsyncLocalStorage< HarnessTurnParent | undefined >(); - const store = harnessTurnParentStore; - startChannel.bindStore( - store, - (event) => parentFromEvent(event) ?? store.getStore(), + return harnessTurnParentStore.run( + parent ?? harnessTurnParentStore.getStore(), + call, ); - return () => { - startChannel.unbindStore(store); - }; } export function startHarnessTurnChildSpan( diff --git a/js/src/wrappers/anthropic.ts b/js/src/wrappers/anthropic.ts index 38b206fd1..5ead70df7 100644 --- a/js/src/wrappers/anthropic.ts +++ b/js/src/wrappers/anthropic.ts @@ -109,9 +109,11 @@ function betaSessionEventsProxy( if (prop === "stream") { return new TypedApplyProxy(target.stream, { apply(stream, thisArg, argArray) { - return anthropicChannels.betaSessionsEventsStream.tracePromise( - () => Reflect.apply(stream, thisArg, argArray), - { arguments: argArray }, + return anthropicChannels.betaSessionsEventsStream.invoke( + stream, + thisArg, + argArray, + {}, ); }, }); @@ -144,9 +146,11 @@ function betaSessionThreadEventsProxy( if (prop === "stream") { return new TypedApplyProxy(target.stream, { apply(stream, thisArg, argArray) { - return anthropicChannels.betaSessionsThreadsEventsStream.tracePromise( - () => Reflect.apply(stream, thisArg, argArray), - { arguments: argArray }, + return anthropicChannels.betaSessionsThreadsEventsStream.invoke( + stream, + thisArg, + argArray, + {}, ); }, }); @@ -208,12 +212,7 @@ function createProxy( ) { return new TypedApplyProxy(create, { apply(target, thisArg, argArray) { - return channel.tracePromise( - () => Reflect.apply(target, thisArg, argArray), - { - arguments: argArray, - }, - ); + return channel.invoke(target, thisArg, argArray, {}); }, }); } @@ -238,12 +237,7 @@ function toolRunnerProxy( }) : { _client: anthropic }; - return channel.traceSync( - () => Reflect.apply(target, invocationTarget, argArray), - { - arguments: argArray, - }, - ); + return channel.invoke(target, invocationTarget, argArray, {}); }, }); } diff --git a/js/src/wrappers/bedrock-runtime.ts b/js/src/wrappers/bedrock-runtime.ts index 13fcd1844..d3072d0df 100644 --- a/js/src/wrappers/bedrock-runtime.ts +++ b/js/src/wrappers/bedrock-runtime.ts @@ -130,15 +130,14 @@ function wrapSend( return send(command, optionsOrCb, cb); } - return bedrockRuntimeChannels.clientSend.tracePromise( - () => + return bedrockRuntimeChannels.clientSend.invoke( + (command, optionsOrCb) => runWithAutoInstrumentationSuppressed(() => send(command, optionsOrCb), ) as Promise, - { - arguments: [command as BedrockRuntimeCommandLike, optionsOrCb], - span_info: buildBedrockRuntimeSpanInfo(command), - }, + undefined, + [command as BedrockRuntimeCommandLike, optionsOrCb], + { span_info: buildBedrockRuntimeSpanInfo(command) }, ); }; } diff --git a/js/src/wrappers/claude-agent-sdk/claude-agent-sdk.ts b/js/src/wrappers/claude-agent-sdk/claude-agent-sdk.ts index ddc6d678e..6f2c09615 100644 --- a/js/src/wrappers/claude-agent-sdk/claude-agent-sdk.ts +++ b/js/src/wrappers/claude-agent-sdk/claude-agent-sdk.ts @@ -13,7 +13,7 @@ type LocalToolMetadata = { /** * Wraps the Claude Agent SDK with Braintrust tracing. Query calls only publish - * tracing-channel events; the Claude Agent SDK plugin owns all span lifecycle + * invocation hooks; the Claude Agent SDK plugin owns all span lifecycle * work, including root/task spans, LLM spans, tool spans, and sub-agent spans. * * @param sdk - The Claude Agent SDK module diff --git a/js/src/wrappers/cloudflare-agent.test.ts b/js/src/wrappers/cloudflare-agent.test.ts index 54805fe8f..62458211d 100644 --- a/js/src/wrappers/cloudflare-agent.test.ts +++ b/js/src/wrappers/cloudflare-agent.test.ts @@ -1,18 +1,18 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -const { tracePromise } = vi.hoisted(() => ({ - tracePromise: vi.fn((fn: () => Promise, _event?: unknown) => fn()), +const { invoke } = vi.hoisted(() => ({ + invoke: vi.fn( + ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), + ), })); - -vi.mock("../isomorph", () => ({ - default: { - getEnv: vi.fn(), - newTracingChannel: vi.fn(() => ({ - subscribe: vi.fn(), - tracePromise, - unsubscribe: vi.fn(), - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { wrapCloudflareAgent } from "./cloudflare-agent"; @@ -41,8 +41,11 @@ describe("wrapCloudflareAgent", () => { receiver: "receiver", status: "completed", }); - expect(tracePromise).toHaveBeenCalledTimes(1); - expect(tracePromise.mock.calls[0][1]).toEqual({ + expect(invoke).toHaveBeenCalledTimes(1); + expect({ + self: invoke.mock.calls[0]?.[1], + arguments: invoke.mock.calls[0]?.[2], + }).toEqual({ arguments: [ChildAgent, options], self: agent, }); @@ -59,7 +62,7 @@ describe("wrapCloudflareAgent", () => { wrapCloudflareAgent(Agent); await new Agent().runAgentTool(); - expect(tracePromise).toHaveBeenCalledTimes(1); + expect(invoke).toHaveBeenCalledTimes(1); }); it("preserves rejections", async () => { @@ -73,7 +76,7 @@ describe("wrapCloudflareAgent", () => { wrapCloudflareAgent(Agent); await expect(new Agent().runAgentTool()).rejects.toBe(rejection); - expect(tracePromise).toHaveBeenCalledTimes(1); + expect(invoke).toHaveBeenCalledTimes(1); }); it("does not trace detached runs", async () => { @@ -96,13 +99,13 @@ describe("wrapCloudflareAgent", () => { options, ); - expect(tracePromise).not.toHaveBeenCalled(); + expect(invoke).not.toHaveBeenCalled(); expect(inputGetter).not.toHaveBeenCalled(); }); it("returns unsupported values unchanged", () => { expect(wrapCloudflareAgent(undefined)).toBeUndefined(); expect(wrapCloudflareAgent(class Unsupported {})).toBeDefined(); - expect(tracePromise).not.toHaveBeenCalled(); + expect(invoke).not.toHaveBeenCalled(); }); }); diff --git a/js/src/wrappers/cloudflare-agent.ts b/js/src/wrappers/cloudflare-agent.ts index c34bb6134..a6c31c7c6 100644 --- a/js/src/wrappers/cloudflare-agent.ts +++ b/js/src/wrappers/cloudflare-agent.ts @@ -54,9 +54,11 @@ export function wrapCloudflareAgent(Agent: T): T { return Reflect.apply(originalRunAgentTool, this, args); } - return cloudflareAgentsChannels.runAgentTool.tracePromise( - () => Reflect.apply(originalRunAgentTool, this, args), - { arguments: args, self: this }, + return cloudflareAgentsChannels.runAgentTool.invoke( + originalRunAgentTool, + this, + args, + {}, ); }, }); diff --git a/js/src/wrappers/cloudflare-ai-chat.test.ts b/js/src/wrappers/cloudflare-ai-chat.test.ts index 123abe70a..bab75e0cc 100644 --- a/js/src/wrappers/cloudflare-ai-chat.test.ts +++ b/js/src/wrappers/cloudflare-ai-chat.test.ts @@ -1,19 +1,18 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -const { tracePromise } = vi.hoisted(() => ({ - tracePromise: vi.fn((fn: () => Promise, _event?: unknown) => fn()), +const { invoke } = vi.hoisted(() => ({ + invoke: vi.fn( + ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), + ), })); - -vi.mock("../isomorph", () => ({ - default: { - getEnv: vi.fn(() => undefined), - newTracingChannel: vi.fn(() => ({ - subscribe: vi.fn(), - tracePromise, - traceSync: vi.fn((fn: () => unknown) => fn()), - unsubscribe: vi.fn(), - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { wrapCloudflareAIChat } from "./cloudflare-ai-chat"; @@ -65,8 +64,11 @@ describe("wrapCloudflareAIChat", () => { expect(agent.onChatResponse()).toBe("field-hook"); expect(module.AIChatAgent.kind).toBe("ai-chat"); expect(module.untouched).toBe("value"); - expect(tracePromise).toHaveBeenCalledTimes(1); - expect(tracePromise.mock.calls[0][1]).toMatchObject({ + expect(invoke).toHaveBeenCalledTimes(1); + expect({ + self: invoke.mock.calls[0]?.[1], + arguments: invoke.mock.calls[0]?.[2], + }).toMatchObject({ arguments: ["request-1", expect.any(Function)], self: agent, }); @@ -93,6 +95,6 @@ describe("wrapCloudflareAIChat", () => { await expect( agent._runExclusiveChatTurn("request-1", async () => {}), ).rejects.toBe(failure); - expect(tracePromise).toHaveBeenCalledTimes(1); + expect(invoke).toHaveBeenCalledTimes(1); }); }); diff --git a/js/src/wrappers/cloudflare-think.test.ts b/js/src/wrappers/cloudflare-think.test.ts index 1db1b0f5d..05b54e46e 100644 --- a/js/src/wrappers/cloudflare-think.test.ts +++ b/js/src/wrappers/cloudflare-think.test.ts @@ -1,17 +1,18 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -const { tracePromise } = vi.hoisted(() => ({ - tracePromise: vi.fn((fn: () => unknown, _event?: unknown) => fn()), +const { invoke } = vi.hoisted(() => ({ + invoke: vi.fn( + ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), + ), })); - -vi.mock("../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(() => ({ - subscribe: vi.fn(), - tracePromise, - unsubscribe: vi.fn(), - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { wrapCloudflareThink } from "./cloudflare-think"; @@ -25,7 +26,7 @@ describe("wrapCloudflareThink", () => { "returns unsupported module %j unchanged", (sdk) => { expect(wrapCloudflareThink(sdk)).toBe(sdk); - expect(tracePromise).not.toHaveBeenCalled(); + expect(invoke).not.toHaveBeenCalled(); }, ); @@ -46,8 +47,11 @@ describe("wrapCloudflareThink", () => { input, marker: "think-instance", }); - expect(tracePromise).toHaveBeenCalledTimes(1); - expect(tracePromise.mock.calls[0]?.[1]).toEqual({ + expect(invoke).toHaveBeenCalledTimes(1); + expect({ + self: invoke.mock.calls[0]?.[1], + arguments: invoke.mock.calls[0]?.[2], + }).toEqual({ arguments: [input], self: instance, }); @@ -65,7 +69,7 @@ describe("wrapCloudflareThink", () => { wrapCloudflareThink(wrapCloudflareThink(sdk)); await new sdk.Think()._runInferenceLoop("hello"); - expect(tracePromise).toHaveBeenCalledTimes(1); + expect(invoke).toHaveBeenCalledTimes(1); }); it("preserves the original method descriptor", () => { diff --git a/js/src/wrappers/cloudflare-think.ts b/js/src/wrappers/cloudflare-think.ts index 9c04400ce..d14da6fa0 100644 --- a/js/src/wrappers/cloudflare-think.ts +++ b/js/src/wrappers/cloudflare-think.ts @@ -58,12 +58,11 @@ function patchThinkClass(Think: CloudflareThinkConstructor): void { input: CloudflareThinkTurnInput, ) { const args = [input] as [CloudflareThinkTurnInput]; - return cloudflareThinkChannels.runInferenceLoop.tracePromise( - () => Reflect.apply(original, this, args), - { - arguments: args, - self: this, - }, + return cloudflareThinkChannels.runInferenceLoop.invoke( + original, + this, + args, + {}, ); }, }); diff --git a/js/src/wrappers/cohere.ts b/js/src/wrappers/cohere.ts index aa6038030..3c49fccc7 100644 --- a/js/src/wrappers/cohere.ts +++ b/js/src/wrappers/cohere.ts @@ -95,9 +95,7 @@ function wrapChat( ) => Promise, ): NonNullable { return (request, options) => - cohereChannels.chat.tracePromise(() => chat(request, options), { - arguments: [request], - } as Parameters[1]); + cohereChannels.chat.invoke(chat, undefined, [request, options], {}); } function wrapChatStream( @@ -107,9 +105,12 @@ function wrapChatStream( ) => Promise, ): NonNullable { return (request, options) => - cohereChannels.chatStream.tracePromise(() => chatStream(request, options), { - arguments: [request], - } as Parameters[1]); + cohereChannels.chatStream.invoke( + chatStream, + undefined, + [request, options], + {}, + ); } function wrapEmbed( @@ -119,9 +120,7 @@ function wrapEmbed( ) => Promise, ): NonNullable { return (request, options) => - cohereChannels.embed.tracePromise(() => embed(request, options), { - arguments: [request], - }); + cohereChannels.embed.invoke(embed, undefined, [request, options], {}); } function wrapRerank( @@ -131,7 +130,5 @@ function wrapRerank( ) => Promise, ): NonNullable { return (request, options) => - cohereChannels.rerank.tracePromise(() => rerank(request, options), { - arguments: [request], - }); + cohereChannels.rerank.invoke(rerank, undefined, [request, options], {}); } diff --git a/js/src/wrappers/cursor-sdk.test.ts b/js/src/wrappers/cursor-sdk.test.ts index c238fba96..8a63363a9 100644 --- a/js/src/wrappers/cursor-sdk.test.ts +++ b/js/src/wrappers/cursor-sdk.test.ts @@ -1,17 +1,18 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -const { tracePromise } = vi.hoisted(() => ({ - tracePromise: vi.fn((fn: () => Promise) => fn()), +const { invoke } = vi.hoisted(() => ({ + invoke: vi.fn( + ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), + ), })); - -vi.mock("../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(() => ({ - subscribe: vi.fn(), - tracePromise, - unsubscribe: vi.fn(), - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { wrapCursorSDK } from "./cursor-sdk"; @@ -56,7 +57,27 @@ describe("wrapCursorSDK", () => { expect(result).toBe(run); expect(agent.send).toHaveBeenCalledWith("hello", expect.any(Object)); - expect(tracePromise).toHaveBeenCalledTimes(2); + expect(invoke).toHaveBeenCalledTimes(2); + }); + + it("does not wrap send twice when the plugin already patched the returned agent", async () => { + const run = makeRun(); + const agent = { + [Symbol.for("braintrust.cursor-sdk.auto-patched-agent")]: true, + send: vi.fn(async () => run), + }; + const sdk = { + Agent: class { + static async create() { + return agent; + } + }, + }; + const wrapped = wrapCursorSDK(sdk as any) as any; + const created = await wrapped.Agent.create({}); + await expect(created.send("hello")).resolves.toBe(run); + expect(agent.send).toHaveBeenCalledExactlyOnceWith("hello"); + expect(invoke).toHaveBeenCalledOnce(); }); it("wraps Agent.resume and preserves private-field-safe method binding", async () => { @@ -101,7 +122,7 @@ describe("wrapCursorSDK", () => { await expect(wrapped.Agent.prompt("hello")).resolves.toMatchObject({ result: "hello", }); - expect(tracePromise).toHaveBeenCalledTimes(1); + expect(invoke).toHaveBeenCalledTimes(1); }); it("handles module namespace-like objects", async () => { diff --git a/js/src/wrappers/cursor-sdk.ts b/js/src/wrappers/cursor-sdk.ts index e20f41fba..a6dc8a36f 100644 --- a/js/src/wrappers/cursor-sdk.ts +++ b/js/src/wrappers/cursor-sdk.ts @@ -12,7 +12,7 @@ const WRAPPED_AGENT = Symbol.for("braintrust.cursor-sdk.wrapped-agent"); /** * Wraps the Cursor TypeScript SDK with Braintrust tracing. The wrapper emits - * diagnostics-channel events; the Cursor SDK plugin owns span lifecycle. + * invocation hooks; the Cursor SDK plugin owns span lifecycle. */ export function wrapCursorSDK(sdk: T): T { if (!sdk || typeof sdk !== "object") { @@ -74,10 +74,8 @@ function wrapCursorAgentClass(Agent: CursorSDKAgentClass): CursorSDKAgentClass { options: CursorSDKAgentOptions, ): Promise { const args = [options] as [CursorSDKAgentOptions]; - return cursorSDKChannels.create.tracePromise( - async () => - wrapCursorAgent(await Reflect.apply(value, target, args)), - { arguments: args } as never, + return wrapCursorAgent( + await cursorSDKChannels.create.invoke(value, target, args, {}), ); }; cache.set(prop, wrapped); @@ -93,10 +91,8 @@ function wrapCursorAgentClass(Agent: CursorSDKAgentClass): CursorSDKAgentClass { string, Partial | undefined, ]; - return cursorSDKChannels.resume.tracePromise( - async () => - wrapCursorAgent(await Reflect.apply(value, target, args)), - { arguments: args } as never, + return wrapCursorAgent( + await cursorSDKChannels.resume.invoke(value, target, args, {}), ); }; cache.set(prop, wrapped); @@ -112,10 +108,7 @@ function wrapCursorAgentClass(Agent: CursorSDKAgentClass): CursorSDKAgentClass { string | CursorSDKUserMessage, CursorSDKAgentOptions | undefined, ]; - return cursorSDKChannels.prompt.tracePromise( - () => Reflect.apply(value, target, args), - { arguments: args } as never, - ); + return cursorSDKChannels.prompt.invoke(value, target, args, {}); }; cache.set(prop, wrapped); return wrapped; @@ -147,7 +140,13 @@ function wrapCursorAgent(agent: CursorSDKAgent): CursorSDKAgent { } const value = Reflect.get(target, prop, receiver); - if (prop === "send" && typeof value === "function") { + if ( + prop === "send" && + typeof value === "function" && + !(target as Record)[ + Symbol.for("braintrust.cursor-sdk.auto-patched-agent") + ] + ) { return function ( message: string | CursorSDKUserMessage, options?: CursorSDKSendOptions, @@ -156,13 +155,11 @@ function wrapCursorAgent(agent: CursorSDKAgent): CursorSDKAgent { string | CursorSDKUserMessage, CursorSDKSendOptions | undefined, ]; - return cursorSDKChannels.send.tracePromise( - () => Reflect.apply(value, target, args), - { - agent: target, - arguments: args, - operation: "send", - } as never, + return cursorSDKChannels.send.invoke( + value as CursorSDKAgent["send"], + target, + args, + { agent: target, operation: "send" }, ); }; } diff --git a/js/src/wrappers/genkit.test.ts b/js/src/wrappers/genkit.test.ts index 5363a64eb..7e612ecbc 100644 --- a/js/src/wrappers/genkit.test.ts +++ b/js/src/wrappers/genkit.test.ts @@ -1,6 +1,6 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import type { IsoChannelHandlers } from "../isomorph"; -import type { ChannelMessage } from "../instrumentation/core/channel-definitions"; + +import type { ChannelMessage } from "../instrumentation/core/tracing-types"; import { genkitChannels } from "../instrumentation/plugins/genkit-channels"; import { configureNode } from "../node/config"; import type { @@ -16,15 +16,9 @@ try { } describe("wrapGenkit", () => { - const tracingChannel = genkitChannels.actionRun.tracingChannel(); - const handlers: IsoChannelHandlers< - ChannelMessage - >[] = []; - + const removals: Array<() => void> = []; afterEach(() => { - for (const handler of handlers.splice(0)) { - tracingChannel.unsubscribe(handler); - } + for (const remove of removals.splice(0)) remove(); vi.restoreAllMocks(); }); @@ -33,18 +27,22 @@ describe("wrapGenkit", () => { phase: "start" | "asyncEnd"; event: ChannelMessage; }> = []; - const handler: IsoChannelHandlers< - ChannelMessage - > = { - asyncEnd: (event) => { - actionRunEvents.push({ event, phase: "asyncEnd" }); - }, - start: (event) => { - actionRunEvents.push({ event, phase: "start" }); - }, - }; - tracingChannel.subscribe(handler); - handlers.push(handler); + removals.push( + genkitChannels.actionRun.intercept( + (target, receiver, args, additional) => { + const event = { ...additional, arguments: args, self: receiver }; + actionRunEvents.push({ event, phase: "start" }); + const result = Reflect.apply(target, receiver, args); + result.then((value) => + actionRunEvents.push({ + event: { ...event, result: value }, + phase: "asyncEnd", + }), + ); + return result; + }, + ), + ); const originalTool = Object.assign( vi.fn(async (input: unknown) => ({ echoed: input })), diff --git a/js/src/wrappers/genkit.ts b/js/src/wrappers/genkit.ts index fc35a8426..112566ff8 100644 --- a/js/src/wrappers/genkit.ts +++ b/js/src/wrappers/genkit.ts @@ -223,47 +223,52 @@ function wrapGenerate( generate: (input: GenkitGenerateInput) => Promise, ): NonNullable { return (input) => - genkitChannels.generate.tracePromise(() => generate(input), { - arguments: [input], - }); + genkitChannels.generate.invoke(generate, undefined, [input], {}); } function wrapGenerateStream( generateStream: (input: GenkitGenerateInput) => GenkitGenerateStreamResponse, ): NonNullable { return (input) => - genkitChannels.generateStream.traceSync(() => generateStream(input), { - arguments: [input], - } as Parameters[1]); + genkitChannels.generateStream.invoke( + generateStream, + undefined, + [input], + {}, + ); } function wrapEmbed( embed: (params: GenkitEmbedParams) => Promise, ): NonNullable { return (params) => - genkitChannels.embed.tracePromise(() => embed(params), { - arguments: [params], - }) as Promise; + genkitChannels.embed.invoke(embed, undefined, [params], {}) as Promise< + GenkitEmbedding[] + >; } function wrapEmbedMany( embedMany: (params: GenkitEmbedManyParams) => Promise, ): NonNullable { return (params) => - genkitChannels.embedMany.tracePromise(() => embedMany(params), { - arguments: [params], - }) as Promise; + genkitChannels.embedMany.invoke( + embedMany, + undefined, + [params], + {}, + ) as Promise; } function wrapRun( run: NonNullable, ): NonNullable { return (name, inputOrFn, maybeFn) => - genkitChannels.actionRun.tracePromise(() => run(name, inputOrFn, maybeFn), { - arguments: [name, inputOrFn, maybeFn], - } as Parameters< - typeof genkitChannels.actionRun.tracePromise - >[1]) as Promise; + genkitChannels.actionRun.invoke( + run, + undefined, + [name, inputOrFn, maybeFn], + {}, + ) as Promise; } function wrapGenkitAction(action: GenkitAction): GenkitAction; @@ -318,12 +323,12 @@ function traceActionRun( run: (input?: unknown, options?: unknown) => Promise, ): (input?: unknown, options?: unknown) => Promise { return (input, options) => - genkitChannels.actionRun.tracePromise(() => run(input, options), { - arguments: [input, options], - self: action, - } as Parameters< - typeof genkitChannels.actionRun.tracePromise - >[1]) as Promise; + genkitChannels.actionRun.invoke( + run, + action, + [input, options], + {}, + ) as Promise; } function traceActionStream( @@ -331,10 +336,7 @@ function traceActionStream( stream: NonNullable, ): NonNullable { return (input, options) => - genkitChannels.actionStream.traceSync(() => stream(input, options), { - arguments: [input, options], - self: action, - } as Parameters[1]); + genkitChannels.actionStream.invoke(stream, action, [input, options], {}); } function hasWrappedFlag(value: object): boolean { diff --git a/js/src/wrappers/github-copilot.ts b/js/src/wrappers/github-copilot.ts index 35e68fa87..ae8813297 100644 --- a/js/src/wrappers/github-copilot.ts +++ b/js/src/wrappers/github-copilot.ts @@ -91,9 +91,11 @@ function wrappedCreateSession( client: GitHubCopilotClient, ): (config: GitHubCopilotSessionConfig) => Promise { return (config: GitHubCopilotSessionConfig) => - gitHubCopilotChannels.createSession.tracePromise( - () => client.createSession(config), - { arguments: [config] }, + gitHubCopilotChannels.createSession.invoke( + client.createSession, + client, + [config], + {}, ); } @@ -104,8 +106,10 @@ function wrappedResumeSession( config: GitHubCopilotResumeSessionConfig, ) => Promise { return (sessionId: string, config: GitHubCopilotResumeSessionConfig) => - gitHubCopilotChannels.resumeSession.tracePromise( - () => client.resumeSession(sessionId, config), - { arguments: [sessionId, config] }, + gitHubCopilotChannels.resumeSession.invoke( + client.resumeSession, + client, + [sessionId, config], + {}, ); } diff --git a/js/src/wrappers/google-adk.test.ts b/js/src/wrappers/google-adk.test.ts index cda99043f..77256f6ba 100644 --- a/js/src/wrappers/google-adk.test.ts +++ b/js/src/wrappers/google-adk.test.ts @@ -1,20 +1,19 @@ import { describe, it, expect, vi, afterEach } from "vitest"; -// Mock iso's newTracingChannel -vi.mock("../isomorph", () => { - const mockTraceSync = vi.fn((fn: () => any) => fn()); - const mockTracePromise = vi.fn((fn: () => any) => fn()); - return { - default: { - newTracingChannel: vi.fn(() => ({ - subscribe: vi.fn(), - unsubscribe: vi.fn(), - traceSync: mockTraceSync, - tracePromise: mockTracePromise, - })), - }, - }; -}); +const { invoke } = vi.hoisted(() => ({ + invoke: vi.fn( + ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), + ), +})); +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), +})); import { wrapGoogleADK } from "./google-adk"; @@ -188,7 +187,7 @@ describe("wrapGoogleADK", () => { expect(events[0].id).toBe("1"); }); - it("should wrap FunctionTool.runAsync to call tracePromise channel", async () => { + it("should wrap FunctionTool.runAsync to call invoke channel", async () => { class FakeFunctionTool { name: string; constructor(config: any) { diff --git a/js/src/wrappers/google-adk.ts b/js/src/wrappers/google-adk.ts index 28af7395e..15a1bfb0d 100644 --- a/js/src/wrappers/google-adk.ts +++ b/js/src/wrappers/google-adk.ts @@ -1,11 +1,11 @@ import { googleADKChannels } from "../instrumentation/plugins/google-adk-channels"; import type { - GoogleADKRunner, - GoogleADKRunnerConstructor, - GoogleADKInMemoryRunnerConstructor, GoogleADKBaseAgent, GoogleADKBaseTool, + GoogleADKInMemoryRunnerConstructor, GoogleADKRunAsyncParams, + GoogleADKRunner, + GoogleADKRunnerConstructor, GoogleADKToolRunRequest, } from "../vendor-sdk-types/google-adk"; @@ -119,10 +119,12 @@ function wrapRunnerRunAsync( ): (params: GoogleADKRunAsyncParams) => AsyncGenerator { const original = runner.runAsync.bind(runner); return function (params: GoogleADKRunAsyncParams) { - return googleADKChannels.runnerRunAsync.traceSync(() => original(params), { - arguments: [params], - self: runner, - } as Parameters[1]); + return googleADKChannels.runnerRunAsync.invoke( + original, + runner, + [params], + {}, + ); }; } @@ -155,11 +157,11 @@ function wrapAgentRunAsync( ): (parentContext: unknown) => AsyncGenerator { const original = agent.runAsync.bind(agent); return function (parentContext: unknown) { - return googleADKChannels.agentRunAsync.traceSync( - () => original(parentContext), - { arguments: [parentContext], self: agent } as Parameters< - typeof googleADKChannels.agentRunAsync.traceSync - >[1], + return googleADKChannels.agentRunAsync.invoke( + original, + agent, + [parentContext], + {}, ); }; } @@ -191,9 +193,6 @@ function wrapToolRunAsync( ): (req: GoogleADKToolRunRequest) => Promise { const original = tool.runAsync.bind(tool); return function (req: GoogleADKToolRunRequest) { - return googleADKChannels.toolRunAsync.tracePromise(() => original(req), { - arguments: [req], - self: tool, - } as Parameters[1]); + return googleADKChannels.toolRunAsync.invoke(original, tool, [req], {}); }; } diff --git a/js/src/wrappers/google-genai.test.ts b/js/src/wrappers/google-genai.test.ts index 0eae7ee46..64964bef7 100644 --- a/js/src/wrappers/google-genai.test.ts +++ b/js/src/wrappers/google-genai.test.ts @@ -1,25 +1,18 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -const { invoke, tracePromise } = vi.hoisted(() => ({ +const { invoke } = vi.hoisted(() => ({ invoke: vi.fn( ( - target: (...args: unknown[]) => unknown, - thisArg: unknown, + target: (...args: any[]) => any, + receiver: unknown, args: unknown[], - ) => Reflect.apply(target, thisArg, args), + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), ), - tracePromise: vi.fn((fn: () => Promise) => fn()), })); - -vi.mock("../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(() => ({ - subscribe: vi.fn(), - invoke, - tracePromise, - unsubscribe: vi.fn(), - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { wrapGoogleGenAI } from "./google-genai"; @@ -152,12 +145,15 @@ describe("wrapGoogleGenAI", () => { expect(result).toEqual({ options, params }); expect(interactionsGetCount).toBe(1); expect(create).toHaveBeenCalledWith(params, options); - expect(tracePromise).toHaveBeenCalledWith(expect.any(Function), { - arguments: [params, options], - }); + expect(invoke).toHaveBeenCalledWith( + expect.any(Function), + undefined, + [params, options], + expect.any(Object), + ); }); - it("does not trace background interaction tasks", async () => { + it("wraps background interaction tasks without making tracing decisions", async () => { const create = vi.fn(async (params: unknown, options?: unknown) => ({ options, params, @@ -191,7 +187,7 @@ describe("wrapGoogleGenAI", () => { expect(result).toEqual({ options, params }); expect(create).toHaveBeenCalledWith(params, options); - expect(tracePromise).not.toHaveBeenCalled(); + expect(invoke).toHaveBeenCalledOnce(); }); it("leaves clients without interactions unchanged", () => { @@ -210,6 +206,6 @@ describe("wrapGoogleGenAI", () => { const client = new wrapped.GoogleGenAI(); expect((client as any).interactions).toBeUndefined(); - expect(tracePromise).not.toHaveBeenCalled(); + expect(invoke).not.toHaveBeenCalled(); }); }); diff --git a/js/src/wrappers/google-genai.ts b/js/src/wrappers/google-genai.ts index 264e8a6bc..b0e3bbe87 100644 --- a/js/src/wrappers/google-genai.ts +++ b/js/src/wrappers/google-genai.ts @@ -3,13 +3,13 @@ import { isObject } from "../util"; import type { GoogleGenAIClient, GoogleGenAIConstructor, - GoogleGenAIEmbedContentParams, GoogleGenAIEditImageParams, + GoogleGenAIEmbedContentParams, GoogleGenAIGenerateContentParams, GoogleGenAIGenerateImagesParams, GoogleGenAIGenerateVideosParams, - GoogleGenAIInteractionCreateParams, GoogleGenAIHttpResponse, + GoogleGenAIInteractionCreateParams, GoogleGenAIInteractions, GoogleGenAIModels, } from "../vendor-sdk-types/google-genai"; @@ -225,11 +225,11 @@ function wrapGenerateContent( original: GoogleGenAIModels["generateContent"], ): GoogleGenAIModels["generateContent"] { return function (params: GoogleGenAIGenerateContentParams) { - return googleGenAIChannels.generateContent.tracePromise( - () => original(params), - { arguments: [params] } as Parameters< - typeof googleGenAIChannels.generateContent.tracePromise - >[1], + return googleGenAIChannels.generateContent.invoke( + original, + undefined, + [params], + {}, ); }; } @@ -238,9 +238,11 @@ function wrapGenerateContentStream( original: GoogleGenAIModels["generateContentStream"], ): GoogleGenAIModels["generateContentStream"] { return function (params: GoogleGenAIGenerateContentParams) { - return googleGenAIChannels.generateContentStream.tracePromise( - () => original(params), - { arguments: [params] }, + return googleGenAIChannels.generateContentStream.invoke( + original, + undefined, + [params], + {}, ); }; } @@ -249,11 +251,11 @@ function wrapEmbedContent( original: GoogleGenAIModels["embedContent"], ): GoogleGenAIModels["embedContent"] { return function (params: GoogleGenAIEmbedContentParams) { - return googleGenAIChannels.embedContent.tracePromise( - () => original(params), - { arguments: [params] } as Parameters< - typeof googleGenAIChannels.embedContent.tracePromise - >[1], + return googleGenAIChannels.embedContent.invoke( + original, + undefined, + [params], + {}, ); }; } @@ -265,22 +267,11 @@ function wrapInteractionCreate( params: GoogleGenAIInteractionCreateParams, options?: Record, ) { - if (params.background === true) { - return options === undefined - ? original(params) - : original(params, options); - } - - const traceContext = - options === undefined - ? { arguments: [params] } - : { arguments: [params, options] }; - return googleGenAIChannels.interactionsCreate.tracePromise( - () => - options === undefined ? original(params) : original(params, options), - traceContext as Parameters< - typeof googleGenAIChannels.interactionsCreate.tracePromise - >[1], + return googleGenAIChannels.interactionsCreate.invoke( + original, + undefined, + options === undefined ? [params] : [params, options], + {}, ); }; } diff --git a/js/src/wrappers/groq.ts b/js/src/wrappers/groq.ts index 7b4b7a192..de1a293fd 100644 --- a/js/src/wrappers/groq.ts +++ b/js/src/wrappers/groq.ts @@ -211,9 +211,11 @@ function wrapChatCompletionsCreate( ) => Promise, ): GroqChat["completions"]["create"] { return (request, options) => - groqChannels.chatCompletionsCreate.tracePromise( - () => create(request, options), - { arguments: [request, options] }, + groqChannels.chatCompletionsCreate.invoke( + create, + undefined, + [request, options], + {}, ) as ReturnType; } @@ -224,18 +226,23 @@ function wrapEmbeddingsCreate( ) => Promise, ): GroqEmbeddings["create"] { return (request, options) => - groqChannels.embeddingsCreate.tracePromise(() => create(request, options), { - arguments: [request, options], - }) as ReturnType; + groqChannels.embeddingsCreate.invoke( + create, + undefined, + [request, options], + {}, + ) as ReturnType; } function wrapAudioSpeechCreate( create: GroqAudioSpeech["create"], ): GroqAudioSpeech["create"] { return (request, options) => - groqChannels.audioSpeechCreate.tracePromise( - () => create(request, options), - { arguments: [request, options] }, + groqChannels.audioSpeechCreate.invoke( + create, + undefined, + [request, options], + {}, ) as ReturnType; } @@ -243,9 +250,11 @@ function wrapAudioTranscriptionsCreate( create: GroqAudioTranscriptions["create"], ): GroqAudioTranscriptions["create"] { return (request, options) => - groqChannels.audioTranscriptionsCreate.tracePromise( - () => create(request, options), - { arguments: [request, options] }, + groqChannels.audioTranscriptionsCreate.invoke( + create, + undefined, + [request, options], + {}, ) as ReturnType; } @@ -253,8 +262,10 @@ function wrapAudioTranslationsCreate( create: GroqAudioTranslations["create"], ): GroqAudioTranslations["create"] { return (request, options) => - groqChannels.audioTranslationsCreate.tracePromise( - () => create(request, options), - { arguments: [request, options] }, + groqChannels.audioTranslationsCreate.invoke( + create, + undefined, + [request, options], + {}, ) as ReturnType; } diff --git a/js/src/wrappers/huggingface-transformers.ts b/js/src/wrappers/huggingface-transformers.ts index 3965f1234..93f1f7076 100644 --- a/js/src/wrappers/huggingface-transformers.ts +++ b/js/src/wrappers/huggingface-transformers.ts @@ -101,12 +101,10 @@ function wrapPipelineFactory( ) { const [task] = args; const context: Parameters< - typeof huggingFaceTransformersChannels.pipeline.tracePromise - >[1] = { - arguments: args, - }; + typeof huggingFaceTransformersChannels.pipeline.invoke + >[3] = {}; return huggingFaceTransformersChannels.pipeline - .tracePromise(() => Reflect.apply(factory, this, args), context) + .invoke(factory, this, args, context) .then((pipeline) => { if (isSupportedHuggingFaceTransformersTask(pipeline.task ?? task)) { return wrapPipeline(pipeline); @@ -150,13 +148,14 @@ function wrapPipeline( const proxy = new Proxy(pipeline, { apply(target, thisArg, args) { const context: Parameters< - typeof huggingFaceTransformersChannels.pipelineCall.tracePromise - >[1] = { - arguments: args as [unknown, ...unknown[]], - self: target, + typeof huggingFaceTransformersChannels.pipelineCall.invoke + >[3] = { + pipeline: target, }; - return huggingFaceTransformersChannels.pipelineCall.tracePromise( - () => Reflect.apply(target, thisArg, args), + return huggingFaceTransformersChannels.pipelineCall.invoke( + target, + thisArg, + args, context, ); }, diff --git a/js/src/wrappers/huggingface.ts b/js/src/wrappers/huggingface.ts index 6ff98a259..8af9723f0 100644 --- a/js/src/wrappers/huggingface.ts +++ b/js/src/wrappers/huggingface.ts @@ -1,5 +1,5 @@ -import { huggingFaceChannels } from "../instrumentation/plugins/huggingface-channels"; import { isObject } from "../../util"; +import { huggingFaceChannels } from "../instrumentation/plugins/huggingface-channels"; import type { HuggingFaceChatCompletion, HuggingFaceChatCompletionChunk, @@ -202,20 +202,6 @@ function clientProxyWithContext( }); } -function withEndpointUrl>( - params: T, - endpointUrl?: string, -): T { - if (!endpointUrl || params.endpointUrl !== undefined) { - return params; - } - - return { - ...params, - endpointUrl, - }; -} - function wrapChatCompletion( original: ( params: HuggingFaceChatCompletionParams, @@ -224,14 +210,15 @@ function wrapChatCompletion( endpointUrl?: string, ): HuggingFaceClient["chatCompletion"] { return (params, options) => { - const traceParams = withEndpointUrl(params, endpointUrl); const context: Parameters< - typeof huggingFaceChannels.chatCompletion.tracePromise - >[1] = { - arguments: [traceParams], + typeof huggingFaceChannels.chatCompletion.invoke + >[3] = { + endpointUrl, }; - return huggingFaceChannels.chatCompletion.tracePromise( - () => original(params, options), + return huggingFaceChannels.chatCompletion.invoke( + original, + undefined, + [params, options], context, ); }; @@ -245,11 +232,11 @@ function wrapChatCompletionStream( endpointUrl?: string, ): HuggingFaceClient["chatCompletionStream"] { return (params, options) => - huggingFaceChannels.chatCompletionStream.traceSync( - () => original(params, options), - { - arguments: [withEndpointUrl(params, endpointUrl)], - }, + huggingFaceChannels.chatCompletionStream.invoke( + original, + undefined, + [params, options], + { endpointUrl }, ); } @@ -261,14 +248,15 @@ function wrapTextGeneration( endpointUrl?: string, ): HuggingFaceClient["textGeneration"] { return (params, options) => { - const traceParams = withEndpointUrl(params, endpointUrl); const context: Parameters< - typeof huggingFaceChannels.textGeneration.tracePromise - >[1] = { - arguments: [traceParams], + typeof huggingFaceChannels.textGeneration.invoke + >[3] = { + endpointUrl, }; - return huggingFaceChannels.textGeneration.tracePromise( - () => original(params, options), + return huggingFaceChannels.textGeneration.invoke( + original, + undefined, + [params, options], context, ); }; @@ -282,11 +270,11 @@ function wrapTextGenerationStream( endpointUrl?: string, ): HuggingFaceClient["textGenerationStream"] { return (params, options) => - huggingFaceChannels.textGenerationStream.traceSync( - () => original(params, options), - { - arguments: [withEndpointUrl(params, endpointUrl)], - }, + huggingFaceChannels.textGenerationStream.invoke( + original, + undefined, + [params, options], + { endpointUrl }, ); } @@ -298,14 +286,15 @@ function wrapFeatureExtraction( endpointUrl?: string, ): HuggingFaceClient["featureExtraction"] { return (params, options) => { - const traceParams = withEndpointUrl(params, endpointUrl); const context: Parameters< - typeof huggingFaceChannels.featureExtraction.tracePromise - >[1] = { - arguments: [traceParams], + typeof huggingFaceChannels.featureExtraction.invoke + >[3] = { + endpointUrl, }; - return huggingFaceChannels.featureExtraction.tracePromise( - () => original(params, options), + return huggingFaceChannels.featureExtraction.invoke( + original, + undefined, + [params, options], context, ); }; diff --git a/js/src/wrappers/langsmith.test.ts b/js/src/wrappers/langsmith.test.ts index 5c2588696..9d0be8c57 100644 --- a/js/src/wrappers/langsmith.test.ts +++ b/js/src/wrappers/langsmith.test.ts @@ -1,17 +1,18 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -const { tracePromise } = vi.hoisted(() => ({ - tracePromise: vi.fn((fn: () => Promise, _event?: unknown) => fn()), +const { invoke } = vi.hoisted(() => ({ + invoke: vi.fn( + ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), + ), })); - -vi.mock("../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(() => ({ - subscribe: vi.fn(), - tracePromise, - unsubscribe: vi.fn(), - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { @@ -59,10 +60,8 @@ describe("LangSmith namespace wrappers", () => { await expect(traced("hello")).resolves.toBe("hello!"); expect(existingOnEnd).toHaveBeenCalledWith(run); expect(wrapped.helper).toBe(helper); - expect(tracePromise).toHaveBeenCalledTimes(1); - expect(tracePromise.mock.calls[0]?.[1]).toMatchObject({ - arguments: ["run-1", run], - }); + expect(invoke).toHaveBeenCalledTimes(1); + expect(invoke.mock.calls[0]?.[2]).toEqual(["run-1", run]); }); it("recursively wraps RunTree children without changing class identity", async () => { @@ -109,7 +108,7 @@ describe("LangSmith namespace wrappers", () => { tree.postRun = replacement; await expect(tree.postRun()).resolves.toBe("replacement"); expect(replacement).toHaveBeenCalledOnce(); - expect(tracePromise).toHaveBeenCalledTimes(3); + expect(invoke).toHaveBeenCalledTimes(3); }); it("wraps Client lifecycle methods and safely binds other methods", async () => { @@ -148,7 +147,7 @@ describe("LangSmith namespace wrappers", () => { id: "replacement", }); expect(replacement).toHaveBeenCalledOnce(); - expect(tracePromise).toHaveBeenCalledTimes(4); + expect(invoke).toHaveBeenCalledTimes(4); }); it("preserves Client method errors", async () => { @@ -172,7 +171,7 @@ describe("LangSmith namespace wrappers", () => { const wrapped = wrapLangSmithClient(wrapLangSmithClient({ Client })); await new wrapped.Client().createRun({ id: "one" }); - expect(tracePromise).toHaveBeenCalledTimes(1); + expect(invoke).toHaveBeenCalledTimes(1); }); it("preserves namespace keys for module-shaped objects", () => { diff --git a/js/src/wrappers/langsmith.ts b/js/src/wrappers/langsmith.ts index a8486c573..dcdad62ef 100644 --- a/js/src/wrappers/langsmith.ts +++ b/js/src/wrappers/langsmith.ts @@ -220,22 +220,15 @@ function wrapRunTreeInstance(runTree: LangSmithRunTree): LangSmithRunTree { } else if (prop === "postRun") { const method = value as (...args: unknown[]) => Promise; wrapped = (...args: unknown[]) => - langSmithChannels.createRun.tracePromise( - () => Reflect.apply(method, target, args), - { arguments: [target] }, - ); + langSmithChannels.createRun.invoke(method, target, args, { + runTree: target, + }); } else if (prop === "patchRun") { const method = value as (...args: unknown[]) => Promise; wrapped = (...args: unknown[]) => - langSmithChannels.updateRun.tracePromise( - () => Reflect.apply(method, target, args), - { - arguments: [ - typeof target.id === "string" ? target.id : "", - target, - ], - }, - ); + langSmithChannels.updateRun.invoke(method, target, args, { + runTree: target, + }); } else { wrapped = value.bind(target); } @@ -287,29 +280,17 @@ function wrapClientInstance(client: LangSmithClient): LangSmithClient { const method = value as (...args: unknown[]) => Promise; wrapped = ( ...args: Parameters> - ) => - langSmithChannels.createRun.tracePromise( - () => Reflect.apply(method, target, args), - { arguments: args }, - ); + ) => langSmithChannels.createRun.invoke(method, target, args, {}); } else if (prop === "updateRun") { const method = value as (...args: unknown[]) => Promise; wrapped = ( ...args: Parameters> - ) => - langSmithChannels.updateRun.tracePromise( - () => Reflect.apply(method, target, args), - { arguments: args }, - ); + ) => langSmithChannels.updateRun.invoke(method, target, args, {}); } else if (prop === "batchIngestRuns") { const method = value as (...args: unknown[]) => Promise; wrapped = ( ...args: Parameters> - ) => - langSmithChannels.batchIngestRuns.tracePromise( - () => Reflect.apply(method, target, args), - { arguments: args }, - ); + ) => langSmithChannels.batchIngestRuns.invoke(method, target, args, {}); } else { wrapped = value.bind(target); } @@ -325,9 +306,12 @@ function publishRunUpdate(runTree: LangSmithRunTree | undefined): void { try { void langSmithChannels.updateRun - .tracePromise(() => Promise.resolve(undefined), { - arguments: [runTree.id, runTree], - }) + .invoke( + () => Promise.resolve(undefined), + undefined, + [runTree.id, runTree], + {}, + ) .catch((error) => { debugLogger.error("LangSmith traceable instrumentation failed:", error); }); diff --git a/js/src/wrappers/mistral.ts b/js/src/wrappers/mistral.ts index c449997f7..7efa02f28 100644 --- a/js/src/wrappers/mistral.ts +++ b/js/src/wrappers/mistral.ts @@ -198,11 +198,11 @@ function wrapChatComplete( ) => Promise, ): MistralChat["complete"] { return (request, options) => - mistralChannels.chatComplete.tracePromise( - () => complete(request, options), - { - arguments: [request], - } as Parameters[1], + mistralChannels.chatComplete.invoke( + complete, + undefined, + [request, options], + {}, ); } @@ -213,9 +213,12 @@ function wrapChatStream( ) => Promise, ): MistralChat["stream"] { return (request, options) => - mistralChannels.chatStream.tracePromise(() => stream(request, options), { - arguments: [request], - } as Parameters[1]); + mistralChannels.chatStream.invoke( + stream, + undefined, + [request, options], + {}, + ); } function wrapEmbeddingsCreate( @@ -225,9 +228,11 @@ function wrapEmbeddingsCreate( ) => Promise, ): MistralEmbeddings["create"] { return (request, options) => - mistralChannels.embeddingsCreate.tracePromise( - () => create(request, options), - { arguments: [request] }, + mistralChannels.embeddingsCreate.invoke( + create, + undefined, + [request, options], + {}, ); } @@ -238,9 +243,11 @@ function wrapClassifiersModerate( ) => Promise, ): MistralClassifiers["moderate"] { return (request, options) => - mistralChannels.classifiersModerate.tracePromise( - () => moderate(request, options), - { arguments: [request] }, + mistralChannels.classifiersModerate.invoke( + moderate, + undefined, + [request, options], + {}, ); } @@ -251,9 +258,11 @@ function wrapClassifiersModerateChat( ) => Promise, ): MistralClassifiers["moderateChat"] { return (request, options) => - mistralChannels.classifiersModerateChat.tracePromise( - () => moderateChat(request, options), - { arguments: [request] }, + mistralChannels.classifiersModerateChat.invoke( + moderateChat, + undefined, + [request, options], + {}, ); } @@ -264,9 +273,11 @@ function wrapClassifiersClassify( ) => Promise, ): NonNullable { return (request, options) => - mistralChannels.classifiersClassify.tracePromise( - () => classify(request, options), - { arguments: [request] }, + mistralChannels.classifiersClassify.invoke( + classify, + undefined, + [request, options], + {}, ); } @@ -277,9 +288,11 @@ function wrapClassifiersClassifyChat( ) => Promise, ): NonNullable { return (request, options) => - mistralChannels.classifiersClassifyChat.tracePromise( - () => classifyChat(request, options), - { arguments: [request] }, + mistralChannels.classifiersClassifyChat.invoke( + classifyChat, + undefined, + [request, options], + {}, ); } @@ -290,9 +303,12 @@ function wrapFimComplete( ) => Promise, ): MistralFim["complete"] { return (request, options) => - mistralChannels.fimComplete.tracePromise(() => complete(request, options), { - arguments: [request], - } as Parameters[1]); + mistralChannels.fimComplete.invoke( + complete, + undefined, + [request, options], + {}, + ); } function wrapFimStream( @@ -302,9 +318,7 @@ function wrapFimStream( ) => Promise, ): MistralFim["stream"] { return (request, options) => - mistralChannels.fimStream.tracePromise(() => stream(request, options), { - arguments: [request], - } as Parameters[1]); + mistralChannels.fimStream.invoke(stream, undefined, [request, options], {}); } function wrapAgentsComplete( @@ -314,11 +328,11 @@ function wrapAgentsComplete( ) => Promise, ): MistralAgents["complete"] { return (request, options) => - mistralChannels.agentsComplete.tracePromise( - () => complete(request, options), - { - arguments: [request], - } as Parameters[1], + mistralChannels.agentsComplete.invoke( + complete, + undefined, + [request, options], + {}, ); } @@ -329,7 +343,10 @@ function wrapAgentsStream( ) => Promise, ): MistralAgents["stream"] { return (request, options) => - mistralChannels.agentsStream.tracePromise(() => stream(request, options), { - arguments: [request], - } as Parameters[1]); + mistralChannels.agentsStream.invoke( + stream, + undefined, + [request, options], + {}, + ); } diff --git a/js/src/wrappers/oai.ts b/js/src/wrappers/oai.ts index de8cd552b..0f40762d1 100644 --- a/js/src/wrappers/oai.ts +++ b/js/src/wrappers/oai.ts @@ -4,18 +4,15 @@ import type { OpenAIMediaParams, } from "../vendor-sdk-types/openai-media"; /* eslint-disable @typescript-eslint/no-explicit-any */ +import type { ArgsOf } from "../instrumentation/core/channel-definitions"; +import type { ResultOf } from "../instrumentation/core/tracing-types"; +import { openAIChannels } from "../instrumentation/plugins/openai-channels"; import type { CompiledPrompt } from "../logger"; import { LEGACY_CACHED_HEADER, parseCachedHeader, X_CACHED_HEADER, } from "../openai-utils"; -import { responsesProxy } from "./oai_responses"; -import type { - ArgsOf, - ResultOf, -} from "../instrumentation/core/channel-definitions"; -import { openAIChannels } from "../instrumentation/plugins/openai-channels"; import type { OpenAIChatCompletion, OpenAIChatCreateParams, @@ -26,16 +23,17 @@ import type { OpenAIModerationCreateParams, OpenAIModerationResponse, } from "../vendor-sdk-types/openai"; +import { OpenAIV4Client } from "../vendor-sdk-types/openai-v4"; +import { responsesProxy } from "./oai_responses"; import { APIPromise, createChannelContext, createLazyAPIPromise, EnhancedResponse, splitSpanInfo, - tracePromiseAsResponse, - tracePromiseWithResponse, + invokeAsResponse, + invokeWithResponse, } from "./openai-promise-utils"; -import { OpenAIV4Client } from "../vendor-sdk-types/openai-v4"; declare global { var __inherited_braintrust_wrap_openai: ((openai: any) => any) | undefined; @@ -280,9 +278,11 @@ function wrapBetaChatCompletionParse< const { span_info, params } = splitSpanInfo( allParams, ); - return openAIChannels.betaChatCompletionsParse.tracePromise( - async () => await completion(params), - { arguments: [params], span_info }, + return openAIChannels.betaChatCompletionsParse.invoke( + completion, + undefined, + [params], + { span_info }, ); }; } @@ -294,9 +294,11 @@ function wrapBetaChatCompletionStream

( const { span_info, params } = splitSpanInfo( allParams, ); - return openAIChannels.betaChatCompletionsStream.traceSync( - () => completion(params), - { arguments: [params], span_info }, + return openAIChannels.betaChatCompletionsStream.invoke( + completion, + undefined, + [params], + { span_info }, ); }; } @@ -337,12 +339,11 @@ function wrapChatCompletion< const completionPromise = // eslint-disable-next-line @typescript-eslint/consistent-type-assertions getAPIPromise() as APIPromise; - const { data, response, request_id } = - await tracePromiseWithResponse( - openAIChannels.chatCompletionsCreate, - traceContext, - completionPromise, - ); + const { data, response, request_id } = await invokeWithResponse( + openAIChannels.chatCompletionsCreate, + traceContext, + completionPromise, + ); // eslint-disable-next-line @typescript-eslint/consistent-type-assertions return { data: data as C, response, request_id }; } @@ -350,7 +351,7 @@ function wrapChatCompletion< const completionResponse = // eslint-disable-next-line @typescript-eslint/consistent-type-assertions getAPIPromise() as APIPromise; - const { data, response, request_id } = await tracePromiseWithResponse( + const { data, response, request_id } = await invokeWithResponse( openAIChannels.chatCompletionsCreate, traceContext, completionResponse, @@ -365,7 +366,7 @@ function wrapChatCompletion< return createLazyAPIPromise( ensureExecuted, () => - tracePromiseAsResponse( + invokeAsResponse( openAIChannels.chatCompletionsCreate, createChannelContext( openAIChannels.chatCompletionsCreate, @@ -427,11 +428,7 @@ function wrapApiCreateWithChannel< if (!executionPromise) { executionPromise = (async () => { const traceContext = createChannelContext(channel, params, span_info); - return tracePromiseWithResponse( - channel, - traceContext, - getAPIPromise(), - ); + return invokeWithResponse(channel, traceContext, getAPIPromise()); })(); } return executionPromise; @@ -439,7 +436,7 @@ function wrapApiCreateWithChannel< return createLazyAPIPromise( ensureExecuted, () => - tracePromiseAsResponse( + invokeAsResponse( channel, createChannelContext(channel, params, span_info), getAPIPromise(), diff --git a/js/src/wrappers/oai_responses.ts b/js/src/wrappers/oai_responses.ts index 466856286..d6b1dc7a6 100644 --- a/js/src/wrappers/oai_responses.ts +++ b/js/src/wrappers/oai_responses.ts @@ -1,7 +1,5 @@ -import type { - ArgsOf, - ResultOf, -} from "../instrumentation/core/channel-definitions"; +import type { ArgsOf } from "../instrumentation/core/channel-definitions"; +import type { ResultOf } from "../instrumentation/core/tracing-types"; import type { ChannelSpanInfo } from "../instrumentation/core/types"; import { openAIChannels } from "../instrumentation/plugins/openai-channels"; import { parseMetricsFromUsage } from "../openai-utils"; @@ -11,8 +9,8 @@ import { createLazyAPIPromise, EnhancedResponse, splitSpanInfo, - tracePromiseAsResponse, - tracePromiseWithResponse, + invokeAsResponse, + invokeWithResponse, } from "./openai-promise-utils"; type SpanInfo = { @@ -92,11 +90,7 @@ function wrapResponsesAsync< if (!executionPromise) { executionPromise = (async () => { const traceContext = createChannelContext(channel, params, span_info); - return tracePromiseWithResponse( - channel, - traceContext, - getAPIPromise(), - ); + return invokeWithResponse(channel, traceContext, getAPIPromise()); })(); } @@ -106,7 +100,7 @@ function wrapResponsesAsync< return createLazyAPIPromise( ensureExecuted, () => - tracePromiseAsResponse( + invokeAsResponse( channel, createChannelContext(channel, params, span_info), getAPIPromise(), @@ -134,10 +128,7 @@ function wrapResponsesSyncStream( ArgsOf[0], SpanInfo["span_info"] >(allParams); - return channel.traceSync(() => target(params, options), { - arguments: [params], - span_info, - }); + return channel.invoke(target, undefined, [params, options], { span_info }); }; } diff --git a/js/src/wrappers/ollama.test.ts b/js/src/wrappers/ollama.test.ts index 9a8ebff46..92c0dc624 100644 --- a/js/src/wrappers/ollama.test.ts +++ b/js/src/wrappers/ollama.test.ts @@ -5,7 +5,7 @@ import type { OllamaClient } from "../vendor-sdk-types/ollama"; import { wrapOllama } from "./ollama"; describe("wrapOllama", () => { - it("emits channel events for every supported generation surface", async () => { + it("invokes wrapping hooks for every supported generation surface", async () => { const client: OllamaClient = { chat: vi.fn(async () => ({ message: { role: "assistant", content: "OK" }, @@ -15,14 +15,20 @@ describe("wrapOllama", () => { embed: vi.fn(async () => ({ embeddings: [[0.1, 0.2]] })), }; const chatSpy = vi - .spyOn(ollamaChannels.chat, "tracePromise") - .mockImplementation((fn) => fn()); + .spyOn(ollamaChannels.chat, "invoke") + .mockImplementation((fn, receiver, args) => + Reflect.apply(fn, receiver, args), + ); const generateSpy = vi - .spyOn(ollamaChannels.generate, "tracePromise") - .mockImplementation((fn) => fn()); + .spyOn(ollamaChannels.generate, "invoke") + .mockImplementation((fn, receiver, args) => + Reflect.apply(fn, receiver, args), + ); const embedSpy = vi - .spyOn(ollamaChannels.embed, "tracePromise") - .mockImplementation((fn) => fn()); + .spyOn(ollamaChannels.embed, "invoke") + .mockImplementation((fn, receiver, args) => + Reflect.apply(fn, receiver, args), + ); const wrapped = wrapOllama(client); expect(wrapped.chat).toBe(wrapped.chat); expect(wrapped.generate).toBe(wrapped.generate); @@ -60,9 +66,11 @@ describe("wrapOllama", () => { done: true, })); const client: OllamaClient = { chat: originalChat }; - const tracePromise = vi - .spyOn(ollamaChannels.chat, "tracePromise") - .mockImplementation((fn) => fn()); + const invoke = vi + .spyOn(ollamaChannels.chat, "invoke") + .mockImplementation((fn, receiver, args) => + Reflect.apply(fn, receiver, args), + ); const wrapped = wrapOllama(client); const firstWrappedChat = wrapped.chat; @@ -82,6 +90,6 @@ describe("wrapOllama", () => { client.chat = undefined; expect(wrapped.chat).toBeUndefined(); - tracePromise.mockRestore(); + invoke.mockRestore(); }); }); diff --git a/js/src/wrappers/ollama.ts b/js/src/wrappers/ollama.ts index 3f90067a9..ee11e6100 100644 --- a/js/src/wrappers/ollama.ts +++ b/js/src/wrappers/ollama.ts @@ -1,6 +1,6 @@ +import { isObject } from "../../util"; import { debugLogger } from "../debug-logger"; import { ollamaChannels } from "../instrumentation/plugins/ollama-channels"; -import { isObject } from "../../util"; import type { OllamaChatRequest, OllamaChatResult, @@ -98,25 +98,19 @@ function wrapChat( chat: (request: OllamaChatRequest) => Promise, ): NonNullable { return (request) => - ollamaChannels.chat.tracePromise(() => chat(request), { - arguments: [request], - }); + ollamaChannels.chat.invoke(chat, undefined, [request], {}); } function wrapGenerate( generate: (request: OllamaGenerateRequest) => Promise, ): NonNullable { return (request) => - ollamaChannels.generate.tracePromise(() => generate(request), { - arguments: [request], - }); + ollamaChannels.generate.invoke(generate, undefined, [request], {}); } function wrapEmbed( embed: (request: OllamaEmbedRequest) => Promise, ): NonNullable { return (request) => - ollamaChannels.embed.tracePromise(() => embed(request), { - arguments: [request], - }); + ollamaChannels.embed.invoke(embed, undefined, [request], {}); } diff --git a/js/src/wrappers/openai-codex.ts b/js/src/wrappers/openai-codex.ts index 8b5a435b3..f6e2d17f8 100644 --- a/js/src/wrappers/openai-codex.ts +++ b/js/src/wrappers/openai-codex.ts @@ -145,14 +145,10 @@ function wrapCodexThread(thread: OpenAICodexThread): OpenAICodexThread { OpenAICodexInput, OpenAICodexTurnOptions | undefined, ]; - return openAICodexChannels.run.tracePromise( - () => Reflect.apply(value, target, args), - { - arguments: args, - operation: "run", - thread: target, - }, - ); + return openAICodexChannels.run.invoke(value, target, args, { + operation: "run", + thread: target, + }); }; } if (prop === "runStreamed" && typeof value === "function") { @@ -164,14 +160,10 @@ function wrapCodexThread(thread: OpenAICodexThread): OpenAICodexThread { OpenAICodexInput, OpenAICodexTurnOptions | undefined, ]; - return openAICodexChannels.runStreamed.tracePromise( - () => Reflect.apply(value, target, args), - { - arguments: args, - operation: "runStreamed", - thread: target, - }, - ); + return openAICodexChannels.runStreamed.invoke(value, target, args, { + operation: "runStreamed", + thread: target, + }); }; } if (typeof value === "function") { diff --git a/js/src/wrappers/openai-promise-utils.ts b/js/src/wrappers/openai-promise-utils.ts index 9ba00d0bd..ac477d33e 100644 --- a/js/src/wrappers/openai-promise-utils.ts +++ b/js/src/wrappers/openai-promise-utils.ts @@ -1,11 +1,11 @@ import type { ArgsOf, - ResultOf, + InvocationAdditionalOf, } from "../instrumentation/core/channel-definitions"; +import type { ResultOf } from "../instrumentation/core/tracing-types"; import type { OpenAIAsyncChannel, OpenAIChannel, - OpenAIStartContext, } from "../instrumentation/plugins/openai-channels"; export type EnhancedResponse = { @@ -20,7 +20,7 @@ export interface APIPromise extends Promise { } type ChannelContext = - OpenAIStartContext; + InvocationAdditionalOf & { arguments: ArgsOf }; type ChannelParam = ArgsOf[0]; @@ -44,10 +44,11 @@ export function createChannelContext( // eslint-disable-next-line @typescript-eslint/consistent-type-assertions [params] as ArgsOf, span_info, + responseInfo: {}, } as ChannelContext; } -export async function tracePromiseWithResponse< +export async function invokeWithResponse< TChannel extends OpenAIAsyncChannel, TResult extends ResultOf, >( @@ -56,18 +57,23 @@ export async function tracePromiseWithResponse< apiPromise: APIPromise, ): Promise> { let enhancedResponse: EnhancedResponse | undefined; - const tracePromise = - // eslint-disable-next-line @typescript-eslint/consistent-type-assertions - channel.tracePromise as unknown as >( - fn: () => TReturn, - context: ChannelContext, - ) => TReturn; - - const data = await tracePromise(async () => { - enhancedResponse = await apiPromise.withResponse(); - traceContext.response = enhancedResponse.response; - return enhancedResponse.data; - }, traceContext); + const invoke = channel.invoke as ( + call: () => T, + receiver: undefined, + args: unknown[], + additional: ChannelContext, + ) => T; + + const data = await invoke( + async () => { + enhancedResponse = await apiPromise.withResponse(); + traceContext.responseInfo!.response = enhancedResponse.response; + return enhancedResponse.data; + }, + undefined, + traceContext.arguments, + traceContext, + ); if (!enhancedResponse) { throw new Error("Expected withResponse() to provide response"); @@ -80,7 +86,7 @@ export async function tracePromiseWithResponse< }; } -export async function tracePromiseAsResponse< +export async function invokeAsResponse< TChannel extends OpenAIAsyncChannel, TResult extends ResultOf, >( @@ -88,19 +94,24 @@ export async function tracePromiseAsResponse< traceContext: ChannelContext, apiPromise: APIPromise, ): Promise { - const tracePromise = - // eslint-disable-next-line @typescript-eslint/consistent-type-assertions - channel.tracePromise as unknown as ( - fn: () => Promise, - context: ChannelContext, - ) => Promise; + const invoke = channel.invoke as ( + call: () => T, + receiver: undefined, + args: unknown[], + additional: ChannelContext, + ) => T; let response: Response | undefined; - await tracePromise(async () => { - response = await apiPromise.asResponse(); - traceContext.response = response; - return undefined; - }, traceContext); + await invoke( + async () => { + response = await apiPromise.asResponse(); + traceContext.responseInfo!.response = response; + return undefined; + }, + undefined, + traceContext.arguments, + traceContext, + ); if (!response) { throw new Error("Expected asResponse() to provide response"); diff --git a/js/src/wrappers/openrouter-agent.test.ts b/js/src/wrappers/openrouter-agent.test.ts index 02f01e040..46c030c37 100644 --- a/js/src/wrappers/openrouter-agent.test.ts +++ b/js/src/wrappers/openrouter-agent.test.ts @@ -17,8 +17,8 @@ describe("wrapOpenRouterAgent", () => { ); }); - it("emits callModel tracing events and clones the request", () => { - const traceSpy = vi.spyOn(openRouterAgentChannels.callModel, "traceSync"); + it("invokes the callModel hook and clones the request", () => { + const traceSpy = vi.spyOn(openRouterAgentChannels.callModel, "invoke"); const sdk = { name: "agent-sdk", callModel(request: Record, options?: unknown) { @@ -35,14 +35,12 @@ describe("wrapOpenRouterAgent", () => { const result = wrapped.callModel(request); expect(traceSpy).toHaveBeenCalledTimes(1); - const traceContext = traceSpy.mock.calls[0]?.[1] as { - arguments: unknown[]; - }; - expect(traceContext.arguments[0]).toMatchObject(request); - expect(traceContext.arguments[0]).not.toBe(request); + const callArgs = traceSpy.mock.calls[0]![2]; + expect(callArgs[0]).toMatchObject(request); + expect(callArgs[0]).not.toBe(request); expect(result).toMatchObject({ options: undefined, - request: traceContext.arguments[0], + request: callArgs[0], thisName: "agent-sdk", }); }); diff --git a/js/src/wrappers/openrouter-agent.ts b/js/src/wrappers/openrouter-agent.ts index 07328249b..1ef0b69af 100644 --- a/js/src/wrappers/openrouter-agent.ts +++ b/js/src/wrappers/openrouter-agent.ts @@ -1,7 +1,7 @@ import { openRouterAgentChannels } from "../instrumentation/plugins/openrouter-agent-channels"; import type { - OpenRouterAgentClient, OpenRouterAgentCallModelRequest, + OpenRouterAgentClient, } from "../vendor-sdk-types/openrouter-agent"; /** @@ -70,11 +70,11 @@ function wrapCallModel( const invocationTarget = thisArg === undefined ? (defaultThis ?? thisArg) : thisArg; - return openRouterAgentChannels.callModel.traceSync( - () => Reflect.apply(target, invocationTarget, [request, options]), - { - arguments: [request], - } as Parameters[1], + return openRouterAgentChannels.callModel.invoke( + target, + invocationTarget, + [request, options], + {}, ); }, }); diff --git a/js/src/wrappers/openrouter.ts b/js/src/wrappers/openrouter.ts index bfd39d2c1..19ac570ed 100644 --- a/js/src/wrappers/openrouter.ts +++ b/js/src/wrappers/openrouter.ts @@ -3,6 +3,8 @@ import type { OpenRouterBeta, OpenRouterCallModelRequest, OpenRouterChat, + OpenRouterChatCreateParams, + OpenRouterChatResult, OpenRouterClient, OpenRouterEmbeddingCreateParams, OpenRouterEmbeddingResponse, @@ -13,8 +15,6 @@ import type { OpenRouterResponses, OpenRouterResponsesCreateParams, OpenRouterResponsesResult, - OpenRouterChatCreateParams, - OpenRouterChatResult, } from "../vendor-sdk-types/openrouter"; /** @@ -137,9 +137,7 @@ function wrapChatSend( ) => Promise, ): OpenRouterChat["send"] { return (request, options) => - openRouterChannels.chatSend.tracePromise(() => send(request, options), { - arguments: [request], - } as Parameters[1]); + openRouterChannels.chatSend.invoke(send, undefined, [request, options], {}); } function wrapEmbeddingsGenerate( @@ -149,9 +147,11 @@ function wrapEmbeddingsGenerate( ) => Promise, ): OpenRouterEmbeddings["generate"] { return (request, options) => - openRouterChannels.embeddingsGenerate.tracePromise( - () => generate(request, options), - { arguments: [request] }, + openRouterChannels.embeddingsGenerate.invoke( + generate, + undefined, + [request, options], + {}, ); } @@ -162,9 +162,11 @@ function wrapResponsesSend( ) => Promise, ): OpenRouterResponses["send"] { return (request, options) => - openRouterChannels.betaResponsesSend.tracePromise( - () => send(request, options), - { arguments: [request] }, + openRouterChannels.betaResponsesSend.invoke( + send, + undefined, + [request, options], + {}, ); } @@ -175,9 +177,11 @@ function wrapRerank( ) => Promise, ): OpenRouterRerank["rerank"] { return (request, options) => - openRouterChannels.rerankRerank.tracePromise( - () => rerank(request, options), - { arguments: [request] }, + openRouterChannels.rerankRerank.invoke( + rerank, + undefined, + [request, options], + {}, ); } @@ -189,11 +193,11 @@ function wrapCallModel( ): NonNullable { return (request, options) => { const tracedRequest = { ...request }; - return openRouterChannels.callModel.traceSync( - () => callModel(tracedRequest, options), - { - arguments: [tracedRequest], - } as Parameters[1], + return openRouterChannels.callModel.invoke( + callModel, + undefined, + [tracedRequest, options], + {}, ); }; } diff --git a/js/src/wrappers/pi-coding-agent.test.ts b/js/src/wrappers/pi-coding-agent.test.ts index 04e469653..7c6438fb1 100644 --- a/js/src/wrappers/pi-coding-agent.test.ts +++ b/js/src/wrappers/pi-coding-agent.test.ts @@ -2,19 +2,17 @@ import { afterEach, describe, expect, it, vi } from "vitest"; const { invoke } = vi.hoisted(() => ({ invoke: vi.fn( - (target: (...args: any[]) => unknown, thisArg: unknown, args: unknown[]) => - Reflect.apply(target, thisArg, args), + ( + target: (...args: any[]) => any, + receiver: unknown, + args: unknown[], + _additional?: unknown, + ) => Reflect.apply(target, receiver, args), ), })); - -vi.mock("../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(() => ({ - hasInterceptors: true, - hasSubscribers: false, - invoke, - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { wrapPiCodingAgentSDK } from "./pi-coding-agent"; diff --git a/js/src/wrappers/strands-agent-sdk.test.ts b/js/src/wrappers/strands-agent-sdk.test.ts index 66d2ad46c..5300ea3ae 100644 --- a/js/src/wrappers/strands-agent-sdk.test.ts +++ b/js/src/wrappers/strands-agent-sdk.test.ts @@ -3,20 +3,16 @@ import { afterEach, describe, expect, it, vi } from "vitest"; const { invoke } = vi.hoisted(() => ({ invoke: vi.fn( ( - target: Function, - thisArg: unknown, + target: (...args: any[]) => any, + receiver: unknown, args: unknown[], _additional?: unknown, - ) => Reflect.apply(target, thisArg, args), + ) => Reflect.apply(target, receiver, args), ), })); - -vi.mock("../isomorph", () => ({ - default: { - newTracingChannel: vi.fn(() => ({ - invoke, - })), - }, +vi.mock("../global-instrumentation-hooks", async (importOriginal) => ({ + ...(await importOriginal()), + newGlobalInvocationHook: vi.fn(() => ({ invoke })), })); import { wrapStrandsAgentSDK } from "./strands-agent-sdk"; diff --git a/js/src/wrappers/strands-agent-sdk.ts b/js/src/wrappers/strands-agent-sdk.ts index 5cf339789..db82658f0 100644 --- a/js/src/wrappers/strands-agent-sdk.ts +++ b/js/src/wrappers/strands-agent-sdk.ts @@ -221,9 +221,7 @@ function wrapMultiAgentInstance( value as StrandsMultiAgent["stream"], target, callArgs, - { - orchestrator: proxy, - }, + { orchestrator: proxy }, ); }; } diff --git a/js/tests/auto-instrumentations/error-handling.test.ts b/js/tests/auto-instrumentations/error-handling.test.ts index 2c31588a4..161c3447a 100644 --- a/js/tests/auto-instrumentations/error-handling.test.ts +++ b/js/tests/auto-instrumentations/error-handling.test.ts @@ -96,9 +96,9 @@ describe("Error Handling", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify error event was emitted - expect(collector.error.length).toBeGreaterThan(0); - expect(collector.error[0].error).toBeDefined(); - expect(collector.error[0].error.message).toBe("Test error"); + expect(collector.failures.length).toBeGreaterThan(0); + expect(collector.failures[0].error).toBeDefined(); + expect(collector.failures[0].error.message).toBe("Test error"); }); it("should emit error event with correct error details", async () => { @@ -155,8 +155,8 @@ describe("Error Handling", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify error details are captured - expect(collector.error.length).toBeGreaterThan(0); - const errorEvent = collector.error[0]; + expect(collector.failures.length).toBeGreaterThan(0); + const errorEvent = collector.failures[0]; expect(errorEvent.error.message).toBe("API failure"); expect(errorEvent.error.name).toBe("CustomError"); expect(errorEvent.error.code).toBe("ERR_API_FAILURE"); @@ -207,7 +207,7 @@ describe("Error Handling", () => { // Also verify error event was emitted await new Promise((resolve) => setImmediate(resolve)); - expect(collector.error.length).toBeGreaterThan(0); + expect(collector.failures.length).toBeGreaterThan(0); }); it("should handle errors in promise rejections", async () => { @@ -251,8 +251,8 @@ describe("Error Handling", () => { // Verify error event was emitted await new Promise((resolve) => setImmediate(resolve)); - expect(collector.error.length).toBeGreaterThan(0); - expect(collector.error[0].error.message).toBe("Promise rejection"); + expect(collector.failures.length).toBeGreaterThan(0); + expect(collector.failures[0].error.message).toBe("Promise rejection"); }); }); @@ -303,12 +303,12 @@ describe("Error Handling", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify both start and error events were emitted - expect(collector.start.length).toBeGreaterThan(0); - expect(collector.error.length).toBeGreaterThan(0); + expect(collector.calls.length).toBeGreaterThan(0); + expect(collector.failures.length).toBeGreaterThan(0); // Verify start event came before error event - expect(collector.start[0].timestamp).toBeLessThanOrEqual( - collector.error[0].timestamp, + expect(collector.calls[0].timestamp).toBeLessThanOrEqual( + collector.failures[0].timestamp, ); }); @@ -358,14 +358,14 @@ describe("Error Handling", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify error event was emitted - expect(collector.error.length).toBeGreaterThan(0); + expect(collector.failures.length).toBeGreaterThan(0); // Verify end event was NOT emitted (or if emitted, came before or at the same time as error due to asyncEnd) // Note: For async functions, asyncEnd might still fire, but end should not - if (collector.end.length > 0) { + if (collector.returns.length > 0) { // If end event exists, it should be before or at the same time as the error - expect(collector.end[0].timestamp).toBeLessThanOrEqual( - collector.error[0].timestamp, + expect(collector.returns[0].timestamp).toBeLessThanOrEqual( + collector.failures[0].timestamp, ); } }); diff --git a/js/tests/auto-instrumentations/event-content.test.ts b/js/tests/auto-instrumentations/event-content.test.ts index 5a2fdd100..08090df9a 100644 --- a/js/tests/auto-instrumentations/event-content.test.ts +++ b/js/tests/auto-instrumentations/event-content.test.ts @@ -97,8 +97,8 @@ describe("Event Content Validation", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify start event was emitted with arguments - expect(collector.start.length).toBeGreaterThan(0); - const startEvent = collector.start[0]; + expect(collector.calls.length).toBeGreaterThan(0); + const startEvent = collector.calls[0]; expect(startEvent.arguments).toBeDefined(); expect(startEvent.arguments!.length).toBeGreaterThan(0); @@ -171,8 +171,8 @@ describe("Event Content Validation", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify all arguments were captured - expect(collector.start.length).toBeGreaterThan(0); - const startEvent = collector.start[0]; + expect(collector.calls.length).toBeGreaterThan(0); + const startEvent = collector.calls[0]; expect(startEvent.arguments).toBeDefined(); expect(startEvent.arguments!.length).toBeGreaterThanOrEqual(1); @@ -239,15 +239,15 @@ describe("Event Content Validation", () => { // Verify end event was emitted with result // Note: For async functions, result appears in asyncEnd event - const hasAsyncEnd = collector.asyncEnd.length > 0; - const hasEnd = collector.end.length > 0; + const hasAsyncEnd = collector.resolutions.length > 0; + const hasEnd = collector.returns.length > 0; expect(hasAsyncEnd || hasEnd).toBe(true); // Check the appropriate event type const resultEvent = hasAsyncEnd - ? collector.asyncEnd[0] - : collector.end[0]; + ? collector.resolutions[0] + : collector.returns[0]; expect(resultEvent.result).toBeDefined(); expect(resultEvent.result.id).toBe("chatcmpl-123"); expect(resultEvent.result.model).toBe("gpt-4"); @@ -318,14 +318,14 @@ describe("Event Content Validation", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify complex result structure is captured - const hasAsyncEnd = collector.asyncEnd.length > 0; - const hasEnd = collector.end.length > 0; + const hasAsyncEnd = collector.resolutions.length > 0; + const hasEnd = collector.returns.length > 0; expect(hasAsyncEnd || hasEnd).toBe(true); const resultEvent = hasAsyncEnd - ? collector.asyncEnd[0] - : collector.end[0]; + ? collector.resolutions[0] + : collector.returns[0]; expect(resultEvent.result).toBeDefined(); expect(resultEvent.result.usage.total_tokens).toBe(75); expect(resultEvent.result.choices[0].message.tool_calls).toHaveLength(1); @@ -379,8 +379,8 @@ describe("Event Content Validation", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify self context was captured - expect(collector.start.length).toBeGreaterThan(0); - const startEvent = collector.start[0]; + expect(collector.calls.length).toBeGreaterThan(0); + const startEvent = collector.calls[0]; expect(startEvent.self).toBeDefined(); // self should be the Completions instance @@ -433,16 +433,16 @@ describe("Event Content Validation", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify events were emitted - expect(collector.start.length).toBeGreaterThan(0); + expect(collector.calls.length).toBeGreaterThan(0); // For async functions, we expect asyncStart and asyncEnd - if (collector.asyncStart.length > 0 && collector.asyncEnd.length > 0) { + if (collector.promises.length > 0 && collector.resolutions.length > 0) { // Verify order: start <= asyncStart <= asyncEnd - expect(collector.start[0].timestamp).toBeLessThanOrEqual( - collector.asyncStart[0].timestamp, + expect(collector.calls[0].timestamp).toBeLessThanOrEqual( + collector.promises[0].timestamp, ); - expect(collector.asyncStart[0].timestamp).toBeLessThanOrEqual( - collector.asyncEnd[0].timestamp, + expect(collector.promises[0].timestamp).toBeLessThanOrEqual( + collector.resolutions[0].timestamp, ); } }); @@ -489,12 +489,12 @@ describe("Event Content Validation", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify we got 3 start events (one per call) - expect(collector.start.length).toBe(3); + expect(collector.calls.length).toBe(3); // Verify each call had different arguments - expect(collector.start[0].arguments![0].model).toBe("gpt-4"); - expect(collector.start[1].arguments![0].model).toBe("gpt-3.5-turbo"); - expect(collector.start[2].arguments![0].model).toBe("gpt-4-turbo"); + expect(collector.calls[0].arguments![0].model).toBe("gpt-4"); + expect(collector.calls[1].arguments![0].model).toBe("gpt-3.5-turbo"); + expect(collector.calls[2].arguments![0].model).toBe("gpt-4-turbo"); }); }); @@ -547,12 +547,12 @@ describe("Event Content Validation", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify correct channel received events - expect(correctCollector.start.length).toBeGreaterThan(0); + expect(correctCollector.calls.length).toBeGreaterThan(0); // Verify wrong channel did NOT receive events - expect(wrongCollector.start.length).toBe(0); - expect(wrongCollector.end.length).toBe(0); - expect(wrongCollector.asyncEnd.length).toBe(0); + expect(wrongCollector.calls.length).toBe(0); + expect(wrongCollector.returns.length).toBe(0); + expect(wrongCollector.resolutions.length).toBe(0); }); }); @@ -679,12 +679,12 @@ describe("Event Content Validation", () => { expect(result.isEventEmitter).toBe(true); // Verify start event captured arguments - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.start[0].arguments![0].model).toBe("gpt-4"); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.calls[0].arguments![0].model).toBe("gpt-4"); // Verify end event (stream returned synchronously) - expect(streamCollector.end.length).toBeGreaterThan(0); - expect(streamCollector.end[0].result).toBeDefined(); + expect(streamCollector.returns.length).toBeGreaterThan(0); + expect(streamCollector.returns[0].result).toBeDefined(); }); it("should return stream object synchronously for responses.stream", async () => { @@ -753,11 +753,11 @@ describe("Event Content Validation", () => { expect(result.hasEmit).toBe(true); // Verify start event captured arguments - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.start[0].arguments![0].model).toBe("gpt-4"); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.calls[0].arguments![0].model).toBe("gpt-4"); // Verify end event (stream returned synchronously) - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); }); @@ -858,8 +858,8 @@ describe("Event Content Validation", () => { expect(result.endCalled).toBe(true); // Verify start and end events were emitted - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); it("should verify handlers are called when stream emits events", async () => { @@ -969,8 +969,8 @@ describe("Event Content Validation", () => { expect(result.completions[0].id).toBe("chatcmpl-123"); // Verify instrumentation events - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); }); @@ -1071,8 +1071,8 @@ describe("Event Content Validation", () => { expect(result.chunkCount).toBe(2); // Verify instrumentation captured the start - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); // Note: Actual time_to_first_token logging happens in the instrumentation wrapper // This test verifies the stream emits chunks in the correct sequence @@ -1157,8 +1157,8 @@ describe("Event Content Validation", () => { expect(result.chunkCount).toBe(5); // Verify instrumentation events - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); }); @@ -1263,8 +1263,8 @@ describe("Event Content Validation", () => { expect(result.capturedCompletion.usage.total_tokens).toBe(15); // Verify instrumentation events - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); it("should log output on response.completed event for responses.stream", async () => { @@ -1369,8 +1369,8 @@ describe("Event Content Validation", () => { expect(result.eventCount).toBe(2); // Verify instrumentation events - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); }); @@ -1441,11 +1441,11 @@ describe("Event Content Validation", () => { expect(result.callDuration).toBeLessThan(50); // Verify start event was emitted immediately - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.start[0].arguments![0].model).toBe("gpt-4"); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.calls[0].arguments![0].model).toBe("gpt-4"); // Verify end event was also emitted (stream returned synchronously) - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); it("should NOT end span before stream completes", async () => { @@ -1541,8 +1541,8 @@ describe("Event Content Validation", () => { expect(result.streamEnded).toBe(true); // Verify instrumentation captured the entire lifecycle - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); // The start and end events are for the synchronous method call // The actual stream completion happens asynchronously via event listeners @@ -1629,8 +1629,8 @@ describe("Event Content Validation", () => { expect(result.errorMessage).toBe("Stream error occurred"); // Verify instrumentation events - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); it("should log error and end span when stream fails", async () => { @@ -1726,8 +1726,8 @@ describe("Event Content Validation", () => { expect(result.errorMessage).toBe("Network failure"); // Verify instrumentation events - expect(streamCollector.start.length).toBeGreaterThan(0); - expect(streamCollector.end.length).toBeGreaterThan(0); + expect(streamCollector.calls.length).toBeGreaterThan(0); + expect(streamCollector.returns.length).toBeGreaterThan(0); }); }); }); diff --git a/js/tests/auto-instrumentations/fixtures/configurable-global-hook-registry.cjs b/js/tests/auto-instrumentations/fixtures/configurable-global-hook-registry.cjs index c9393ef15..6c47cb187 100644 --- a/js/tests/auto-instrumentations/fixtures/configurable-global-hook-registry.cjs +++ b/js/tests/auto-instrumentations/fixtures/configurable-global-hook-registry.cjs @@ -1,26 +1,25 @@ const { parentPort } = require("node:worker_threads"); -const registryKey = "__braintrust_instrumentation_hooks"; +const registryKey = "__braintrust_invocation_hooks_v2"; const registryBrand = Symbol.for( "braintrust.global-instrumentation-hooks.registry", ); const foreignRegistry = new Map([["foreign", "entry"]]); globalThis[registryKey] = foreignRegistry; -const { newGlobalTracingChannel } = require( +const { newGlobalInvocationHook } = require( process.env.BRAINTRUST_TEST_GLOBAL_HOOK_RUNTIME, ); -const channel = newGlobalTracingChannel( +const channel = newGlobalInvocationHook( "orchestrion:test:configurable-registry", ); let subscriberCalls = 0; -channel.subscribe({ - start() { - subscriberCalls += 1; - }, +channel.intercept((target, receiver, args) => { + subscriberCalls += 1; + return Reflect.apply(target, receiver, args); }); -const result = channel.traceSync(() => "result"); +const result = channel.invoke(() => "result", undefined, [], {}); const descriptor = Object.getOwnPropertyDescriptor(globalThis, registryKey); parentPort?.postMessage({ diff --git a/js/tests/auto-instrumentations/fixtures/global-hook-listener.cjs b/js/tests/auto-instrumentations/fixtures/global-hook-listener.cjs index 2fed72dea..278c3c760 100644 --- a/js/tests/auto-instrumentations/fixtures/global-hook-listener.cjs +++ b/js/tests/auto-instrumentations/fixtures/global-hook-listener.cjs @@ -1,7 +1,7 @@ -const { newGlobalTracingChannel } = require( +const { newGlobalInvocationHook } = require( process.env.BRAINTRUST_TEST_GLOBAL_HOOK_RUNTIME, ); module.exports = { - getTracingHook: newGlobalTracingChannel, + getInvocationHook: newGlobalInvocationHook, }; diff --git a/js/tests/auto-instrumentations/fixtures/incompatible-global-hook-registry.cjs b/js/tests/auto-instrumentations/fixtures/incompatible-global-hook-registry.cjs index ab760cfba..476f0d2ce 100644 --- a/js/tests/auto-instrumentations/fixtures/incompatible-global-hook-registry.cjs +++ b/js/tests/auto-instrumentations/fixtures/incompatible-global-hook-registry.cjs @@ -1,35 +1,39 @@ const { parentPort } = require("node:worker_threads"); -Object.defineProperty(globalThis, "__braintrust_instrumentation_hooks", { +Object.defineProperty(globalThis, "__braintrust_invocation_hooks_v2", { configurable: false, enumerable: false, value: {}, writable: false, }); -const { newGlobalTracingChannel } = require( +const { newGlobalInvocationHook } = require( process.env.BRAINTRUST_TEST_GLOBAL_HOOK_RUNTIME, ); -const channel = newGlobalTracingChannel( +const channel = newGlobalInvocationHook( "orchestrion:test:incompatible-registry", ); let subscriberCalls = 0; let providerCalls = 0; -channel.subscribe({ - start() { - subscriberCalls += 1; - }, -}); -const result = channel.traceSync(() => { - providerCalls += 1; - return "result"; +channel.intercept((target, receiver, args) => { + subscriberCalls += 1; + return Reflect.apply(target, receiver, args); }); +const result = channel.invoke( + () => { + providerCalls += 1; + return "result"; + }, + undefined, + [], + {}, +); parentPort?.postMessage({ type: "incompatible-registry", result: { - hasSubscribers: channel.hasSubscribers, + hasInterceptors: channel.hasInterceptors, providerCalls, result, subscriberCalls, diff --git a/js/tests/auto-instrumentations/fixtures/listener-cjs.cjs b/js/tests/auto-instrumentations/fixtures/listener-cjs.cjs index 825c04309..a424037e1 100644 --- a/js/tests/auto-instrumentations/fixtures/listener-cjs.cjs +++ b/js/tests/auto-instrumentations/fixtures/listener-cjs.cjs @@ -1,29 +1,29 @@ const { parentPort } = require("worker_threads"); -const { getTracingHook } = require("./global-hook-listener.cjs"); +const { getInvocationHook } = require("./global-hook-listener.cjs"); const events = { start: [], end: [], error: [] }; // NOTE: code-transformer prepends "orchestrion:openai:" to the channel name const expectedChannel = "orchestrion:openai:chat.completions.create"; -// Subscribe to the global hook and accumulate events -const channel = getTracingHook(expectedChannel); -channel.subscribe({ - start: (ctx) => { - // Convert arguments to array for serialization - events.start.push({ - args: Array.from(ctx.arguments || []), - self: !!ctx.self, - }); - }, - end: (ctx) => { - // Only send serializable result data - events.end.push({ - result: ctx.result ? JSON.parse(JSON.stringify(ctx.result)) : null, - }); - }, - error: (ctx) => { - events.error.push({ error: String(ctx.error) }); - }, +getInvocationHook(expectedChannel).intercept((target, receiver, args) => { + events.start.push({ args, self: !!receiver }); + try { + const result = Reflect.apply(target, receiver, args); + Promise.resolve(result).then( + (value) => { + events.end.push({ + result: value ? JSON.parse(JSON.stringify(value)) : null, + }); + }, + (error) => { + events.error.push({ error: String(error) }); + }, + ); + return result; + } catch (error) { + events.error.push({ error: String(error) }); + throw error; + } }); // Send all accumulated events on exit diff --git a/js/tests/auto-instrumentations/fixtures/listener-esm.mjs b/js/tests/auto-instrumentations/fixtures/listener-esm.mjs index c735b77fe..207f24d38 100644 --- a/js/tests/auto-instrumentations/fixtures/listener-esm.mjs +++ b/js/tests/auto-instrumentations/fixtures/listener-esm.mjs @@ -1,7 +1,7 @@ import { createRequire } from "node:module"; import { parentPort } from "node:worker_threads"; -const { getTracingHook } = createRequire(import.meta.url)( +const { getInvocationHook } = createRequire(import.meta.url)( "./global-hook-listener.cjs", ); @@ -9,24 +9,25 @@ const events = { start: [], end: [], error: [] }; // NOTE: code-transformer prepends "orchestrion:openai:" to the channel name const expectedChannel = "orchestrion:openai:chat.completions.create"; -// Subscribe to the global hook and accumulate events -const channel = getTracingHook(expectedChannel); -channel.subscribe({ - start: (ctx) => { - events.start.push({ - args: ctx.arguments ? Array.from(ctx.arguments) : [], - self: !!ctx.self, - }); - }, - asyncEnd: (ctx) => { - // Only send serializable result data - events.end.push({ - result: ctx.result ? JSON.parse(JSON.stringify(ctx.result)) : null, - }); - }, - error: (ctx) => { - events.error.push({ error: String(ctx.error) }); - }, +getInvocationHook(expectedChannel).intercept((target, receiver, args) => { + events.start.push({ args, self: !!receiver }); + try { + const result = Reflect.apply(target, receiver, args); + Promise.resolve(result).then( + (value) => { + events.end.push({ + result: value ? JSON.parse(JSON.stringify(value)) : null, + }); + }, + (error) => { + events.error.push({ error: String(error) }); + }, + ); + return result; + } catch (error) { + events.error.push({ error: String(error) }); + throw error; + } }); // Send all accumulated events on exit diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/arguments_mutation/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/arguments_mutation/test.js index 9a7917da9..f868c5ba2 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/arguments_mutation/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/arguments_mutation/test.js @@ -3,25 +3,24 @@ * This product includes software developed at Datadog (https://www.datadoghq.com/). Copyright 2025 Datadog, Inc. **/ const { fetch_simple, fetch_complex } = require("./instrumented.js"); -const { assert, getTracingHook } = require("../common/preamble.js"); +const { assert, getInvocationHook } = require("../common/preamble.js"); -const handler = { - start(message) { - const originalCb = message.arguments[1]; - const wrappedCb = function (a, b) { - assert.strictEqual(this.this, "this"); - assert.strictEqual(a, "arg1"); - assert.strictEqual(b, "arg2"); - arguments[1] = "arg2_mutated"; - return originalCb.apply(this, arguments); - }; +const handler = (target, receiver, args) => { + const originalCb = args[1]; + const wrappedCb = function (a, b) { + assert.strictEqual(this.this, "this"); + assert.strictEqual(a, "arg1"); + assert.strictEqual(b, "arg2"); + arguments[1] = "arg2_mutated"; + return originalCb.apply(this, arguments); + }; - message.arguments[1] = wrappedCb; - }, + args[1] = wrappedCb; + return Reflect.apply(target, receiver, args); }; -getTracingHook("orchestrion:undici:fetch_simple").subscribe(handler); -getTracingHook("orchestrion:undici:fetch.complex").subscribe(handler); +getInvocationHook("orchestrion:undici:fetch_simple").intercept(handler); +getInvocationHook("orchestrion:undici:fetch.complex").intercept(handler); assert.strictEqual(fetch_simple.length, 2); assert.strictEqual(fetch_complex.length, 2); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/ast_query_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/ast_query_cjs/test.js index 781fd1c7d..fb805775f 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/ast_query_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/ast_query_cjs/test.js @@ -10,9 +10,7 @@ const context = getContext("orchestrion:undici:fetch_ast_query"); const result = await fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/callback_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/callback_cjs/test.js index a51ecbc66..d8d0fa50e 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/callback_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/callback_cjs/test.js @@ -15,9 +15,7 @@ const context = getContext("orchestrion:undici:fetch.cb"); }); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_expression_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_expression_cjs/test.js index 457789642..b2da80cb5 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_expression_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_expression_cjs/test.js @@ -10,9 +10,7 @@ const context = getContext("orchestrion:undici:Undici:fetch"); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_method_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_method_cjs/test.js index 457789642..b2da80cb5 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_method_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/class_method_cjs/test.js @@ -10,9 +10,7 @@ const context = getContext("orchestrion:undici:Undici:fetch"); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/common/preamble.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/common/preamble.js index f395843bd..18844e7e5 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/common/preamble.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/common/preamble.js @@ -6,25 +6,31 @@ const assert = require("node:assert"); const runtimePath = process.env.BRAINTRUST_TEST_GLOBAL_HOOK_RUNTIME; assert(runtimePath, "BRAINTRUST_TEST_GLOBAL_HOOK_RUNTIME must be set"); -const { newGlobalTracingChannel } = require(runtimePath); +const { newGlobalInvocationHook } = require(runtimePath); function getContext(channelName) { - const channel = newGlobalTracingChannel(channelName); const context = {}; - channel.subscribe({ - start(message) { - message.context = context; - context.start = true; - }, - end(message) { - message.context.end = message.result ?? true; - }, - asyncStart(message) { - message.context.asyncStart = message.result; - }, - asyncEnd(message) { - message.context.asyncEnd = message.result; - }, + newGlobalInvocationHook(channelName).intercept((target, receiver, args) => { + context.called = true; + const callbackIndex = args.findLastIndex( + (arg) => typeof arg === "function", + ); + if (callbackIndex !== -1) { + const callback = args[callbackIndex]; + args[callbackIndex] = function (...callbackArgs) { + context.result = callbackArgs[1]; + return Reflect.apply(callback, this, callbackArgs); + }; + } + const result = Reflect.apply(target, receiver, args); + if (result && typeof result.then === "function") { + result.then((value) => { + context.result = value; + }); + } else if (callbackIndex === -1) { + context.result = result; + } + return result; }); return context; } @@ -32,5 +38,5 @@ function getContext(channelName) { module.exports = { assert, getContext, - getTracingHook: newGlobalTracingChannel, + getInvocationHook: newGlobalInvocationHook, }; diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/const_class_export_alias_mjs/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/const_class_export_alias_mjs/test.mjs index 8fba82870..02ff37c0a 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/const_class_export_alias_mjs/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/const_class_export_alias_mjs/test.mjs @@ -9,8 +9,6 @@ const undici = new Undici(); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_cjs/test.js index f39df4014..efb6db492 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_cjs/test.js @@ -9,9 +9,7 @@ const context = getContext("orchestrion:undici:fetch.decl"); const result = await fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs/test.mjs index 07b7c2d35..926900efd 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs/test.mjs @@ -8,8 +8,6 @@ const context = getContext("orchestrion:undici:fetch_decl"); const result = await fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs_mismatched_type/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs_mismatched_type/test.mjs index 07b7c2d35..926900efd 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs_mismatched_type/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/decl_mjs_mismatched_type/test.mjs @@ -8,8 +8,6 @@ const context = getContext("orchestrion:undici:fetch_decl"); const result = await fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_class_mjs/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_class_mjs/test.mjs index 8fba82870..02ff37c0a 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_class_mjs/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_class_mjs/test.mjs @@ -9,8 +9,6 @@ const undici = new Undici(); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_mjs/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_mjs/test.mjs index 77f065a29..1b38714a8 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_mjs/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/export_alias_mjs/test.mjs @@ -8,8 +8,6 @@ const context = getContext("orchestrion:undici:fetch_alias"); const result = await fetchAliased("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/iife_nested_class/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/iife_nested_class/test.js index 939cd6f1f..be26781a9 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/iife_nested_class/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/iife_nested_class/test.js @@ -10,7 +10,7 @@ const context = getContext("orchestrion:undici:register"); const result = server.register(); assert.strictEqual(result, 1); assert.deepStrictEqual(context, { - start: true, - end: 1, + called: true, + result: 1, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/index_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/index_cjs/test.js index 8e48e0f26..b018e004b 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/index_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/index_cjs/test.js @@ -7,24 +7,20 @@ const { assert, getContext } = require("../common/preamble.js"); const context = getContext("orchestrion:undici:Undici_fetch"); async function testOne(Undici, num, expectedCtx) { + delete context.called; + delete context.result; const undici = new Undici(); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, num); assert.deepStrictEqual(context, expectedCtx); - delete context.start; - delete context.end; - delete context.asyncStart; - delete context.asyncEnd; } (async () => { await testOne(undicis.Undici0, 0, {}); await testOne(undicis.Undici1, 1, {}); await testOne(undicis.Undici2, 2, { - start: true, - end: true, - asyncStart: 2, - asyncEnd: 2, + called: true, + result: 2, }); await testOne(undicis.Undici3, 3, {}); await testOne(undicis.Undici4, 4, {}); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/instance_method_subclass_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/instance_method_subclass_cjs/test.js index e94a36d5a..f0cf493f5 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/instance_method_subclass_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/instance_method_subclass_cjs/test.js @@ -10,9 +10,7 @@ const context = getContext("orchestrion:undici:Base_fetch"); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/let_class_export_alias_mjs/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/let_class_export_alias_mjs/test.mjs index 8fba82870..02ff37c0a 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/let_class_export_alias_mjs/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/let_class_export_alias_mjs/test.mjs @@ -9,8 +9,6 @@ const undici = new Undici(); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_class_method_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_class_method_cjs/test.js index 98e5f38a2..e3e23b1ac 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_class_method_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_class_method_cjs/test.js @@ -12,17 +12,13 @@ const context2 = getContext("orchestrion:undici:Undici_fetch2"); const result1 = await undici.fetch1("https://example.com"); assert.strictEqual(result1, 42); assert.deepStrictEqual(context1, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); const result2 = await undici.fetch2("https://example.com"); assert.strictEqual(result2, 43); assert.deepStrictEqual(context2, { - start: true, - end: true, - asyncStart: 43, - asyncEnd: 43, + called: true, + result: 43, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_load_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_load_cjs/test.js index bb545a5ef..c83144877 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_load_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/multiple_load_cjs/test.js @@ -10,9 +10,7 @@ const context = getContext("orchestrion:undici:Undici_fetch"); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/nested_functions/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/nested_functions/test.js index b04cbd031..e7d7d02e5 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/nested_functions/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/nested_functions/test.js @@ -10,7 +10,7 @@ const context = getContext("orchestrion:undici:nested_fn"); const result = f.addHook(); assert.strictEqual(result, "Hook added"); assert.deepStrictEqual(context, { - start: true, - end: "Hook added", + called: true, + result: "Hook added", }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_method_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_method_cjs/test.js index bf22412ba..4cb68c56f 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_method_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_method_cjs/test.js @@ -9,9 +9,7 @@ const context = getContext("orchestrion:undici:Undici_fetch"); const result = await fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_named_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_named_cjs/test.js index 204595b5a..fd8fb1d03 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_named_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_named_cjs/test.js @@ -12,9 +12,7 @@ const context = getContext("orchestrion:undici:conn_query"); const result = await conn.query(); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_this_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_this_cjs/test.js index 1b895d60c..124654776 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_this_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/object_property_this_cjs/test.js @@ -13,9 +13,7 @@ const context = getContext("orchestrion:undici:Connection_query"); const result = await conn._query(); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/private_method_cjs/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/private_method_cjs/test.js index d8ebd2a6c..8be36651a 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/private_method_cjs/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/private_method_cjs/test.js @@ -10,9 +10,7 @@ const context = getContext("orchestrion:undici:TestClass:testMe"); const result = await test.testMe(); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/promise_subclass/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/promise_subclass/test.js index daa4db7cc..8d0ef7b8c 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/promise_subclass/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/promise_subclass/test.js @@ -17,9 +17,7 @@ const context = getContext("orchestrion:undici:fetch_subclass"); const { data } = await promise.withResponse(); assert.strictEqual(data, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_class_export_alias_mjs/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_class_export_alias_mjs/test.mjs index 8fba82870..02ff37c0a 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_class_export_alias_mjs/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_class_export_alias_mjs/test.mjs @@ -9,8 +9,6 @@ const undici = new Undici(); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_named_class_export_alias_mjs/test.mjs b/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_named_class_export_alias_mjs/test.mjs index 8fba82870..02ff37c0a 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_named_class_export_alias_mjs/test.mjs +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/var_named_class_export_alias_mjs/test.mjs @@ -9,8 +9,6 @@ const undici = new Undici(); const result = await undici.fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/windows_path/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/windows_path/test.js index 0c66e7227..50328a5a9 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/windows_path/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/windows_path/test.js @@ -9,9 +9,7 @@ const context = getContext("orchestrion:undici:fetch_decl"); const result = await fetch("https://example.com"); assert.strictEqual(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); })(); diff --git a/js/tests/auto-instrumentations/fixtures/orchestrion-js/wrap_promise_non_promise/test.js b/js/tests/auto-instrumentations/fixtures/orchestrion-js/wrap_promise_non_promise/test.js index 60bc3612f..3cf5c2625 100644 --- a/js/tests/auto-instrumentations/fixtures/orchestrion-js/wrap_promise_non_promise/test.js +++ b/js/tests/auto-instrumentations/fixtures/orchestrion-js/wrap_promise_non_promise/test.js @@ -8,8 +8,6 @@ const context = getContext("orchestrion:undici:fetch_nonpromise"); const result = fetch("https://example.com"); assert.equal(result, 42); assert.deepStrictEqual(context, { - start: true, - end: true, - asyncStart: 42, - asyncEnd: 42, + called: true, + result: 42, }); diff --git a/js/tests/auto-instrumentations/loader-hook.test.ts b/js/tests/auto-instrumentations/loader-hook.test.ts index 688e1aa97..6b5adc569 100644 --- a/js/tests/auto-instrumentations/loader-hook.test.ts +++ b/js/tests/auto-instrumentations/loader-hook.test.ts @@ -156,7 +156,7 @@ describe("Unified Loader Hook Integration Tests", () => { }); expect(result).toEqual({ - hasSubscribers: false, + hasInterceptors: false, providerCalls: 1, result: "result", subscriberCalls: 0, diff --git a/js/tests/auto-instrumentations/multiple-instrumentations.test.ts b/js/tests/auto-instrumentations/multiple-instrumentations.test.ts index 2b48ff0a6..aaea4d691 100644 --- a/js/tests/auto-instrumentations/multiple-instrumentations.test.ts +++ b/js/tests/auto-instrumentations/multiple-instrumentations.test.ts @@ -127,12 +127,12 @@ describe("Multiple Instrumentations", () => { expect(result.embeddingResult.type).toBe("embeddings"); // Verify chat completions channel received events - expect(chatCollector.start.length).toBeGreaterThan(0); - expect(chatCollector.start[0].arguments![0].model).toBe("gpt-4"); + expect(chatCollector.calls.length).toBeGreaterThan(0); + expect(chatCollector.calls[0].arguments![0].model).toBe("gpt-4"); // Verify embeddings channel received events - expect(embeddingsCollector.start.length).toBeGreaterThan(0); - expect(embeddingsCollector.start[0].arguments![0].model).toBe( + expect(embeddingsCollector.calls.length).toBeGreaterThan(0); + expect(embeddingsCollector.calls[0].arguments![0].model).toBe( "text-embedding-ada-002", ); }); @@ -191,12 +191,12 @@ describe("Multiple Instrumentations", () => { await new Promise((resolve) => setImmediate(resolve)); // Chat completions channel should have events - expect(chatCollector.start.length).toBeGreaterThan(0); + expect(chatCollector.calls.length).toBeGreaterThan(0); // Embeddings channel should NOT have events (we didn't call it) - expect(embeddingsCollector.start.length).toBe(0); - expect(embeddingsCollector.end.length).toBe(0); - expect(embeddingsCollector.asyncEnd.length).toBe(0); + expect(embeddingsCollector.calls.length).toBe(0); + expect(embeddingsCollector.returns.length).toBe(0); + expect(embeddingsCollector.resolutions.length).toBe(0); }); }); @@ -256,16 +256,16 @@ describe("Multiple Instrumentations", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify correct number of calls to each method - expect(chatCollector.start.length).toBe(3); - expect(embeddingsCollector.start.length).toBe(2); + expect(chatCollector.calls.length).toBe(3); + expect(embeddingsCollector.calls.length).toBe(2); // Verify the right models were passed - expect(chatCollector.start[0].arguments![0].model).toBe("gpt-4"); - expect(chatCollector.start[1].arguments![0].model).toBe("gpt-3.5-turbo"); - expect(chatCollector.start[2].arguments![0].model).toBe("gpt-4-turbo"); + expect(chatCollector.calls[0].arguments![0].model).toBe("gpt-4"); + expect(chatCollector.calls[1].arguments![0].model).toBe("gpt-3.5-turbo"); + expect(chatCollector.calls[2].arguments![0].model).toBe("gpt-4-turbo"); - expect(embeddingsCollector.start[0].arguments![0].model).toBe("ada-002"); - expect(embeddingsCollector.start[1].arguments![0].model).toBe("ada-003"); + expect(embeddingsCollector.calls[0].arguments![0].model).toBe("ada-002"); + expect(embeddingsCollector.calls[1].arguments![0].model).toBe("ada-003"); }); }); @@ -339,16 +339,16 @@ describe("Multiple Instrumentations", () => { expect(callOrder).toHaveLength(4); // Verify we got events for all calls - expect(chatCollector.start.length).toBe(2); - expect(embeddingsCollector.start.length).toBe(2); + expect(chatCollector.calls.length).toBe(2); + expect(embeddingsCollector.calls.length).toBe(2); // Verify the events contain the right data (order may vary due to random delays) - const chatModels = chatCollector.start + const chatModels = chatCollector.calls .map((e) => e.arguments![0].model) .sort(); expect(chatModels).toEqual(["gpt-3.5", "gpt-4"]); - const embedInputs = embeddingsCollector.start + const embedInputs = embeddingsCollector.calls .map((e) => e.arguments![0].input) .sort(); expect(embedInputs).toEqual(["embed1", "embed2"]); @@ -447,8 +447,8 @@ describe("Multiple Instrumentations", () => { expect(result.processed).toBe("test"); // Verify custom instrumentation emitted events - expect(customCollector.start.length).toBeGreaterThan(0); - expect(customCollector.start[0].arguments![0].data).toBe("test"); + expect(customCollector.calls.length).toBeGreaterThan(0); + expect(customCollector.calls[0].arguments![0].data).toBe("test"); // Clean up fs.rmSync(customSdkDir, { recursive: true, force: true }); @@ -514,11 +514,11 @@ describe("Multiple Instrumentations", () => { expect(result2.client).toBe("client2"); // Verify both calls emitted events - expect(collector.start.length).toBe(2); + expect(collector.calls.length).toBe(2); // Verify each call had the right self context - expect(collector.start[0].self._client.name).toBe("client1"); - expect(collector.start[1].self._client.name).toBe("client2"); + expect(collector.calls[0].self._client.name).toBe("client1"); + expect(collector.calls[1].self._client.name).toBe("client2"); }); }); }); diff --git a/js/tests/auto-instrumentations/streaming-and-responses.test.ts b/js/tests/auto-instrumentations/streaming-and-responses.test.ts index e994de00e..663d32fa4 100644 --- a/js/tests/auto-instrumentations/streaming-and-responses.test.ts +++ b/js/tests/auto-instrumentations/streaming-and-responses.test.ts @@ -177,12 +177,12 @@ describe("Streaming Methods and Responses API", () => { await new Promise((resolve) => setTimeout(resolve, 100)); // Verify start event (method called) - expect(collector.start.length).toBeGreaterThan(0); - const startEvent = collector.start[0]; + expect(collector.calls.length).toBeGreaterThan(0); + const startEvent = collector.calls[0]; expect(startEvent.arguments![0].model).toBe("gpt-4"); // Verify end event (stream returned synchronously) - expect(collector.end.length).toBeGreaterThan(0); + expect(collector.returns.length).toBeGreaterThan(0); }); }); @@ -238,11 +238,11 @@ describe("Streaming Methods and Responses API", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify events were captured - expect(collector.start.length).toBeGreaterThan(0); - expect(collector.asyncEnd.length).toBeGreaterThan(0); + expect(collector.calls.length).toBeGreaterThan(0); + expect(collector.resolutions.length).toBeGreaterThan(0); // Verify input was captured - const startEvent = collector.start[0]; + const startEvent = collector.calls[0]; expect(startEvent.arguments![0].model).toBe("gpt-4"); // Verify result @@ -335,8 +335,8 @@ describe("Streaming Methods and Responses API", () => { expect(result).toHaveLength(3); // Verify events were captured - expect(collector.start.length).toBeGreaterThan(0); - expect(collector.asyncEnd.length).toBeGreaterThan(0); + expect(collector.calls.length).toBeGreaterThan(0); + expect(collector.resolutions.length).toBeGreaterThan(0); }); }); @@ -426,8 +426,8 @@ describe("Streaming Methods and Responses API", () => { expect(result.events[1].type).toBe("response.completed"); // Verify instrumentation events - expect(collector.start.length).toBeGreaterThan(0); - expect(collector.end.length).toBeGreaterThan(0); + expect(collector.calls.length).toBeGreaterThan(0); + expect(collector.returns.length).toBeGreaterThan(0); }); }); @@ -483,11 +483,11 @@ describe("Streaming Methods and Responses API", () => { await new Promise((resolve) => setImmediate(resolve)); // Verify events were captured - expect(collector.start.length).toBeGreaterThan(0); - expect(collector.asyncEnd.length).toBeGreaterThan(0); + expect(collector.calls.length).toBeGreaterThan(0); + expect(collector.resolutions.length).toBeGreaterThan(0); // Verify input was captured - const startEvent = collector.start[0]; + const startEvent = collector.calls[0]; expect(startEvent.arguments![0].model).toBe("gpt-4"); // Verify result @@ -580,8 +580,8 @@ describe("Streaming Methods and Responses API", () => { expect(result[0].choices[0].delta.role).toBe("assistant"); // Verify events were captured - expect(collector.start.length).toBeGreaterThan(0); - expect(collector.asyncEnd.length).toBeGreaterThan(0); + expect(collector.calls.length).toBeGreaterThan(0); + expect(collector.resolutions.length).toBeGreaterThan(0); }); }); }); diff --git a/js/tests/auto-instrumentations/test-helpers.ts b/js/tests/auto-instrumentations/test-helpers.ts index f2ddd23e0..3d0c5be6d 100644 --- a/js/tests/auto-instrumentations/test-helpers.ts +++ b/js/tests/auto-instrumentations/test-helpers.ts @@ -1,12 +1,5 @@ -/** - * Test helpers for functional testing of instrumented code. - */ - -import { - newGlobalTracingChannel, - type GlobalHookHandlers, - type GlobalTracingChannel, -} from "../../src/global-instrumentation-hooks"; +/** Test-local recording through ordinary invocation interceptors. */ +import { newGlobalInvocationHook } from "../../src/global-instrumentation-hooks"; export interface CapturedEvent { arguments?: any[]; @@ -15,95 +8,68 @@ export interface CapturedEvent { error?: any; timestamp: number; } - -export interface EventCollector { - start: CapturedEvent[]; - end: CapturedEvent[]; - asyncStart: CapturedEvent[]; - asyncEnd: CapturedEvent[]; - error: CapturedEvent[]; - clear: () => void; - subscribe: (channelName: string) => void; - unsubscribe: () => void; -} - -/** - * Creates an event collector for capturing global instrumentation hook events. - */ -export function createEventCollector(): EventCollector { - const subscriptions: Array<{ - channel: GlobalTracingChannel; - handlers: GlobalHookHandlers; - }> = []; - const collector: EventCollector = { - start: [], - end: [], - asyncStart: [], - asyncEnd: [], - error: [], +export function createEventCollector() { + const removals: Array<() => void> = []; + return { + calls: [] as CapturedEvent[], + returns: [] as CapturedEvent[], + promises: [] as CapturedEvent[], + resolutions: [] as CapturedEvent[], + failures: [] as CapturedEvent[], clear() { - this.start = []; - this.end = []; - this.asyncStart = []; - this.asyncEnd = []; - this.error = []; + this.calls = []; + this.returns = []; + this.promises = []; + this.resolutions = []; + this.failures = []; }, subscribe(channelName: string) { - const channel = newGlobalTracingChannel(channelName); - const handlers = { - start: (ctx: any) => { - this.start.push({ - arguments: ctx.arguments ? Array.from(ctx.arguments) : undefined, - self: ctx.self, - timestamp: Date.now(), - }); - }, - end: (ctx: any) => { - this.end.push({ - result: ctx.result, - timestamp: Date.now(), - }); - }, - asyncStart: (ctx: any) => { - this.asyncStart.push({ - timestamp: Date.now(), - }); - }, - asyncEnd: (ctx: any) => { - this.asyncEnd.push({ - result: ctx.result, - timestamp: Date.now(), - }); - }, - error: (ctx: any) => { - this.error.push({ - error: ctx.error, - timestamp: Date.now(), - }); - }, - }; - channel.subscribe(handlers); - subscriptions.push({ channel, handlers }); + removals.push( + newGlobalInvocationHook(channelName).intercept( + (target, receiver, args) => { + this.calls.push({ + arguments: args, + self: receiver, + timestamp: Date.now(), + }); + let result; + try { + result = Reflect.apply(target, receiver, args); + } catch (error) { + this.failures.push({ error, timestamp: Date.now() }); + throw error; + } + this.returns.push({ result, timestamp: Date.now() }); + if (result && typeof result.then === "function") { + this.promises.push({ result, timestamp: Date.now() }); + result.then( + (value: unknown) => { + this.resolutions.push({ + result: value, + timestamp: Date.now(), + }); + }, + (error: unknown) => { + this.failures.push({ error, timestamp: Date.now() }); + }, + ); + } + return result; + }, + ), + ); }, unsubscribe() { - for (const { channel, handlers } of subscriptions.splice(0)) { - channel.unsubscribe(handlers); - } + for (const remove of removals.splice(0)) remove(); }, }; - - return collector; } - -/** - * Helper to run a function and wait for all events to be emitted. - */ +export type EventCollector = ReturnType; export async function runAndCollectEvents( fn: () => T | Promise, - collector: EventCollector, + _collector: EventCollector, ): Promise { const result = await fn(); - // Give event handlers a chance to run await new Promise((resolve) => setImmediate(resolve)); return result; } diff --git a/js/tests/auto-instrumentations/transformation.test.ts b/js/tests/auto-instrumentations/transformation.test.ts index 0a21456cd..d95558880 100644 --- a/js/tests/auto-instrumentations/transformation.test.ts +++ b/js/tests/auto-instrumentations/transformation.test.ts @@ -7,13 +7,13 @@ * IMPORTANT: Tests use a mock OpenAI package structure in test/fixtures/node_modules/openai. */ -import { describe, it, expect, beforeAll, afterAll, vi } from "vitest"; import * as esbuild from "esbuild"; -import { build as viteBuild } from "vite"; import * as fs from "node:fs"; import * as path from "node:path"; import { fileURLToPath } from "node:url"; import { runInNewContext } from "node:vm"; +import { build as viteBuild } from "vite"; +import { afterAll, beforeAll, describe, expect, it, vi } from "vitest"; import { create, type InstrumentationConfig, @@ -22,7 +22,7 @@ import { GLOBAL_INSTRUMENTATION_HOOKS_KEY, GLOBAL_INSTRUMENTATION_HOOKS_PROTOCOL_VERSION, GLOBAL_INSTRUMENTATION_HOOKS_REGISTRY_BRAND, - newGlobalTracingChannel, + newGlobalInvocationHook, } from "../../src/global-instrumentation-hooks"; const __dirname = path.dirname(fileURLToPath(import.meta.url)); @@ -82,14 +82,16 @@ function transformTestCode( } function expectGlobalHookTransform(output: string): void { - expect(output).toContain("__braintrust_instrumentation_hooks"); + expect(output).toContain("__braintrust_invocation_hooks_v2"); expect(output).toContain("braintrust.global-instrumentation-hooks.registry"); - expect(output).toContain("braintrust.global-instrumentation-hooks.hook"); + expect(output).toContain( + "braintrust.global-instrumentation-hooks.invocation-hook", + ); expect(output).toContain("orchestrion:openai:chat.completions.create"); - expect(output).toContain("__bt$hook.traceInvocation"); + expect(output).toContain("__bt$hook.invoke"); expect(output).not.toContain("__bt$hook.hasSubscribers"); expect(output).not.toContain("__bt$hook.hasInterceptors"); - expect(output).not.toContain("__bt$hook.invoke"); + expect(output).not.toContain("__bt$hook.traceInvocation"); expect(output).not.toContain("__bt$hook.tracePromise"); expect(output).not.toContain("__apm$"); expect(output).not.toContain("tr_ch_apm$"); @@ -170,7 +172,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports method-only configs", () => { @@ -186,7 +188,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports function declaration configs", () => { @@ -200,7 +202,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("ignores malformed global hook entries at runtime", () => { @@ -256,7 +258,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports export-alias class method configs", () => { @@ -278,7 +280,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports private class method configs", () => { @@ -300,7 +302,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports object/property configs", () => { @@ -318,7 +320,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports callback configs", () => { @@ -332,7 +334,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports raw AST query configs", () => { @@ -348,7 +350,7 @@ describe("Orchestrion Transformation Tests", () => { ); expect(result.code).toContain("orchestrion:test-sdk:test"); - expect(result.code).toContain("__bt$hook.traceInvocation"); + expect(result.code).toContain("__bt$hook.invoke"); }); it("supports index selection", () => { @@ -368,12 +370,10 @@ describe("Orchestrion Transformation Tests", () => { `, ); - const wrapperCount = result.code.match( - /return __bt\$hook\.traceInvocation/g, - ); + const wrapperCount = result.code.match(/return __bt\$hook\.invoke/g); expect(wrapperCount).toHaveLength(1); expect(result.code.indexOf("secondCreate")).toBeLessThan( - result.code.indexOf("return __bt$hook.traceInvocation"), + result.code.indexOf("return __bt$hook.invoke"), ); }); @@ -415,11 +415,9 @@ describe("Orchestrion Transformation Tests", () => { ).toHaveLength(1); expect(result.code).toContain("orchestrion:test-sdk:first"); expect(result.code).toContain("orchestrion:test-sdk:second"); - expect( - result.code.match(/return __bt\$hook\.traceInvocation/g), - ).toHaveLength(2); - expect(result.code).toContain('"traceSync"'); - expect(result.code).toContain('"tracePromise"'); + expect(result.code.match(/return __bt\$hook\.invoke/g)).toHaveLength(2); + expect(result.code).not.toContain('"traceSync"'); + expect(result.code).not.toContain('"tracePromise"'); }); it("generates source maps", () => { @@ -462,16 +460,18 @@ describe("Orchestrion Transformation Tests", () => { expect(loadedModule.exports.query("before")).toBe("before"); const events: unknown[] = []; - const hook = newGlobalTracingChannel("orchestrion:test-sdk:test"); - const handlers = { start: (event: unknown) => events.push(event) }; - hook.subscribe(handlers); + const hook = newGlobalInvocationHook("orchestrion:test-sdk:test"); + const remove = hook.intercept((target, receiver, args) => { + events.push(args); + return Reflect.apply(target, receiver, args); + }); expect(loadedModule.exports.query("after")).toBe("after"); expect(events).toHaveLength(1); - hook.unsubscribe(handlers); + remove(); }); - it("invokes generic interceptors inside tracing hooks", () => { + it("invokes generic interceptors with receiver, arguments, and version metadata", () => { const result = transformTestCode( { className: "Client", methodName: "query", kind: "Sync" }, ` @@ -498,19 +498,10 @@ describe("Orchestrion Transformation Tests", () => { result.code, )(loadedModule, loadedModule.exports); - const hook = newGlobalTracingChannel>( + const hook = newGlobalInvocationHook<{ moduleVersion: string }>( "orchestrion:test-sdk:test", ); const lifecycle: string[] = []; - const handlers = { - start: (event: Record) => - lifecycle.push( - `start:${Array.from(event.arguments as ArrayLike)[0]}`, - ), - end: (event: Record) => - lifecycle.push(`end:${event.result}`), - }; - hook.subscribe(handlers); const removeInterceptor = hook.intercept( (target, _thisArg, args, additional: { moduleVersion: string }) => { lifecycle.push(`intercept:${additional.moduleVersion}`); @@ -522,14 +513,9 @@ describe("Orchestrion Transformation Tests", () => { const client = new loadedModule.exports("original"); expect(client.query("value")).toBe("patched:VALUE!"); - expect(lifecycle).toEqual([ - "start:value", - "intercept:1.0.0", - "end:patched:VALUE!", - ]); + expect(lifecycle).toEqual(["intercept:1.0.0"]); removeInterceptor(); - hook.unsubscribe(handlers); }); it("supports asynchronous invocation interceptors", async () => { @@ -552,7 +538,7 @@ describe("Orchestrion Transformation Tests", () => { "exports", result.code, )(loadedModule, loadedModule.exports); - const hook = newGlobalTracingChannel("orchestrion:test-sdk:test"); + const hook = newGlobalInvocationHook("orchestrion:test-sdk:test"); const removeInterceptor = hook.intercept(async (target, thisArg, args) => String(await target.apply(thisArg, [String(args[0]).toUpperCase()])), ); @@ -586,7 +572,7 @@ describe("Orchestrion Transformation Tests", () => { "exports", result.code, )(loadedModule, loadedModule.exports); - const hook = newGlobalTracingChannel("orchestrion:test-sdk:test"); + const hook = newGlobalInvocationHook("orchestrion:test-sdk:test"); const removeInterceptor = hook.intercept((target, thisArg, args) => { const callback = args[1] as (error: unknown, value: string) => void; return target.apply(thisArg, [String(args[0]).toUpperCase(), callback]);