@shanepadgett/tau-agent 0.45.1 → 0.46.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/shared/events.ts CHANGED
@@ -81,19 +81,14 @@ interface TauEventAPI extends EmitEventAPI {
81
81
 
82
82
  type TauEventHandler<Name extends keyof TauAgentEvents> = (data: TauAgentEvents[Name]) => void | Promise<void>;
83
83
 
84
- interface TauEventSubscription {
85
- stop(): void;
86
- }
87
-
88
- type TauEventSubscriptionRegistry = WeakMap<
89
- ExtensionAPI["events"],
90
- Map<string, Map<keyof TauAgentEvents, TauEventSubscription>>
91
- >;
84
+ // `pi.events` is a fresh wrapper per extension, so subscriptions are matched through the shared bus itself.
85
+ // Each copy of this module (global and project installs) announces its subscription and stops any older one.
86
+ const CLAIM_CHANNEL = "tau:subscription.claim";
92
87
 
93
- // Global and project installs load separate module instances but share one pi.events bus.
94
- const registryKey = Symbol.for("tau-agent.eventSubscriptions");
95
- const registryHost = globalThis as typeof globalThis & { [registryKey]?: TauEventSubscriptionRegistry };
96
- const tauEventSubscriptions: TauEventSubscriptionRegistry = (registryHost[registryKey] ??= new WeakMap());
88
+ interface TauSubscriptionClaim {
89
+ owner: string;
90
+ name: string;
91
+ }
97
92
 
98
93
  export function emitTauEvent<Name extends keyof TauAgentEvents>(
99
94
  pi: EmitEventAPI,
@@ -130,25 +125,21 @@ function subscribeToTauEvent<Name extends keyof TauAgentEvents>(
130
125
  ): () => void {
131
126
  if (owner.length === 0) throw new Error("Tau event owner is required.");
132
127
 
133
- const subscriptions = getOwnerSubscriptions(pi.events, owner);
134
- subscriptions.get(name)?.stop();
135
-
136
128
  let unsubscribe: (() => void) | undefined;
137
129
  let disposed = false;
130
+ const claim: TauSubscriptionClaim = { owner, name };
138
131
 
139
132
  function detach(): void {
140
133
  unsubscribe?.();
141
134
  unsubscribe = undefined;
142
135
  }
143
136
 
144
- const subscription: TauEventSubscription = {
145
- stop() {
146
- if (disposed) return;
147
- disposed = true;
148
- detach();
149
- if (subscriptions.get(name) === subscription) subscriptions.delete(name);
150
- },
151
- };
137
+ function stop(): void {
138
+ if (disposed) return;
139
+ disposed = true;
140
+ detach();
141
+ stopClaimListener();
142
+ }
152
143
 
153
144
  function attach(): void {
154
145
  if (disposed) return;
@@ -156,32 +147,17 @@ function subscribeToTauEvent<Name extends keyof TauAgentEvents>(
156
147
  unsubscribe = pi.events.on(name, handler as (data: unknown) => void);
157
148
  }
158
149
 
159
- subscriptions.set(name, subscription);
150
+ pi.events.emit(CLAIM_CHANNEL, claim);
151
+ const stopClaimListener = pi.events.on(CLAIM_CHANNEL, (other) => {
152
+ const { owner: otherOwner, name: otherName } = other as TauSubscriptionClaim;
153
+ if (other !== claim && otherOwner === owner && otherName === name) stop();
154
+ });
160
155
  if (attachImmediately) attach();
161
156
  pi.on("session_start", attach);
162
157
  pi.on("session_shutdown", detach);
163
- return subscription.stop;
158
+ return stop;
164
159
  }
165
160
 
166
161
  export function setTauFooterItem(pi: EmitEventAPI, item: TauFooterItem): void {
167
162
  emitTauEvent(pi, "tau:footer-item", item);
168
163
  }
169
-
170
- function getOwnerSubscriptions(
171
- events: ExtensionAPI["events"],
172
- owner: string,
173
- ): Map<keyof TauAgentEvents, TauEventSubscription> {
174
- let busSubscriptions = tauEventSubscriptions.get(events);
175
- if (!busSubscriptions) {
176
- busSubscriptions = new Map();
177
- tauEventSubscriptions.set(events, busSubscriptions);
178
- }
179
-
180
- let ownerSubscriptions = busSubscriptions.get(owner);
181
- if (!ownerSubscriptions) {
182
- ownerSubscriptions = new Map();
183
- busSubscriptions.set(owner, ownerSubscriptions);
184
- }
185
-
186
- return ownerSubscriptions;
187
- }
@@ -24,6 +24,7 @@ interface GenerationContext {
24
24
  }
25
25
 
26
26
  interface ModelFallbackOptions {
27
+ sessionId?: string;
27
28
  statusKey?: string;
28
29
  notifyOnFallback?: boolean;
29
30
  maxAttempts?: number;
@@ -78,12 +79,13 @@ export async function generateValidated<T>(
78
79
  export async function generateToolValidated<T>(
79
80
  ctx: GenerationContext,
80
81
  candidates: readonly ModelCandidate[],
81
- prompt: string,
82
+ prompt: string | readonly Message[],
82
83
  tool: Tool,
83
84
  validate: (input: unknown) => T,
84
85
  correctionPrompt?: (error: Error, output: string) => string,
85
86
  options?: ModelFallbackOptions,
86
87
  ): Promise<{ value: T; candidate: ModelCandidate }> {
88
+ const sessionId = options?.sessionId ?? randomUUID();
87
89
  return withModelFallback(ctx, candidates, options, (candidate) =>
88
90
  requestToolValidated(
89
91
  ctx,
@@ -91,6 +93,7 @@ export async function generateToolValidated<T>(
91
93
  prompt,
92
94
  tool,
93
95
  validate,
96
+ sessionId,
94
97
  correctionPrompt,
95
98
  options?.maxAttempts ?? MAX_TOOL_ATTEMPTS,
96
99
  ),
@@ -207,14 +210,17 @@ function validateSingleToolCall(tool: Tool, toolCalls: readonly { name: string;
207
210
  async function requestToolValidated<T>(
208
211
  ctx: GenerationContext,
209
212
  candidate: ModelCandidate,
210
- prompt: string,
213
+ prompt: string | readonly Message[],
211
214
  tool: Tool,
212
215
  validate: (input: unknown) => T,
216
+ sessionId: string,
213
217
  correctionPrompt?: (error: Error, output: string) => string,
214
218
  maxAttempts = MAX_TOOL_ATTEMPTS,
215
219
  ): Promise<T> {
216
- const messages: Message[] = [{ role: "user", content: [{ type: "text", text: prompt }], timestamp: Date.now() }];
217
- const sessionId = randomUUID();
220
+ const messages: Message[] =
221
+ typeof prompt === "string"
222
+ ? [{ role: "user", content: [{ type: "text", text: prompt }], timestamp: Date.now() }]
223
+ : [...prompt];
218
224
 
219
225
  for (let attempt = 1; attempt <= maxAttempts; attempt++) {
220
226
  const response = await completeCandidate(ctx, candidate, messages, sessionId, [tool]);