Skip to content

Commit 897f4b5

Browse files
committed
[feat] Add streaming support for AI tasks
- Introduced a new base class, StreamingAiTask, to facilitate streaming output for AI tasks, enhancing the execution model for tasks like text generation, rewriting, and summarization. - Implemented streaming functions for existing tasks such as TextGenerationTask, TextRewriterTask, and TextSummaryTask, allowing for real-time output delivery. - Updated AI providers (Anthropic, Google Gemini, Ollama, and Hugging Face) to support streaming capabilities, improving responsiveness and user experience. - Enhanced the AiJob class with a new executeStream method to handle streaming execution, providing a fallback to non-streaming execution when necessary. - Refactored task registration to include streaming functions, ensuring seamless integration with the existing architecture.
1 parent fbceede commit 897f4b5

50 files changed

Lines changed: 4217 additions & 62 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

packages/ai-provider/src/anthropic/AnthropicProvider.ts

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
* SPDX-License-Identifier: Apache-2.0
55
*/
66

7-
import { AiProvider, type AiProviderRunFn } from "@workglow/ai";
7+
import { AiProvider, type AiProviderRunFn, type AiProviderStreamFn } from "@workglow/ai";
88
import { ANTHROPIC } from "./common/Anthropic_Constants";
99
import type { AnthropicModelConfig } from "./common/Anthropic_ModelSchema";
1010

