feat: 拆分 Models Extension
This commit is contained in:
parent
59903d6d5e
commit
1644731a5e
@ -1,5 +1,6 @@
|
|||||||
export const ExtensionId = {
|
export const ExtensionId = {
|
||||||
Workspace: "workspace",
|
Workspace: "workspace",
|
||||||
|
Models: "models",
|
||||||
DeepSeek: "deepseek",
|
DeepSeek: "deepseek",
|
||||||
Agent: "agent",
|
Agent: "agent",
|
||||||
Cli: "cli",
|
Cli: "cli",
|
||||||
@ -7,8 +8,9 @@ export const ExtensionId = {
|
|||||||
|
|
||||||
export const Hook = {
|
export const Hook = {
|
||||||
Workspace: "workspace",
|
Workspace: "workspace",
|
||||||
|
Models: "models",
|
||||||
Agent: "agent",
|
Agent: "agent",
|
||||||
ModelProviders: "model.providers",
|
ModelProviders: "models.providers",
|
||||||
};
|
};
|
||||||
|
|
||||||
export const Event = {
|
export const Event = {
|
||||||
|
|||||||
@ -11,6 +11,7 @@ import {
|
|||||||
type AgentService,
|
type AgentService,
|
||||||
} from "./shared/agent";
|
} from "./shared/agent";
|
||||||
import { createDeepSeekExtension } from "./shared/deepseek";
|
import { createDeepSeekExtension } from "./shared/deepseek";
|
||||||
|
import { createModelsExtension } from "./shared/models";
|
||||||
import {
|
import {
|
||||||
createWorkspaceExtension,
|
createWorkspaceExtension,
|
||||||
type WorkspaceService,
|
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(
|
const firstKernel = new Kernel({ extensionConfigs }).use(
|
||||||
createWorkspaceExtension,
|
createWorkspaceExtension,
|
||||||
|
createModelsExtension,
|
||||||
createDeepSeekExtension,
|
createDeepSeekExtension,
|
||||||
createAgentExtension,
|
createAgentExtension,
|
||||||
);
|
);
|
||||||
@ -73,6 +75,7 @@ test("streams a reply, saves it, and restores the conversation after restart", a
|
|||||||
|
|
||||||
const secondKernel = new Kernel({ extensionConfigs }).use(
|
const secondKernel = new Kernel({ extensionConfigs }).use(
|
||||||
createWorkspaceExtension,
|
createWorkspaceExtension,
|
||||||
|
createModelsExtension,
|
||||||
createDeepSeekExtension,
|
createDeepSeekExtension,
|
||||||
createAgentExtension,
|
createAgentExtension,
|
||||||
);
|
);
|
||||||
|
|||||||
@ -5,15 +5,8 @@ import type {
|
|||||||
ExtensionRuntimeContext,
|
ExtensionRuntimeContext,
|
||||||
} from "../../../kernel";
|
} from "../../../kernel";
|
||||||
import { Event, ExtensionId, Hook } from "../../catalog";
|
import { Event, ExtensionId, Hook } from "../../catalog";
|
||||||
import {
|
import type { ModelService } from "../models";
|
||||||
type ChatMessage,
|
import type { WorkspaceService } from "../workspace";
|
||||||
type WorkspaceService,
|
|
||||||
} from "../workspace";
|
|
||||||
|
|
||||||
export interface ModelProvider {
|
|
||||||
id: string;
|
|
||||||
chat(messages: ChatMessage[], signal?: AbortSignal): AsyncIterable<string>;
|
|
||||||
}
|
|
||||||
|
|
||||||
export interface AgentService {
|
export interface AgentService {
|
||||||
chat(input: string, signal?: AbortSignal): AsyncIterable<string>;
|
chat(input: string, signal?: AbortSignal): AsyncIterable<string>;
|
||||||
@ -22,17 +15,17 @@ export interface AgentService {
|
|||||||
export function createAgentExtension(): Extension {
|
export function createAgentExtension(): Extension {
|
||||||
let runtime: ExtensionRuntimeContext | undefined;
|
let runtime: ExtensionRuntimeContext | undefined;
|
||||||
let workspace: WorkspaceService | undefined;
|
let workspace: WorkspaceService | undefined;
|
||||||
let model: ModelProvider | undefined;
|
let models: ModelService | undefined;
|
||||||
|
|
||||||
const agent: AgentService = {
|
const agent: AgentService = {
|
||||||
async *chat(input, signal) {
|
async *chat(input, signal) {
|
||||||
if (!runtime || !workspace || !model) {
|
if (!runtime || !workspace || !models) {
|
||||||
throw new Error("Agent is not running.");
|
throw new Error("Agent is not running.");
|
||||||
}
|
}
|
||||||
|
|
||||||
const activeRuntime = runtime;
|
const activeRuntime = runtime;
|
||||||
const activeWorkspace = workspace;
|
const activeWorkspace = workspace;
|
||||||
const activeModel = model;
|
const activeModels = models;
|
||||||
const runId = randomUUID();
|
const runId = randomUUID();
|
||||||
|
|
||||||
const userMessage = await activeWorkspace.append("user", input);
|
const userMessage = await activeWorkspace.append("user", input);
|
||||||
@ -46,9 +39,11 @@ export function createAgentExtension(): Extension {
|
|||||||
let answer = "";
|
let answer = "";
|
||||||
|
|
||||||
try {
|
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;
|
answer += chunk;
|
||||||
yield chunk;
|
yield chunk;
|
||||||
}
|
}
|
||||||
@ -73,21 +68,15 @@ export function createAgentExtension(): Extension {
|
|||||||
},
|
},
|
||||||
|
|
||||||
start(context) {
|
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;
|
runtime = context;
|
||||||
workspace = context.get<WorkspaceService>(Hook.Workspace);
|
workspace = context.get<WorkspaceService>(Hook.Workspace);
|
||||||
model = providers[0];
|
models = context.get<ModelService>(Hook.Models);
|
||||||
},
|
},
|
||||||
|
|
||||||
stop() {
|
stop() {
|
||||||
runtime = undefined;
|
runtime = undefined;
|
||||||
workspace = undefined;
|
workspace = undefined;
|
||||||
model = undefined;
|
models = undefined;
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|||||||
@ -3,7 +3,7 @@ import test from "node:test";
|
|||||||
|
|
||||||
import { Kernel } from "../../../kernel";
|
import { Kernel } from "../../../kernel";
|
||||||
import { ExtensionId, Hook } from "../../catalog";
|
import { ExtensionId, Hook } from "../../catalog";
|
||||||
import type { ModelProvider } from "../agent";
|
import { createModelsExtension, type ModelService } from "../models";
|
||||||
import { createDeepSeekExtension } from ".";
|
import { createDeepSeekExtension } from ".";
|
||||||
|
|
||||||
test("reads DeepSeek settings from extension config", async (t) => {
|
test("reads DeepSeek settings from extension config", async (t) => {
|
||||||
@ -38,15 +38,12 @@ test("reads DeepSeek settings from extension config", async (t) => {
|
|||||||
request,
|
request,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}).use(createDeepSeekExtension);
|
}).use(createModelsExtension, createDeepSeekExtension);
|
||||||
await kernel.start();
|
await kernel.start();
|
||||||
|
|
||||||
const provider = kernel.all<ModelProvider>(Hook.ModelProviders)[0];
|
for await (const _ of kernel
|
||||||
assert.ok(provider);
|
.get<ModelService>(Hook.Models)
|
||||||
|
.chat([{ role: "user", content: "你好" }])) {
|
||||||
for await (const _ of provider.chat([
|
|
||||||
{ role: "user", content: "你好", createdAt: "2026-08-04T00:00:00.000Z" },
|
|
||||||
])) {
|
|
||||||
// The test response contains no text chunks.
|
// The test response contains no text chunks.
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@ -1,6 +1,6 @@
|
|||||||
import type { Extension, ExtensionConfig } from "../../../kernel";
|
import type { Extension, ExtensionConfig } from "../../../kernel";
|
||||||
import { ExtensionId, Hook } from "../../catalog";
|
import { ExtensionId, Hook } from "../../catalog";
|
||||||
import type { ModelProvider } from "../agent";
|
import type { ModelProvider } from "../models";
|
||||||
|
|
||||||
export interface DeepSeekOptions {
|
export interface DeepSeekOptions {
|
||||||
apiKey?: string;
|
apiKey?: string;
|
||||||
|
|||||||
60
src/extensions/shared/models/index.test.ts
Normal file
60
src/extensions/shared/models/index.test.ts
Normal 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/);
|
||||||
|
});
|
||||||
63
src/extensions/shared/models/index.ts
Normal file
63
src/extensions/shared/models/index.ts
Normal 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;
|
||||||
@ -1,10 +1,12 @@
|
|||||||
import type { ExtensionFactory } from "../kernel";
|
import type { ExtensionFactory } from "../kernel";
|
||||||
import { createAgentExtension } from "../extensions/shared/agent";
|
import { createAgentExtension } from "../extensions/shared/agent";
|
||||||
import { createDeepSeekExtension } from "../extensions/shared/deepseek";
|
import { createDeepSeekExtension } from "../extensions/shared/deepseek";
|
||||||
|
import { createModelsExtension } from "../extensions/shared/models";
|
||||||
import { createWorkspaceExtension } from "../extensions/shared/workspace";
|
import { createWorkspaceExtension } from "../extensions/shared/workspace";
|
||||||
|
|
||||||
export const sharedExtensions: ExtensionFactory[] = [
|
export const sharedExtensions: ExtensionFactory[] = [
|
||||||
createWorkspaceExtension,
|
createWorkspaceExtension,
|
||||||
|
createModelsExtension,
|
||||||
createDeepSeekExtension,
|
createDeepSeekExtension,
|
||||||
createAgentExtension,
|
createAgentExtension,
|
||||||
];
|
];
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user