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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/calm-facets-reply.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
"agents": patch
---

Route asynchronous callable and streaming responses through the facet WebSocket frame that originated each RPC.
165 changes: 149 additions & 16 deletions packages/agents/src/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -283,6 +283,92 @@ function sendRpcResponseIfOpen(
}
}

type RPCReplyTarget = {
send(message: string | ArrayBuffer | ArrayBufferView): void | Promise<void>;
};

type FacetRPCResponseDelivery = {
sent: boolean;
completion: Promise<void>;
};

type SubAgentRpcReplyInvocationContext = {
bridge?: SubAgentConnectionBridge;
};

const subAgentRpcReplyContext =
new AsyncLocalStorage<SubAgentRpcReplyInvocationContext>();

function sendFacetRpcResponseIfOpen(
target: RPCReplyTarget,
response: RPCResponse
): FacetRPCResponseDelivery {
try {
const completion = Promise.resolve(
target.send(JSON.stringify(response))
).catch((error: unknown) => {
if (!isClosedWebSocketSendError(error)) {
console.error("[Agent] Facet RPC response delivery failed:", error);
}
});
return { sent: true, completion };
} catch (error) {
if (isClosedWebSocketSendError(error)) {
return { sent: false, completion: Promise.resolve() };
}
throw error;
}
}

type FacetStreamingResponseDeliveryState = {
replyTarget: RPCReplyTarget;
pending: Set<Promise<void>>;
};

const facetStreamingResponseDeliveryStates = new WeakMap<
StreamingResponse,
FacetStreamingResponseDeliveryState
>();

function createStreamingResponse(
connection: Connection,
id: string,
facetReplyTarget?: RPCReplyTarget
): StreamingResponse {
const stream = new StreamingResponse(connection, id);
if (facetReplyTarget) {
facetStreamingResponseDeliveryStates.set(stream, {
replyTarget: facetReplyTarget,
pending: new Set()
});
}
return stream;
}

function trackFacetStreamingResponseDelivery(
stream: StreamingResponse,
completion: Promise<void>
): void {
const state = facetStreamingResponseDeliveryStates.get(stream);
if (!state) return;

state.pending.add(completion);
void completion.finally(() => state.pending.delete(completion));
}

async function waitForFacetStreamingResponseDeliveries(
stream: StreamingResponse
): Promise<void> {
const state = facetStreamingResponseDeliveryStates.get(stream);
if (!state) return;

try {
await Promise.all(state.pending);
} finally {
facetStreamingResponseDeliveryStates.delete(stream);
}
}