@@ -39,7 +39,10 @@ export class AnthropicProvider extends AiProvider<AnthropicModelConfig> {
3939

4040
readonly taskTypes = ["TextGenerationTask", "TextRewriterTask", "TextSummaryTask"] as const;
4141

42-
constructor(tasks?: Record<string, AiProviderRunFn<any, any, AnthropicModelConfig>>) {
43-
super(tasks);
42+
constructor(
43+
tasks?: Record<string, AiProviderRunFn<any, any, AnthropicModelConfig>>,
44+
streamTasks?: Record<string, AiProviderStreamFn<any, any, AnthropicModelConfig>>
45+
) {
46+
super(tasks, streamTasks);
4447
}
4548
}

packages/ai-provider/src/anthropic/Anthropic_Worker.ts

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,12 +5,14 @@
55
*/
66

77
import { globalServiceRegistry, parentPort, WORKER_SERVER } from "@workglow/util";
8-
import { ANTHROPIC_TASKS } from "./common/Anthropic_JobRunFns";
98
import { AnthropicProvider } from "./AnthropicProvider";
9+
import { ANTHROPIC_STREAM_TASKS, ANTHROPIC_TASKS } from "./common/Anthropic_JobRunFns";
1010

1111
export function ANTHROPIC_WORKER_JOBRUN_REGISTER() {
1212
const workerServer = globalServiceRegistry.get(WORKER_SERVER);
13-
new AnthropicProvider(ANTHROPIC_TASKS).registerOnWorkerServer(workerServer);
13+
new AnthropicProvider(ANTHROPIC_TASKS, ANTHROPIC_STREAM_TASKS).registerOnWorkerServer(
14+
workerServer
15+
);
1416
parentPort.postMessage({ type: "ready" });
1517
console.log("ANTHROPIC_WORKER_JOBRUN registered");
1618
}

packages/ai-provider/src/anthropic/common/Anthropic_JobRunFns.ts

Lines changed: 98 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,13 +7,15 @@
77
import Anthropic from "@anthropic-ai/sdk";
88
import type {
99
AiProviderRunFn,
10+
AiProviderStreamFn,
1011
TextGenerationTaskInput,
1112
TextGenerationTaskOutput,
1213
TextRewriterTaskInput,
1314
TextRewriterTaskOutput,
1415
TextSummaryTaskInput,
1516
TextSummaryTaskOutput,
1617
} from "@workglow/ai";
18+
import type { StreamEvent } from "@workglow/task-graph";
1719
import type { AnthropicModelConfig } from "./Anthropic_ModelSchema";
1820

1921
function getClient(model: AnthropicModelConfig | undefined): Anthropic {
@@ -123,8 +125,104 @@ export const Anthropic_TextSummary: AiProviderRunFn<
123125
return { text };
124126
};
125127

128+
// ========================================================================
129+
// Streaming implementations (append mode)
130+
// ========================================================================
131+
132+
export const Anthropic_TextGeneration_Stream: AiProviderStreamFn<
133+
TextGenerationTaskInput,
134+
TextGenerationTaskOutput,
135+
AnthropicModelConfig
136+
> = async function* (input, model, signal): AsyncIterable<StreamEvent<TextGenerationTaskOutput>> {
137+
const client = getClient(model);
138+
const modelName = getModelName(model);
139+
140+
const stream = client.messages.stream(
141+
{
142+
model: modelName,
143+
messages: [{ role: "user", content: input.prompt }],
144+
max_tokens: getMaxTokens(input, model),
145+
temperature: input.temperature,
146+
top_p: input.topP,
147+
},
148+
{ signal }
149+
);
150+
151+
for await (const event of stream) {
152+
if (event.type === "content_block_delta" && event.delta.type === "text_delta") {
153+
yield { type: "text-delta", textDelta: event.delta.text };
154+
}
155+
}
156+
yield { type: "finish", data: {} as TextGenerationTaskOutput };
157+
};
158+
159+
export const Anthropic_TextRewriter_Stream: AiProviderStreamFn<
160+
TextRewriterTaskInput,
161+
TextRewriterTaskOutput,
162+
AnthropicModelConfig
163+
> = async function* (input, model, signal): AsyncIterable<StreamEvent<TextRewriterTaskOutput>> {
164+
const client = getClient(model);
165+
const modelName = getModelName(model);
166+
167+
const stream = client.messages.stream(
168+
{
169+
model: modelName,
170+
system: input.prompt,
171+
messages: [{ role: "user", content: input.text }],
172+
max_tokens: getMaxTokens({}, model),
173+
},
174+
{ signal }
175+
);
176+
177+
for await (const event of stream) {
178+
if (event.type === "content_block_delta" && event.delta.type === "text_delta") {
179+
yield { type: "text-delta", textDelta: event.delta.text };
180+
}
181+
}
182+
yield { type: "finish", data: {} as TextRewriterTaskOutput };
183+
};
184+
185+
export const Anthropic_TextSummary_Stream: AiProviderStreamFn<
186+
TextSummaryTaskInput,
187+
TextSummaryTaskOutput,
188+
AnthropicModelConfig
189+
> = async function* (input, model, signal): AsyncIterable<StreamEvent<TextSummaryTaskOutput>> {
190+
const client = getClient(model);
191+
const modelName = getModelName(model);
192+
193+
const stream = client.messages.stream(
194+
{
195+
model: modelName,
196+
system: "Summarize the following text concisely.",
197+
messages: [{ role: "user", content: input.text }],
198+
max_tokens: getMaxTokens({}, model),
199+
},
200+
{ signal }
201+
);
202+
203+
for await (const event of stream) {
204+
if (event.type === "content_block_delta" && event.delta.type === "text_delta") {
205+
yield { type: "text-delta", textDelta: event.delta.text };
206+
}
207+
}
208+
yield { type: "finish", data: {} as TextSummaryTaskOutput };
209+
};
210+
211+
// ========================================================================
212+
// Task registries
213+
// ========================================================================
214+
126215
export const ANTHROPIC_TASKS: Record<string, AiProviderRunFn<any, any, AnthropicModelConfig>> = {
127216
TextGenerationTask: Anthropic_TextGeneration,
128217
TextRewriterTask: Anthropic_TextRewriter,
129218
TextSummaryTask: Anthropic_TextSummary,
130219
};
220+
221+
export const ANTHROPIC_STREAM_TASKS: Record<
222+
string,
223+
AiProviderStreamFn<any, any, AnthropicModelConfig>
224+
> = {
225+
TextGenerationTask: Anthropic_TextGeneration_Stream,
226+
TextRewriterTask: Anthropic_TextRewriter_Stream,
227+
TextSummaryTask: Anthropic_TextSummary_Stream,
228+
};

packages/ai-provider/src/google-gemini/Gemini_Worker.ts

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,12 +5,12 @@
55
*/
66

77
import { globalServiceRegistry, parentPort, WORKER_SERVER } from "@workglow/util";
8-
import { GEMINI_TASKS } from "./common/Gemini_JobRunFns";
8+
import { GEMINI_STREAM_TASKS, GEMINI_TASKS } from "./common/Gemini_JobRunFns";
99
import { GoogleGeminiProvider } from "./GoogleGeminiProvider";
1010

1111
export function GEMINI_WORKER_JOBRUN_REGISTER() {
1212
const workerServer = globalServiceRegistry.get(WORKER_SERVER);
13-
new GoogleGeminiProvider(GEMINI_TASKS).registerOnWorkerServer(workerServer);
13+
new GoogleGeminiProvider(GEMINI_TASKS, GEMINI_STREAM_TASKS).registerOnWorkerServer(workerServer);
1414
parentPort.postMessage({ type: "ready" });
1515
console.log("GEMINI_WORKER_JOBRUN registered");
1616
}

packages/ai-provider/src/google-gemini/GoogleGeminiProvider.ts

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
* SPDX-License-Identifier: Apache-2.0
55
*/
66

7-
import { AiProvider, type AiProviderRunFn } from "@workglow/ai";
7+
import { AiProvider, type AiProviderRunFn, type AiProviderStreamFn } from "@workglow/ai";
88
import { GOOGLE_GEMINI } from "./common/Gemini_Constants";
99
import type { GeminiModelConfig } from "./common/Gemini_ModelSchema";
1010

@@ -41,7 +41,10 @@ export class GoogleGeminiProvider extends AiProvider<GeminiModelConfig> {
4141
"TextSummaryTask",
4242
] as const;
4343

44-
constructor(tasks?: Record<string, AiProviderRunFn<any, any, GeminiModelConfig>>) {
45-
super(tasks);
44+
constructor(
45+
tasks?: Record<string, AiProviderRunFn<any, any, GeminiModelConfig>>,
46+
streamTasks?: Record<string, AiProviderStreamFn<any, any, GeminiModelConfig>>
47+
) {
48+
super(tasks, streamTasks);
4649
}
4750
}

packages/ai-provider/src/google-gemini/common/Gemini_JobRunFns.ts

Lines changed: 95 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import { GoogleGenerativeAI, type TaskType } from "@google/generative-ai";
88
import type {
99
AiProviderRunFn,
10+
AiProviderStreamFn,
1011
TextEmbeddingTaskInput,
1112
TextEmbeddingTaskOutput,
1213
TextGenerationTaskInput,
@@ -16,6 +17,7 @@ import type {
1617
TextSummaryTaskInput,
1718
TextSummaryTaskOutput,
1819
} from "@workglow/ai";
20+
import type { StreamEvent } from "@workglow/task-graph";
1921
import type { GeminiModelConfig } from "./Gemini_ModelSchema";
2022

2123
function getApiKey(model: GeminiModelConfig | undefined): string {
@@ -143,9 +145,102 @@ export const Gemini_TextSummary: AiProviderRunFn<
143145
return { text };
144146
};
145147

148+
// ========================================================================
149+
// Streaming implementations (append mode)
150+
// ========================================================================
151+
152+
export const Gemini_TextGeneration_Stream: AiProviderStreamFn<
153+
TextGenerationTaskInput,
154+
TextGenerationTaskOutput,
155+
GeminiModelConfig
156+
> = async function* (input, model, signal): AsyncIterable<StreamEvent<TextGenerationTaskOutput>> {
157+
const genAI = new GoogleGenerativeAI(getApiKey(model));
158+
const genModel = genAI.getGenerativeModel({
159+
model: getModelName(model),
160+
generationConfig: {
161+
maxOutputTokens: input.maxTokens,
162+
temperature: input.temperature,
163+
topP: input.topP,
164+
},
165+
});
166+
167+
const result = await genModel.generateContentStream({
168+
contents: [{ role: "user", parts: [{ text: input.prompt }] }],
169+
});
170+
171+
for await (const chunk of result.stream) {
172+
const text = chunk.text();
173+
if (text) {
174+
yield { type: "text-delta", textDelta: text };
175+
}
176+
}
177+
yield { type: "finish", data: {} as TextGenerationTaskOutput };
178+
};
179+
180+
export const Gemini_TextRewriter_Stream: AiProviderStreamFn<
181+
TextRewriterTaskInput,
182+
TextRewriterTaskOutput,
183+
GeminiModelConfig
184+
> = async function* (input, model, signal): AsyncIterable<StreamEvent<TextRewriterTaskOutput>> {
185+
const genAI = new GoogleGenerativeAI(getApiKey(model));
186+
const genModel = genAI.getGenerativeModel({
187+
model: getModelName(model),
188+
systemInstruction: input.prompt,
189+
});
190+
191+
const result = await genModel.generateContentStream({
192+
contents: [{ role: "user", parts: [{ text: input.text }] }],
193+
});
194+
195+
for await (const chunk of result.stream) {
196+
const text = chunk.text();
197+
if (text) {
198+
yield { type: "text-delta", textDelta: text };
199+
}
200+
}
201+
yield { type: "finish", data: {} as TextRewriterTaskOutput };
202+
};
203+
204+
export const Gemini_TextSummary_Stream: AiProviderStreamFn<
205+
TextSummaryTaskInput,
206+
TextSummaryTaskOutput,
207+
GeminiModelConfig
208+
> = async function* (input, model, signal): AsyncIterable<StreamEvent<TextSummaryTaskOutput>> {
209+
const genAI = new GoogleGenerativeAI(getApiKey(model));
210+
const genModel = genAI.getGenerativeModel({
211+
model: getModelName(model),
212+
systemInstruction: "Summarize the following text concisely.",
213+
});
214+
215+
const result = await genModel.generateContentStream({
216+
contents: [{ role: "user", parts: [{ text: input.text }] }],
217+
});
218+
219+
for await (const chunk of result.stream) {
220+
const text = chunk.text();
221+
if (text) {
222+
yield { type: "text-delta", textDelta: text };
223+
}
224+
}
225+
yield { type: "finish", data: {} as TextSummaryTaskOutput };
226+
};
227+
228+
// ========================================================================
229+
// Task registries
230+
// ========================================================================
231+
146232
export const GEMINI_TASKS: Record<string, AiProviderRunFn<any, any, GeminiModelConfig>> = {
147233
TextGenerationTask: Gemini_TextGeneration,
148234
TextEmbeddingTask: Gemini_TextEmbedding,
149235
TextRewriterTask: Gemini_TextRewriter,
150236
TextSummaryTask: Gemini_TextSummary,
151237
};
238+
239+
export const GEMINI_STREAM_TASKS: Record<
240+
string,
241+
AiProviderStreamFn<any, any, GeminiModelConfig>
242+
> = {
243+
TextGenerationTask: Gemini_TextGeneration_Stream,
244+
TextRewriterTask: Gemini_TextRewriter_Stream,
245+
TextSummaryTask: Gemini_TextSummary_Stream,
246+
};

packages/ai-provider/src/hf-transformers/HFT_Worker.ts

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,12 +5,14 @@
55
*/
66

77
import { globalServiceRegistry, parentPort, WORKER_SERVER } from "@workglow/util";
8-
import { HFT_TASKS } from "./common/HFT_JobRunFns";
8+
import { HFT_STREAM_TASKS, HFT_TASKS } from "./common/HFT_JobRunFns";
99
import { HuggingFaceTransformersProvider } from "./HuggingFaceTransformersProvider";
1010

1111
export function HFT_WORKER_JOBRUN_REGISTER() {
1212
const workerServer = globalServiceRegistry.get(WORKER_SERVER);
13-
new HuggingFaceTransformersProvider(HFT_TASKS).registerOnWorkerServer(workerServer);
13+
new HuggingFaceTransformersProvider(HFT_TASKS, HFT_STREAM_TASKS).registerOnWorkerServer(
14+
workerServer
15+
);
1416
parentPort.postMessage({ type: "ready" });
1517
console.log("HFT_WORKER_JOBRUN registered");
1618
}

packages/ai-provider/src/hf-transformers/HuggingFaceTransformersProvider.ts

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,12 @@
44
* SPDX-License-Identifier: Apache-2.0
55
*/
66

7-
import { AiProvider, type AiProviderRegisterOptions, type AiProviderRunFn } from "@workglow/ai";
7+
import {
8+
AiProvider,
9+
type AiProviderRegisterOptions,
10+
type AiProviderRunFn,
11+
type AiProviderStreamFn,
12+
} from "@workglow/ai";
813
import { HF_TRANSFORMERS_ONNX } from "./common/HFT_Constants";
914
import type { HfTransformersOnnxModelConfig } from "./common/HFT_ModelSchema";
1015

@@ -58,8 +63,11 @@ export class HuggingFaceTransformersProvider extends AiProvider<HfTransformersOn
5863
"ObjectDetectionTask",
5964
] as const;
6065

61-
constructor(tasks?: Record<string, AiProviderRunFn<any, any, HfTransformersOnnxModelConfig>>) {
62-
super(tasks);
66+
constructor(
67+
tasks?: Record<string, AiProviderRunFn<any, any, HfTransformersOnnxModelConfig>>,
68+
streamTasks?: Record<string, AiProviderStreamFn<any, any, HfTransformersOnnxModelConfig>>
69+
) {
70+
super(tasks, streamTasks);
6371
}
6472

6573
protected override async onInitialize(options: AiProviderRegisterOptions): Promise<void> {

0 commit comments

Comments
 (0)