From 730516a9ffa6fc397a19d173ad566365cce8d67f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9D=8E=E5=B2=A9=E5=B2=A9?= Date: Tue, 4 Aug 2026 17:44:47 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=8B=86=E5=88=86=20Models=20Extension?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/extensions/catalog.ts | 4 +- src/extensions/chat.test.ts | 3 + src/extensions/shared/agent/index.ts | 33 ++++------ src/extensions/shared/deepseek/index.test.ts | 13 ++-- src/extensions/shared/deepseek/index.ts | 2 +- src/extensions/shared/models/index.test.ts | 60 +++++++++++++++++++ src/extensions/shared/models/index.ts | 63 ++++++++++++++++++++ src/products/shared.ts | 2 + 8 files changed, 148 insertions(+), 32 deletions(-) create mode 100644 src/extensions/shared/models/index.test.ts create mode 100644 src/extensions/shared/models/index.ts diff --git a/src/extensions/catalog.ts b/src/extensions/catalog.ts index 736302d..95a4c6c 100644 --- a/src/extensions/catalog.ts +++ b/src/extensions/catalog.ts @@ -1,5 +1,6 @@ export const ExtensionId = { Workspace: "workspace", + Models: "models", DeepSeek: "deepseek", Agent: "agent", Cli: "cli", @@ -7,8 +8,9 @@ export const ExtensionId = { export const Hook = { Workspace: "workspace", + Models: "models", Agent: "agent", - ModelProviders: "model.providers", + ModelProviders: "models.providers", }; export const Event = { diff --git a/src/extensions/chat.test.ts b/src/extensions/chat.test.ts index 4ba34c5..380fe7b 100644 --- a/src/extensions/chat.test.ts +++ b/src/extensions/chat.test.ts @@ -11,6 +11,7 @@ import { type AgentService, } from "./shared/agent"; import { createDeepSeekExtension } from "./shared/deepseek"; +import { createModelsExtension } from "./shared/models"; import { createWorkspaceExtension, type WorkspaceService, @@ -57,6 +58,7 @@ test("streams a reply, saves it, and restores the conversation after restart", a }; const firstKernel = new Kernel({ extensionConfigs }).use( createWorkspaceExtension, + createModelsExtension, createDeepSeekExtension, createAgentExtension, ); @@ -73,6 +75,7 @@ test("streams a reply, saves it, and restores the conversation after restart", a const secondKernel = new Kernel({ extensionConfigs }).use( createWorkspaceExtension, + createModelsExtension, createDeepSeekExtension, createAgentExtension, ); diff --git a/src/extensions/shared/agent/index.ts b/src/extensions/shared/agent/index.ts index 32eb633..0ca560d 100644 --- a/src/extensions/shared/agent/index.ts +++ b/src/extensions/shared/agent/index.ts @@ -5,15 +5,8 @@ import type { ExtensionRuntimeContext, } from "../../../kernel"; import { Event, ExtensionId, Hook } from "../../catalog"; -import { - type ChatMessage, - type WorkspaceService, -} from "../workspace"; - -export interface ModelProvider { - id: string; - chat(messages: ChatMessage[], signal?: AbortSignal): AsyncIterable; -} +import type { ModelService } from "../models"; +import type { WorkspaceService } from "../workspace"; export interface AgentService { chat(input: string, signal?: AbortSignal): AsyncIterable; @@ -22,17 +15,17 @@ export interface AgentService { export function createAgentExtension(): Extension { let runtime: ExtensionRuntimeContext | undefined; let workspace: WorkspaceService | undefined; - let model: ModelProvider | undefined; + let models: ModelService | undefined; const agent: AgentService = { async *chat(input, signal) { - if (!runtime || !workspace || !model) { + if (!runtime || !workspace || !models) { throw new Error("Agent is not running."); } const activeRuntime = runtime; const activeWorkspace = workspace; - const activeModel = model; + const activeModels = models; const runId = randomUUID(); const userMessage = await activeWorkspace.append("user", input); @@ -46,9 +39,11 @@ export function createAgentExtension(): Extension { let answer = ""; try { - const messages = await activeWorkspace.messages(); + const messages = (await activeWorkspace.messages()).map( + ({ role, content }) => ({ role, content }), + ); - for await (const chunk of activeModel.chat(messages, signal)) { + for await (const chunk of activeModels.chat(messages, signal)) { answer += chunk; yield chunk; } @@ -73,21 +68,15 @@ export function createAgentExtension(): Extension { }, start(context) { - const providers = context.all(Hook.ModelProviders); - - if (providers.length === 0) { - throw new Error("Agent needs at least one model provider."); - } - runtime = context; workspace = context.get(Hook.Workspace); - model = providers[0]; + models = context.get(Hook.Models); }, stop() { runtime = undefined; workspace = undefined; - model = undefined; + models = undefined; }, }; } diff --git a/src/extensions/shared/deepseek/index.test.ts b/src/extensions/shared/deepseek/index.test.ts index 06d365c..0ed1648 100644 --- a/src/extensions/shared/deepseek/index.test.ts +++ b/src/extensions/shared/deepseek/index.test.ts @@ -3,7 +3,7 @@ import test from "node:test"; import { Kernel } from "../../../kernel"; import { ExtensionId, Hook } from "../../catalog"; -import type { ModelProvider } from "../agent"; +import { createModelsExtension, type ModelService } from "../models"; import { createDeepSeekExtension } from "."; test("reads DeepSeek settings from extension config", async (t) => { @@ -38,15 +38,12 @@ test("reads DeepSeek settings from extension config", async (t) => { request, }, }, - }).use(createDeepSeekExtension); + }).use(createModelsExtension, createDeepSeekExtension); await kernel.start(); - const provider = kernel.all(Hook.ModelProviders)[0]; - assert.ok(provider); - - for await (const _ of provider.chat([ - { role: "user", content: "你好", createdAt: "2026-08-04T00:00:00.000Z" }, - ])) { + for await (const _ of kernel + .get(Hook.Models) + .chat([{ role: "user", content: "你好" }])) { // The test response contains no text chunks. } diff --git a/src/extensions/shared/deepseek/index.ts b/src/extensions/shared/deepseek/index.ts index a142642..38b0401 100644 --- a/src/extensions/shared/deepseek/index.ts +++ b/src/extensions/shared/deepseek/index.ts @@ -1,6 +1,6 @@ import type { Extension, ExtensionConfig } from "../../../kernel"; import { ExtensionId, Hook } from "../../catalog"; -import type { ModelProvider } from "../agent"; +import type { ModelProvider } from "../models"; export interface DeepSeekOptions { apiKey?: string; diff --git a/src/extensions/shared/models/index.test.ts b/src/extensions/shared/models/index.test.ts new file mode 100644 index 0000000..fab7b7f --- /dev/null +++ b/src/extensions/shared/models/index.test.ts @@ -0,0 +1,60 @@ +import assert from "node:assert/strict"; +import test from "node:test"; + +import { Kernel, type ExtensionSetupContext } from "../../../kernel"; +import { ExtensionId, Hook } from "../../catalog"; +import { + createModelsExtension, + type ModelProvider, + type ModelService, +} from "."; + +test("selects a provider and streams its response", async () => { + const requests: Array<{ role: string; content: string }[]> = []; + const ignoredProvider: ModelProvider = { + id: "ignored", + async *chat() { + yield "错误"; + }, + }; + const provider: ModelProvider = { + id: "test", + async *chat(messages) { + requests.push(messages); + yield "你"; + yield "好"; + }, + }; + function createProviderExtension() { + return { + setup(context: ExtensionSetupContext) { + context.add(Hook.ModelProviders, ignoredProvider); + context.add(Hook.ModelProviders, provider); + }, + }; + } + createProviderExtension.id = "test-model-provider"; + + const kernel = new Kernel({ + extensionConfigs: { + [ExtensionId.Models]: { defaultProvider: "test" }, + }, + }).use(createModelsExtension, createProviderExtension); + await kernel.start(); + + let answer = ""; + for await (const chunk of kernel + .get(Hook.Models) + .chat([{ role: "user", content: "你好" }])) { + answer += chunk; + } + + assert.equal(answer, "你好"); + assert.deepEqual(requests, [[{ role: "user", content: "你好" }]]); + await kernel.stop(); +}); + +test("requires at least one model provider", async () => { + const kernel = new Kernel().use(createModelsExtension); + await assert.rejects(kernel.start(), /at least one provider/); +}); diff --git a/src/extensions/shared/models/index.ts b/src/extensions/shared/models/index.ts new file mode 100644 index 0000000..ba1e014 --- /dev/null +++ b/src/extensions/shared/models/index.ts @@ -0,0 +1,63 @@ +import type { Extension, ExtensionConfig } from "../../../kernel"; +import { ExtensionId, Hook } from "../../catalog"; + +export interface ModelMessage { + role: "system" | "user" | "assistant"; + content: string; +} + +export interface ModelProvider { + id: string; + chat(messages: ModelMessage[], signal?: AbortSignal): AsyncIterable; +} + +export interface ModelService { + chat(messages: ModelMessage[], signal?: AbortSignal): AsyncIterable; +} + +export function createModelsExtension(options: ExtensionConfig = {}): Extension { + if ( + options.defaultProvider !== undefined && + typeof options.defaultProvider !== "string" + ) { + throw new Error("models.defaultProvider must be a string."); + } + + const defaultProvider = options.defaultProvider; + let provider: ModelProvider | undefined; + + const models: ModelService = { + chat(messages, signal) { + if (!provider) throw new Error("Models is not running."); + return provider.chat(messages, signal); + }, + }; + + return { + setup(context) { + context.add(Hook.Models, models); + }, + + start(context) { + const providers = context.all(Hook.ModelProviders); + + if (providers.length === 0) { + throw new Error("Models needs at least one provider."); + } + + provider = defaultProvider + ? providers.find(({ id }) => id === defaultProvider) + : providers[0]; + + if (!provider) { + throw new Error(`Model provider "${defaultProvider}" is not registered.`); + } + }, + + stop() { + provider = undefined; + }, + }; +} + +createModelsExtension.id = ExtensionId.Models; diff --git a/src/products/shared.ts b/src/products/shared.ts index 48e7a8d..284df79 100644 --- a/src/products/shared.ts +++ b/src/products/shared.ts @@ -1,10 +1,12 @@ import type { ExtensionFactory } from "../kernel"; import { createAgentExtension } from "../extensions/shared/agent"; import { createDeepSeekExtension } from "../extensions/shared/deepseek"; +import { createModelsExtension } from "../extensions/shared/models"; import { createWorkspaceExtension } from "../extensions/shared/workspace"; export const sharedExtensions: ExtensionFactory[] = [ createWorkspaceExtension, + createModelsExtension, createDeepSeekExtension, createAgentExtension, ];