




















@@ -1,4 +1,5 @@
11import type { StreamFn } from "@mariozechner/pi-agent-core";
2+import { fireAndForgetBoundedHook } from "../../../hooks/fire-and-forget.js";
23import {
34diagnosticErrorCategory,
45diagnosticProviderRequestIdHash,
@@ -12,6 +13,12 @@ import {
1213freezeDiagnosticTraceContext,
1314type DiagnosticTraceContext,
1415} from "../../../infra/diagnostic-trace-context.js";
16+import { getGlobalHookRunner } from "../../../plugins/hook-runner-global.js";
17+import type {
18+PluginHookAgentContext,
19+PluginHookModelCallEndedEvent,
20+PluginHookModelCallStartedEvent,
21+} from "../../../plugins/hook-types.js";
15221623export { diagnosticErrorCategory };
1724@@ -35,6 +42,10 @@ type ModelCallErrorFields = Pick<
3542Extract<DiagnosticEventInput, { type: "model.call.error" }>,
3643"errorCategory" | "upstreamRequestIdHash"
3744>;
45+type ModelCallEndedHookFields = Pick<
46+PluginHookModelCallEndedEvent,
47+"durationMs" | "outcome" | "errorCategory" | "upstreamRequestIdHash"
48+>;
38493950const MODEL_CALL_STREAM_RETURN_TIMEOUT_MS = 1000;
4051@@ -90,6 +101,102 @@ function modelCallErrorFields(err: unknown): ModelCallErrorFields {
90101};
91102}
92103104+function modelCallHookEventBase(eventBase: ModelCallEventBase): PluginHookModelCallStartedEvent {
105+return {
106+runId: eventBase.runId,
107+callId: eventBase.callId,
108+ ...(eventBase.sessionKey ? { sessionKey: eventBase.sessionKey } : {}),
109+ ...(eventBase.sessionId ? { sessionId: eventBase.sessionId } : {}),
110+provider: eventBase.provider,
111+model: eventBase.model,
112+ ...(eventBase.api ? { api: eventBase.api } : {}),
113+ ...(eventBase.transport ? { transport: eventBase.transport } : {}),
114+};
115+}
116+117+function modelCallHookContext(eventBase: ModelCallEventBase): PluginHookAgentContext {
118+return Object.freeze({
119+runId: eventBase.runId,
120+trace: eventBase.trace,
121+ ...(eventBase.sessionKey ? { sessionKey: eventBase.sessionKey } : {}),
122+ ...(eventBase.sessionId ? { sessionId: eventBase.sessionId } : {}),
123+modelProviderId: eventBase.provider,
124+modelId: eventBase.model,
125+}) as PluginHookAgentContext;
126+}
127+128+function dispatchModelCallStartedHook(eventBase: ModelCallEventBase): void {
129+const hookRunner = getGlobalHookRunner();
130+if (!hookRunner?.hasHooks("model_call_started")) {
131+return;
132+}
133+const event = Object.freeze(modelCallHookEventBase(eventBase)) as PluginHookModelCallStartedEvent;
134+const hookCtx = modelCallHookContext(eventBase);
135+fireAndForgetBoundedHook(
136+() => hookRunner.runModelCallStarted(event, hookCtx),
137+"model_call_started plugin hook failed",
138+);
139+}
140+141+function dispatchModelCallEndedHook(
142+eventBase: ModelCallEventBase,
143+fields: ModelCallEndedHookFields,
144+): void {
145+const hookRunner = getGlobalHookRunner();
146+if (!hookRunner?.hasHooks("model_call_ended")) {
147+return;
148+}
149+const event = Object.freeze({
150+ ...modelCallHookEventBase(eventBase),
151+ ...fields,
152+}) as PluginHookModelCallEndedEvent;
153+const hookCtx = modelCallHookContext(eventBase);
154+fireAndForgetBoundedHook(
155+() => hookRunner.runModelCallEnded(event, hookCtx),
156+"model_call_ended plugin hook failed",
157+);
158+}
159+160+function emitModelCallStarted(eventBase: ModelCallEventBase): void {
161+emitTrustedDiagnosticEvent({
162+type: "model.call.started",
163+ ...eventBase,
164+});
165+dispatchModelCallStartedHook(eventBase);
166+}
167+168+function emitModelCallCompleted(eventBase: ModelCallEventBase, startedAt: number): void {
169+const durationMs = Date.now() - startedAt;
170+emitTrustedDiagnosticEvent({
171+type: "model.call.completed",
172+ ...eventBase,
173+ durationMs,
174+});
175+dispatchModelCallEndedHook(eventBase, {
176+ durationMs,
177+outcome: "completed",
178+});
179+}
180+181+function emitModelCallError(
182+eventBase: ModelCallEventBase,
183+startedAt: number,
184+fields: ModelCallErrorFields,
185+): void {
186+const durationMs = Date.now() - startedAt;
187+emitTrustedDiagnosticEvent({
188+type: "model.call.error",
189+ ...eventBase,
190+ durationMs,
191+ ...fields,
192+});
193+dispatchModelCallEndedHook(eventBase, {
194+ durationMs,
195+outcome: "error",
196+ ...fields,
197+});
198+}
199+93200async function safeReturnIterator(iterator: AsyncIterator<unknown>): Promise<void> {
94201let returnResult: unknown;
95202try {
@@ -137,29 +244,15 @@ async function* observeModelCallIterator<T>(
137244yield next.value;
138245}
139246terminalEmitted = true;
140-emitTrustedDiagnosticEvent({
141-type: "model.call.completed",
142- ...eventBase,
143-durationMs: Date.now() - startedAt,
144-});
247+emitModelCallCompleted(eventBase, startedAt);
145248} catch (err) {
146249terminalEmitted = true;
147-emitTrustedDiagnosticEvent({
148-type: "model.call.error",
149- ...eventBase,
150-durationMs: Date.now() - startedAt,
151- ...modelCallErrorFields(err),
152-});
250+emitModelCallError(eventBase, startedAt, modelCallErrorFields(err));
153251throw err;
154252} finally {
155253if (!terminalEmitted) {
156254await safeReturnIterator(iterator);
157-emitTrustedDiagnosticEvent({
158-type: "model.call.error",
159- ...eventBase,
160-durationMs: Date.now() - startedAt,
161-errorCategory: "StreamAbandoned",
162-});
255+emitModelCallError(eventBase, startedAt, { errorCategory: "StreamAbandoned" });
163256}
164257}
165258}
@@ -209,11 +302,7 @@ function observeModelCallResult(
209302startedAt,
210303);
211304}
212-emitTrustedDiagnosticEvent({
213-type: "model.call.completed",
214- ...eventBase,
215-durationMs: Date.now() - startedAt,
216-});
305+emitModelCallCompleted(eventBase, startedAt);
217306return result;
218307}
219308@@ -225,10 +314,7 @@ export function wrapStreamFnWithDiagnosticModelCallEvents(
225314const callId = ctx.nextCallId();
226315const trace = freezeDiagnosticTraceContext(createChildDiagnosticTraceContext(ctx.trace));
227316const eventBase = baseModelCallEvent(ctx, callId, trace);
228-emitTrustedDiagnosticEvent({
229-type: "model.call.started",
230- ...eventBase,
231-});
317+emitModelCallStarted(eventBase);
232318const startedAt = Date.now();
233319234320try {
@@ -237,24 +323,14 @@ export function wrapStreamFnWithDiagnosticModelCallEvents(
237323return result.then(
238324(resolved) => observeModelCallResult(resolved, eventBase, startedAt),
239325(err) => {
240-emitTrustedDiagnosticEvent({
241-type: "model.call.error",
242- ...eventBase,
243-durationMs: Date.now() - startedAt,
244- ...modelCallErrorFields(err),
245-});
326+emitModelCallError(eventBase, startedAt, modelCallErrorFields(err));
246327throw err;
247328},
248329);
249330}
250331return observeModelCallResult(result, eventBase, startedAt);
251332} catch (err) {
252-emitTrustedDiagnosticEvent({
253-type: "model.call.error",
254- ...eventBase,
255-durationMs: Date.now() - startedAt,
256- ...modelCallErrorFields(err),
257-});
333+emitModelCallError(eventBase, startedAt, modelCallErrorFields(err));
258334throw err;
259335}
260336}) as StreamFn;
此内容由惯性聚合(RSS阅读器)自动聚合整理,仅供阅读参考。 原文来自 — 版权归原作者所有。