64 lines
1.6 KiB
TypeScript
64 lines
1.6 KiB
TypeScript
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;
|