/**
* Type guard for RPC request messages
*/
Expand Down Expand Up @@ -408,7 +494,8 @@ type SubAgentWebSocketEndpoint = {
_cf_handleSubAgentWebSocketMessage(
message: WSMessage,
bridge: SubAgentConnectionBridge,
meta: SubAgentConnectionMeta
meta: SubAgentConnectionMeta,
replyBridge?: SubAgentConnectionBridge
): Promise<void>;
_cf_handleSubAgentWebSocketClose(
code: number,
Expand Down Expand Up @@ -2453,7 +2540,14 @@ export class Agent<

const _onMessage = this.onMessage.bind(this);
this.onMessage = async (connection: Connection, message: WSMessage) => {
if (await this._cf_forwardSubAgentWebSocketMessage(connection, message)) {
const replyBridge = subAgentRpcReplyContext.getStore()?.bridge;
if (
await this._cf_forwardSubAgentWebSocketMessage(
connection,
message,
replyBridge
)
) {
return;
}
this._ensureConnectionWrapped(connection);
Expand Down Expand Up @@ -2518,7 +2612,11 @@ export class Agent<

// For streaming methods, pass a StreamingResponse object
if (metadata?.streaming) {
const stream = new StreamingResponse(connection, id);
const stream = createStreamingResponse(
connection,
id,
replyBridge
);

this._emit("rpc", { method, streaming: true });

Expand All @@ -2537,6 +2635,7 @@ export class Agent<
);
}
}
await waitForFacetStreamingResponseDeliveries(stream);
return;
}

Expand All @@ -2552,17 +2651,27 @@ export class Agent<
success: true,
type: MessageType.RPC
};
sendRpcResponseIfOpen(connection, response);
if (replyBridge) {
await sendFacetRpcResponseIfOpen(replyBridge, response)
.completion;
} else {
sendRpcResponseIfOpen(connection, response);
}
} catch (e) {
// Send error response
const response: RPCResponse = {
error:
e instanceof Error ? e.message : "Unknown error occurred",
id: parsed.id,
success: false,
type: MessageType.RPC
};
sendRpcResponseIfOpen(connection, response);
if (replyBridge) {
await sendFacetRpcResponseIfOpen(replyBridge, response)
.completion;
} else {
sendRpcResponseIfOpen(connection, response);
}

console.error("RPC error:", e);
this._emit("rpc:error", {
method: parsed.method,
Expand Down Expand Up @@ -7337,15 +7446,18 @@ export class Agent<

private async _cf_forwardSubAgentWebSocketMessage(
connection: Connection,
message: WSMessage
message: WSMessage,
replyBridge?: SubAgentConnectionBridge
): Promise<boolean> {
const routed = await this._cf_resolveSubAgentConnection(connection);
if (!routed) return false;

const bridge = this._cf_createSubAgentConnectionBridge(connection);
await routed.child._cf_handleSubAgentWebSocketMessage(
message,
this._cf_createSubAgentConnectionBridge(connection),
routed.meta
bridge,
routed.meta,
replyBridge ?? bridge
);
return true;
}
Expand Down Expand Up @@ -7481,13 +7593,23 @@ export class Agent<
async _cf_handleSubAgentWebSocketMessage(
message: WSMessage,
bridge: SubAgentConnectionBridge,
meta: SubAgentConnectionMeta
meta: SubAgentConnectionMeta,
replyBridge: SubAgentConnectionBridge = bridge
): Promise<void> {
const connection = this._cf_createSubAgentBridgeConnection(bridge, meta);
this._cf_storeVirtualSubAgentConnection(bridge, connection);
await this._cf_runWithSubAgentBridge(bridge, () =>
this.onMessage(connection, message)
);
const replyContext: SubAgentRpcReplyInvocationContext = {
bridge: replyBridge
};
try {
await subAgentRpcReplyContext.run(replyContext, () =>
this._cf_runWithSubAgentBridge(bridge, () =>
this.onMessage(connection, message)
)
);
} finally {
replyContext.bridge = undefined;
}
}

async _cf_handleSubAgentWebSocketClose(
Expand Down Expand Up @@ -13164,6 +13286,17 @@ export class StreamingResponse {
this._id = id;
}

private _send(response: RPCResponse): boolean {
const state = facetStreamingResponseDeliveryStates.get(this);
if (!state) {
return sendRpcResponseIfOpen(this._connection, response);
}

const delivery = sendFacetRpcResponseIfOpen(state.replyTarget, response);
trackFacetStreamingResponseDelivery(this, delivery.completion);
return delivery.sent;
}

/**
* Whether the stream has been closed (via end() or error())
*/
Expand All @@ -13190,7 +13323,7 @@ export class StreamingResponse {
success: true,
type: MessageType.RPC
};
return sendRpcResponseIfOpen(this._connection, response);
return this._send(response);
}

/**
Expand All @@ -13210,7 +13343,7 @@ export class StreamingResponse {
success: true,
type: MessageType.RPC
};
return sendRpcResponseIfOpen(this._connection, response);
return this._send(response);
}

/**
Expand All @@ -13229,6 +13362,6 @@ export class StreamingResponse {
success: false,
type: MessageType.RPC
};
return sendRpcResponseIfOpen(this._connection, response);
return this._send(response);
}
}
1 change: 1 addition & 0 deletions packages/agents/src/tests/agents/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,7 @@ export {
LeafSubAgent,
CallbackSubAgent,
BroadcastSubAgent,
SlowReplySubAgent,
HookingSubAgentParent,
Sub,
SUB,
Expand Down
50 changes: 48 additions & 2 deletions packages/agents/src/tests/agents/sub-agent.ts
Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
import { Agent, getCurrentAgent } from "../../index.ts";
import { Agent, callable, getCurrentAgent } from "../../index.ts";
import type {
FiberInspection,
FiberRecoveryContext,
FiberRecoveryResult
FiberRecoveryResult,
StreamingResponse
} from "../../index.ts";
import { RpcTarget } from "cloudflare:workers";

Expand Down Expand Up @@ -1082,6 +1083,11 @@ export class CustomBoundSubAgentParent extends Agent {
// ── Parent Agent that manages sub-agents ────────────────────────────

export class TestSubAgentParent extends Agent {
async delayedEchoFromParent(value: string): Promise<string> {
await new Promise((resolve) => setTimeout(resolve, 150));
return `parent:${value}`;
}

async onMessage(
connection: { send(message: string): void },
message: string | ArrayBuffer
Expand Down Expand Up @@ -2211,6 +2217,46 @@ class _UnboundParent extends Agent {
}
export { _UnboundParent as TestUnboundParentAgent };

// Regression fixture for issue #1991. The onMessage wrapper is intentional:
// facet RPC replies must retain their originating bridge through application
// and framework middleware before the Agent protocol dispatcher handles them.
export class SlowReplySubAgent extends Agent {
onStart(): void {
const handleMessage = this.onMessage.bind(this);
this.onMessage = async (connection, message) => {
await new Promise((resolve) => setTimeout(resolve, 25));
await handleMessage(connection, message);
};
}

@callable()
async slowEcho(value: string): Promise<string> {
await new Promise((resolve) => setTimeout(resolve, 150));
return `slow:${value}`;
}

@callable()
fastEcho(value: string): string {
return `fast:${value}`;
}

@callable()
async parentEcho(value: string): Promise<string> {
const parent = await this.parentAgent(TestSubAgentParent);
return await parent.delayedEchoFromParent(value);
}

@callable({ streaming: true })
async slowStreamingEcho(
stream: StreamingResponse,
value: string
): Promise<void> {
await new Promise((resolve) => setTimeout(resolve, 150));
stream.send(`slow-stream:${value}:chunk`);
stream.end(`slow-stream:${value}:done`);
}
}

/** Class identifier `_a`, exported as `TestMinifiedNameParentAgent`. */
class _a extends Agent {
async tryToSpawn(name: string): Promise<string> {
Expand Down
Loading
Loading