feat: 拆分 Models Extension

This commit is contained in:
李岩岩 2026-08-04 17:44:47 +08:00 committed by liyy
parent 88d9a8f5d6
commit 730516a9ff
8 changed files with 148 additions and 32 deletions

View File

@ -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 = {

View File

@ -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,
);

View File

@ -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<string>;
}
import type { ModelService } from "../models";
import type { WorkspaceService } from "../workspace";
export interface AgentService {
chat(input: string, signal?: AbortSignal): AsyncIterable<string>;
@ -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<ModelProvider>(Hook.ModelProviders);
if (providers.length === 0) {
throw new Error("Agent needs at least one model provider.");
}
runtime = context;
workspace = context.get<WorkspaceService>(Hook.Workspace);
model = providers[0];
models = context.get<ModelService>(Hook.Models);
},
stop() {
runtime = undefined;
workspace = undefined;
model = undefined;
models = undefined;
},
};
}

View File

@ -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<ModelProvider>(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<ModelService>(Hook.Models)
.chat([{ role: "user", content: "你好" }])) {
// The test response contains no text chunks.
}

View File

@ -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;

View File

@ -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<ModelService>(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/);
});

View File

@ -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<string>;
}
export interface ModelService {
chat(messages: ModelMessage[], signal?: AbortSignal): AsyncIterable<string>;
}
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<ModelProvider>(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;

View File

@ -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,
];