feat(coding-agent): accept native extension providers
Allow extensions to register complete pi-ai Provider objects while preserving models.json composition and resolved provider auth access. Refs #4823 and #4824.
This commit is contained in:
@@ -2,6 +2,10 @@
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- Added extension registration for complete pi-ai providers, including native authentication, model refresh, filtering, and streaming behavior.
|
||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
- Fixed obsolete custom UI, custom tool, and custom editor examples in the extension documentation ([#6735](https://github.com/earendil-works/pi/issues/6735)).
|
- Fixed obsolete custom UI, custom tool, and custom editor examples in the extension documentation ([#6735](https://github.com/earendil-works/pi/issues/6735)).
|
||||||
|
|||||||
@@ -30,10 +30,38 @@ See these complete provider examples:
|
|||||||
|
|
||||||
## Quick Reference
|
## Quick Reference
|
||||||
|
|
||||||
|
Extensions can register either a complete pi-ai `Provider` or use the legacy provider-config form. Prefer a complete provider when custom authentication, filtering, refresh, or streaming behavior is required. Pi composes `models.json` overrides above registered native providers.
|
||||||
|
|
||||||
```typescript
|
```typescript
|
||||||
|
import { createProvider, openAICompletionsApi } from "@earendil-works/pi-ai";
|
||||||
import type { ExtensionAPI } from "@earendil-works/pi-coding-agent";
|
import type { ExtensionAPI } from "@earendil-works/pi-coding-agent";
|
||||||
|
|
||||||
export default function (pi: ExtensionAPI) {
|
export default function (pi: ExtensionAPI) {
|
||||||
|
pi.registerProvider(createProvider({
|
||||||
|
id: "native-local",
|
||||||
|
name: "Native Local",
|
||||||
|
baseUrl: "http://localhost:8080/v1",
|
||||||
|
auth: {
|
||||||
|
apiKey: {
|
||||||
|
name: "Local server API key",
|
||||||
|
async login(interaction) {
|
||||||
|
return {
|
||||||
|
type: "api_key",
|
||||||
|
key: await interaction.prompt({ type: "secret", message: "API key" })
|
||||||
|
};
|
||||||
|
},
|
||||||
|
async resolve({ credential }) {
|
||||||
|
return credential?.key
|
||||||
|
? { auth: { apiKey: credential.key }, source: "stored API key" }
|
||||||
|
: undefined;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
models: [],
|
||||||
|
api: openAICompletionsApi()
|
||||||
|
}));
|
||||||
|
|
||||||
|
// Legacy provider-config form:
|
||||||
// Override baseUrl for existing provider
|
// Override baseUrl for existing provider
|
||||||
pi.registerProvider("anthropic", {
|
pi.registerProvider("anthropic", {
|
||||||
baseUrl: "https://proxy.example.com"
|
baseUrl: "https://proxy.example.com"
|
||||||
|
|||||||
@@ -977,7 +977,7 @@ ctx.sessionManager.getLeafId() // Current leaf entry ID
|
|||||||
|
|
||||||
### ctx.modelRegistry / ctx.model
|
### ctx.modelRegistry / ctx.model
|
||||||
|
|
||||||
Access to models and API keys.
|
Access to models, providers, and resolved authentication. `ctx.modelRegistry.getProvider(id)` returns the effective pi-ai provider, while `getProviderAuth(id)` resolves its current API key, headers, base URL, and provider-scoped environment without requiring a loaded model. `ctx.model` is the active model.
|
||||||
|
|
||||||
### ctx.signal
|
### ctx.signal
|
||||||
|
|
||||||
@@ -1679,7 +1679,37 @@ Calls made during the extension factory function are queued and applied once the
|
|||||||
|
|
||||||
Dynamic providers can implement `refreshModels`. Pi calls it during model refresh, publishes the returned list synchronously through the provider, and passes the canonical credential/store/network/signal context. The extension decides whether to persist the catalog through `context.store`; live servers such as llama.cpp can ignore it.
|
Dynamic providers can implement `refreshModels`. Pi calls it during model refresh, publishes the returned list synchronously through the provider, and passes the canonical credential/store/network/signal context. The extension decides whether to persist the catalog through `context.store`; live servers such as llama.cpp can ignore it.
|
||||||
|
|
||||||
|
Extensions that need native provider auth, filtering, refresh, or stream behavior can register a complete `Provider` from `@earendil-works/pi-ai`. The provider becomes the composition base and `models.json` overrides still apply above it.
|
||||||
|
|
||||||
```typescript
|
```typescript
|
||||||
|
import { createProvider, openAICompletionsApi } from "@earendil-works/pi-ai";
|
||||||
|
|
||||||
|
const provider = createProvider({
|
||||||
|
id: "local-server",
|
||||||
|
name: "Local Server",
|
||||||
|
baseUrl: "http://localhost:8080/v1",
|
||||||
|
auth: {
|
||||||
|
apiKey: {
|
||||||
|
name: "Local server setup",
|
||||||
|
async login(interaction) {
|
||||||
|
return {
|
||||||
|
type: "api_key",
|
||||||
|
key: await interaction.prompt({ type: "secret", message: "API key" }),
|
||||||
|
};
|
||||||
|
},
|
||||||
|
async resolve({ credential }) {
|
||||||
|
return credential?.key
|
||||||
|
? { auth: { apiKey: credential.key }, source: "stored API key" }
|
||||||
|
: undefined;
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
models: [],
|
||||||
|
api: openAICompletionsApi(),
|
||||||
|
});
|
||||||
|
|
||||||
|
pi.registerProvider(provider);
|
||||||
|
|
||||||
// Register a new provider with custom models
|
// Register a new provider with custom models
|
||||||
pi.registerProvider("my-proxy", {
|
pi.registerProvider("my-proxy", {
|
||||||
name: "My Proxy",
|
name: "My Proxy",
|
||||||
@@ -1748,7 +1778,9 @@ pi.registerProvider("corporate-ai", {
|
|||||||
});
|
});
|
||||||
```
|
```
|
||||||
|
|
||||||
**Config options:**
|
The object form accepts a complete pi-ai `Provider`, including native `auth`, `getModels`, `refreshModels`, `filterModels`, `stream`, and `streamSimple` behavior.
|
||||||
|
|
||||||
|
**Legacy config options:**
|
||||||
- `name` - Display name for the provider in UI such as `/login`.
|
- `name` - Display name for the provider in UI such as `/login`.
|
||||||
- `baseUrl` - API endpoint URL. Required when defining models.
|
- `baseUrl` - API endpoint URL. Required when defining models.
|
||||||
- `apiKey` - API key literal, environment interpolation (`$ENV_VAR` or `${ENV_VAR}`), or leading `!command`. Required when defining models (unless `oauth` provided). `$$` escapes `$`, and `$!` escapes a literal `!` without triggering command execution.
|
- `apiKey` - API key literal, environment interpolation (`$ENV_VAR` or `${ENV_VAR}`), or leading `!command`. Required when defining models (unless `oauth` provided). `$$` escapes `$`, and `$!` escapes a literal `!` without triggering command execution.
|
||||||
|
|||||||
@@ -165,6 +165,18 @@ export async function createAgentSessionServices(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
extensionsResult.runtime.pendingProviderRegistrations = [];
|
extensionsResult.runtime.pendingProviderRegistrations = [];
|
||||||
|
for (const { provider, extensionPath } of extensionsResult.runtime.pendingNativeProviderRegistrations) {
|
||||||
|
try {
|
||||||
|
modelRuntime.registerNativeProvider(provider);
|
||||||
|
} catch (error) {
|
||||||
|
const message = error instanceof Error ? error.message : String(error);
|
||||||
|
diagnostics.push({
|
||||||
|
type: "error",
|
||||||
|
message: `Extension "${extensionPath}" error: ${message}`,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
extensionsResult.runtime.pendingNativeProviderRegistrations = [];
|
||||||
await modelRuntime.refresh({ allowNetwork: false });
|
await modelRuntime.refresh({ allowNetwork: false });
|
||||||
diagnostics.push(...applyExtensionFlagValues(resourceLoader, options.extensionFlagValues));
|
diagnostics.push(...applyExtensionFlagValues(resourceLoader, options.extensionFlagValues));
|
||||||
|
|
||||||
|
|||||||
@@ -2419,6 +2419,10 @@ export class AgentSession {
|
|||||||
this._modelRuntime.registerProvider(name, config);
|
this._modelRuntime.registerProvider(name, config);
|
||||||
this._refreshCurrentModelFromRegistry();
|
this._refreshCurrentModelFromRegistry();
|
||||||
},
|
},
|
||||||
|
registerNativeProvider: (provider) => {
|
||||||
|
this._modelRuntime.registerNativeProvider(provider);
|
||||||
|
this._refreshCurrentModelFromRegistry();
|
||||||
|
},
|
||||||
unregisterProvider: (name) => {
|
unregisterProvider: (name) => {
|
||||||
this._modelRuntime.unregisterProvider(name);
|
this._modelRuntime.unregisterProvider(name);
|
||||||
this._refreshCurrentModelFromRegistry();
|
this._refreshCurrentModelFromRegistry();
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import { createRequire } from "node:module";
|
|||||||
import * as path from "node:path";
|
import * as path from "node:path";
|
||||||
import { fileURLToPath } from "node:url";
|
import { fileURLToPath } from "node:url";
|
||||||
import * as _bundledPiAgentCore from "@earendil-works/pi-agent-core";
|
import * as _bundledPiAgentCore from "@earendil-works/pi-agent-core";
|
||||||
|
import type { Provider } from "@earendil-works/pi-ai";
|
||||||
import * as _bundledPiAiCompat from "@earendil-works/pi-ai/compat";
|
import * as _bundledPiAiCompat from "@earendil-works/pi-ai/compat";
|
||||||
import * as _bundledPiAiOauth from "@earendil-works/pi-ai/oauth";
|
import * as _bundledPiAiOauth from "@earendil-works/pi-ai/oauth";
|
||||||
import * as _bundledPiAiProviders from "@earendil-works/pi-ai/providers/all";
|
import * as _bundledPiAiProviders from "@earendil-works/pi-ai/providers/all";
|
||||||
@@ -195,6 +196,7 @@ export function createExtensionRuntime(): ExtensionRuntime {
|
|||||||
setThinkingLevel: notInitialized,
|
setThinkingLevel: notInitialized,
|
||||||
flagValues: new Map(),
|
flagValues: new Map(),
|
||||||
pendingProviderRegistrations: [],
|
pendingProviderRegistrations: [],
|
||||||
|
pendingNativeProviderRegistrations: [],
|
||||||
assertActive,
|
assertActive,
|
||||||
invalidate: (message) => {
|
invalidate: (message) => {
|
||||||
state.staleMessage ??=
|
state.staleMessage ??=
|
||||||
@@ -206,8 +208,14 @@ export function createExtensionRuntime(): ExtensionRuntime {
|
|||||||
registerProvider: (name, config, extensionPath = "<unknown>") => {
|
registerProvider: (name, config, extensionPath = "<unknown>") => {
|
||||||
runtime.pendingProviderRegistrations.push({ name, config, extensionPath });
|
runtime.pendingProviderRegistrations.push({ name, config, extensionPath });
|
||||||
},
|
},
|
||||||
|
registerNativeProvider: (provider, extensionPath = "<unknown>") => {
|
||||||
|
runtime.pendingNativeProviderRegistrations.push({ provider, extensionPath });
|
||||||
|
},
|
||||||
unregisterProvider: (name) => {
|
unregisterProvider: (name) => {
|
||||||
runtime.pendingProviderRegistrations = runtime.pendingProviderRegistrations.filter((r) => r.name !== name);
|
runtime.pendingProviderRegistrations = runtime.pendingProviderRegistrations.filter((r) => r.name !== name);
|
||||||
|
runtime.pendingNativeProviderRegistrations = runtime.pendingNativeProviderRegistrations.filter(
|
||||||
|
(r) => r.provider.id !== name,
|
||||||
|
);
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -363,9 +371,14 @@ function createExtensionAPI(
|
|||||||
runtime.setThinkingLevel(level);
|
runtime.setThinkingLevel(level);
|
||||||
},
|
},
|
||||||
|
|
||||||
registerProvider(name: string, config: ProviderConfig) {
|
registerProvider(providerOrName: Provider | string, config?: ProviderConfig) {
|
||||||
runtime.assertActive();
|
runtime.assertActive();
|
||||||
runtime.registerProvider(name, config, extension.path);
|
if (typeof providerOrName === "string") {
|
||||||
|
if (!config) throw new Error("Provider config is required when registering by name");
|
||||||
|
runtime.registerProvider(providerOrName, config, extension.path);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
runtime.registerNativeProvider(providerOrName, extension.path);
|
||||||
},
|
},
|
||||||
|
|
||||||
unregisterProvider(name: string) {
|
unregisterProvider(name: string) {
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
*/
|
*/
|
||||||
|
|
||||||
import type { AgentMessage } from "@earendil-works/pi-agent-core";
|
import type { AgentMessage } from "@earendil-works/pi-agent-core";
|
||||||
import type { ImageContent, Model, ProviderHeaders } from "@earendil-works/pi-ai";
|
import type { ImageContent, Model, Provider, ProviderHeaders } from "@earendil-works/pi-ai";
|
||||||
import type { KeyId } from "@earendil-works/pi-tui";
|
import type { KeyId } from "@earendil-works/pi-tui";
|
||||||
import { type Theme, theme } from "../../modes/interactive/theme/theme.ts";
|
import { type Theme, theme } from "../../modes/interactive/theme/theme.ts";
|
||||||
import type { ResourceDiagnostic } from "../diagnostics.ts";
|
import type { ResourceDiagnostic } from "../diagnostics.ts";
|
||||||
@@ -313,6 +313,7 @@ export class ExtensionRunner {
|
|||||||
contextActions: ExtensionContextActions,
|
contextActions: ExtensionContextActions,
|
||||||
providerActions?: {
|
providerActions?: {
|
||||||
registerProvider?: (name: string, config: ProviderConfig) => void;
|
registerProvider?: (name: string, config: ProviderConfig) => void;
|
||||||
|
registerNativeProvider?: (provider: Provider) => void;
|
||||||
unregisterProvider?: (name: string) => void;
|
unregisterProvider?: (name: string) => void;
|
||||||
},
|
},
|
||||||
): void {
|
): void {
|
||||||
@@ -363,6 +364,23 @@ export class ExtensionRunner {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
this.runtime.pendingProviderRegistrations = [];
|
this.runtime.pendingProviderRegistrations = [];
|
||||||
|
for (const { provider, extensionPath } of this.runtime.pendingNativeProviderRegistrations) {
|
||||||
|
try {
|
||||||
|
if (providerActions?.registerNativeProvider) {
|
||||||
|
providerActions.registerNativeProvider(provider);
|
||||||
|
} else {
|
||||||
|
this.modelRegistry.registerProvider(provider);
|
||||||
|
}
|
||||||
|
} catch (err) {
|
||||||
|
this.emitError({
|
||||||
|
extensionPath,
|
||||||
|
event: "register_provider",
|
||||||
|
error: err instanceof Error ? err.message : String(err),
|
||||||
|
stack: err instanceof Error ? err.stack : undefined,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
this.runtime.pendingNativeProviderRegistrations = [];
|
||||||
|
|
||||||
// From this point on, provider registration/unregistration takes effect immediately
|
// From this point on, provider registration/unregistration takes effect immediately
|
||||||
// without requiring a /reload.
|
// without requiring a /reload.
|
||||||
@@ -373,6 +391,13 @@ export class ExtensionRunner {
|
|||||||
}
|
}
|
||||||
this.modelRegistry.registerProvider(name, config);
|
this.modelRegistry.registerProvider(name, config);
|
||||||
};
|
};
|
||||||
|
this.runtime.registerNativeProvider = (provider) => {
|
||||||
|
if (providerActions?.registerNativeProvider) {
|
||||||
|
providerActions.registerNativeProvider(provider);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
this.modelRegistry.registerProvider(provider);
|
||||||
|
};
|
||||||
this.runtime.unregisterProvider = (name) => {
|
this.runtime.unregisterProvider = (name) => {
|
||||||
if (providerActions?.unregisterProvider) {
|
if (providerActions?.unregisterProvider) {
|
||||||
providerActions.unregisterProvider(name);
|
providerActions.unregisterProvider(name);
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ import type {
|
|||||||
Model,
|
Model,
|
||||||
OAuthCredentials,
|
OAuthCredentials,
|
||||||
OAuthLoginCallbacks,
|
OAuthLoginCallbacks,
|
||||||
|
Provider,
|
||||||
ProviderHeaders,
|
ProviderHeaders,
|
||||||
RefreshModelsContext,
|
RefreshModelsContext,
|
||||||
SimpleStreamOptions,
|
SimpleStreamOptions,
|
||||||
@@ -1378,6 +1379,7 @@ export interface ExtensionAPI {
|
|||||||
* }
|
* }
|
||||||
* });
|
* });
|
||||||
*/
|
*/
|
||||||
|
registerProvider(provider: Provider): void;
|
||||||
registerProvider(name: string, config: ProviderConfig): void;
|
registerProvider(name: string, config: ProviderConfig): void;
|
||||||
|
|
||||||
/**
|
/**
|
||||||
@@ -1553,8 +1555,10 @@ export type SetLabelHandler = (entryId: string, label: string | undefined) => vo
|
|||||||
*/
|
*/
|
||||||
export interface ExtensionRuntimeState {
|
export interface ExtensionRuntimeState {
|
||||||
flagValues: Map<string, boolean | string>;
|
flagValues: Map<string, boolean | string>;
|
||||||
/** Provider registrations queued during extension loading, processed when runner binds */
|
/** Legacy provider-config registrations queued during extension loading, processed when runner binds. */
|
||||||
pendingProviderRegistrations: Array<{ name: string; config: ProviderConfig; extensionPath: string }>;
|
pendingProviderRegistrations: Array<{ name: string; config: ProviderConfig; extensionPath: string }>;
|
||||||
|
/** Native pi-ai provider registrations queued during extension loading, processed when runner binds. */
|
||||||
|
pendingNativeProviderRegistrations: Array<{ provider: Provider; extensionPath: string }>;
|
||||||
/** Throws when this extension instance is stale after runtime replacement. */
|
/** Throws when this extension instance is stale after runtime replacement. */
|
||||||
assertActive: () => void;
|
assertActive: () => void;
|
||||||
/** Marks this extension instance as stale after runtime replacement or reload. */
|
/** Marks this extension instance as stale after runtime replacement or reload. */
|
||||||
@@ -1566,6 +1570,7 @@ export interface ExtensionRuntimeState {
|
|||||||
* After bindCore(): calls ModelRegistry directly for immediate effect.
|
* After bindCore(): calls ModelRegistry directly for immediate effect.
|
||||||
*/
|
*/
|
||||||
registerProvider: (name: string, config: ProviderConfig, extensionPath?: string) => void;
|
registerProvider: (name: string, config: ProviderConfig, extensionPath?: string) => void;
|
||||||
|
registerNativeProvider: (provider: Provider, extensionPath?: string) => void;
|
||||||
unregisterProvider: (name: string, extensionPath?: string) => void;
|
unregisterProvider: (name: string, extensionPath?: string) => void;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
import type { Api, Model } from "@earendil-works/pi-ai";
|
import type { Api, AuthResult, Model, Provider } from "@earendil-works/pi-ai";
|
||||||
import type { ModelRuntime } from "./model-runtime.ts";
|
import type { ModelRuntime } from "./model-runtime.ts";
|
||||||
import type { AuthStatus, ProviderConfigInput } from "./provider-composer.ts";
|
import type { AuthStatus, ProviderConfigInput } from "./provider-composer.ts";
|
||||||
|
|
||||||
@@ -92,10 +92,18 @@ export class ModelRegistry {
|
|||||||
return this.runtime.getProviderAuthStatus(provider);
|
return this.runtime.getProviderAuthStatus(provider);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
getProvider(provider: string): Provider | undefined {
|
||||||
|
return this.runtime.getProvider(provider);
|
||||||
|
}
|
||||||
|
|
||||||
getProviderDisplayName(provider: string): string {
|
getProviderDisplayName(provider: string): string {
|
||||||
return this.runtime.getProvider(provider)?.name ?? provider;
|
return this.runtime.getProvider(provider)?.name ?? provider;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
getProviderAuth(provider: string): Promise<AuthResult | undefined> {
|
||||||
|
return this.runtime.getAuth(provider);
|
||||||
|
}
|
||||||
|
|
||||||
async getApiKeyForProvider(provider: string): Promise<string | undefined> {
|
async getApiKeyForProvider(provider: string): Promise<string | undefined> {
|
||||||
try {
|
try {
|
||||||
return (await this.runtime.getAuth(provider))?.auth.apiKey;
|
return (await this.runtime.getAuth(provider))?.auth.apiKey;
|
||||||
@@ -108,8 +116,15 @@ export class ModelRegistry {
|
|||||||
return this.runtime.isUsingOAuth(model.provider);
|
return this.runtime.isUsingOAuth(model.provider);
|
||||||
}
|
}
|
||||||
|
|
||||||
registerProvider(providerName: string, config: ProviderConfigInput): void {
|
registerProvider(provider: Provider): void;
|
||||||
this.runtime.registerProvider(providerName, config);
|
registerProvider(providerName: string, config: ProviderConfigInput): void;
|
||||||
|
registerProvider(providerOrName: Provider | string, config?: ProviderConfigInput): void {
|
||||||
|
if (typeof providerOrName === "string") {
|
||||||
|
if (!config) throw new Error("Provider config is required when registering by name");
|
||||||
|
this.runtime.registerProvider(providerOrName, config);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
this.runtime.registerNativeProvider(providerOrName);
|
||||||
}
|
}
|
||||||
|
|
||||||
unregisterProvider(providerName: string): void {
|
unregisterProvider(providerName: string): void {
|
||||||
@@ -120,6 +135,10 @@ export class ModelRegistry {
|
|||||||
return this.runtime.getRegisteredProviderConfig(providerName);
|
return this.runtime.getRegisteredProviderConfig(providerName);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
getRegisteredNativeProvider(providerName: string): Provider | undefined {
|
||||||
|
return this.runtime.getRegisteredNativeProvider(providerName);
|
||||||
|
}
|
||||||
|
|
||||||
getRegisteredProviderIds(): readonly string[] {
|
getRegisteredProviderIds(): readonly string[] {
|
||||||
return this.runtime.getRegisteredProviderIds();
|
return this.runtime.getRegisteredProviderIds();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -94,6 +94,7 @@ export class ModelRuntime implements Models {
|
|||||||
private readonly credentials: RuntimeCredentials;
|
private readonly credentials: RuntimeCredentials;
|
||||||
private readonly defaultBuiltins: ReadonlyMap<string, Provider>;
|
private readonly defaultBuiltins: ReadonlyMap<string, Provider>;
|
||||||
private readonly builtins = new Map<string, Provider>();
|
private readonly builtins = new Map<string, Provider>();
|
||||||
|
private readonly nativeExtensionProviders = new Map<string, Provider>();
|
||||||
private readonly extensionProviders = new Map<string, ProviderConfigInput>();
|
private readonly extensionProviders = new Map<string, ProviderConfigInput>();
|
||||||
private readonly compositionErrors = new Map<string, string>();
|
private readonly compositionErrors = new Map<string, string>();
|
||||||
private readonly modelsPath: string | undefined;
|
private readonly modelsPath: string | undefined;
|
||||||
@@ -182,11 +183,16 @@ export class ModelRuntime implements Models {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private providerIds(): Set<string> {
|
private providerIds(): Set<string> {
|
||||||
return new Set([...this.builtins.keys(), ...this.config.getProviderIds(), ...this.extensionProviders.keys()]);
|
return new Set([
|
||||||
|
...this.builtins.keys(),
|
||||||
|
...this.nativeExtensionProviders.keys(),
|
||||||
|
...this.config.getProviderIds(),
|
||||||
|
...this.extensionProviders.keys(),
|
||||||
|
]);
|
||||||
}
|
}
|
||||||
|
|
||||||
private recomposeProvider(providerId: string): void {
|
private recomposeProvider(providerId: string): void {
|
||||||
const base = this.builtins.get(providerId);
|
const base = this.nativeExtensionProviders.get(providerId) ?? this.builtins.get(providerId);
|
||||||
const extension = this.extensionProviders.get(providerId);
|
const extension = this.extensionProviders.get(providerId);
|
||||||
if (!base && !this.config.getProvider(providerId) && !extension) {
|
if (!base && !this.config.getProvider(providerId) && !extension) {
|
||||||
this.models.deleteProvider(providerId);
|
this.models.deleteProvider(providerId);
|
||||||
@@ -335,7 +341,11 @@ export class ModelRuntime implements Models {
|
|||||||
}
|
}
|
||||||
|
|
||||||
getRegisteredProviderIds(): readonly string[] {
|
getRegisteredProviderIds(): readonly string[] {
|
||||||
return [...this.extensionProviders.keys()];
|
return [...new Set([...this.extensionProviders.keys(), ...this.nativeExtensionProviders.keys()])];
|
||||||
|
}
|
||||||
|
|
||||||
|
getRegisteredNativeProvider(providerId: string): Provider | undefined {
|
||||||
|
return this.nativeExtensionProviders.get(providerId);
|
||||||
}
|
}
|
||||||
|
|
||||||
/** @internal Compatibility fallback for ModelRegistry when provider auth is unconfigured. */
|
/** @internal Compatibility fallback for ModelRegistry when provider auth is unconfigured. */
|
||||||
@@ -520,10 +530,20 @@ export class ModelRuntime implements Models {
|
|||||||
return result;
|
return result;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
registerNativeProvider(provider: Provider): void {
|
||||||
|
if (!provider.id.trim()) throw new Error("Provider id must not be empty.");
|
||||||
|
this.extensionProviders.delete(provider.id);
|
||||||
|
this.nativeExtensionProviders.set(provider.id, provider);
|
||||||
|
this.recomposeProvider(provider.id);
|
||||||
|
this.updateModelSnapshot();
|
||||||
|
void this.refresh({ allowNetwork: false });
|
||||||
|
}
|
||||||
|
|
||||||
registerProvider(providerId: string, config: ProviderConfigInput): void {
|
registerProvider(providerId: string, config: ProviderConfigInput): void {
|
||||||
// Validate the incoming registration on its own, like the legacy registry:
|
// Validate the incoming registration on its own, like the legacy registry:
|
||||||
// a broken re-registration must throw without touching the stored config.
|
// a broken re-registration must throw without touching the stored config.
|
||||||
validateExtensionProvider(providerId, this.builtins.get(providerId), this.config.getProvider(providerId), config);
|
validateExtensionProvider(providerId, this.builtins.get(providerId), this.config.getProvider(providerId), config);
|
||||||
|
this.nativeExtensionProviders.delete(providerId);
|
||||||
// Re-registration merges defined values over the previous registration and
|
// Re-registration merges defined values over the previous registration and
|
||||||
// preserves undefined ones, matching the legacy ModelRegistry contract.
|
// preserves undefined ones, matching the legacy ModelRegistry contract.
|
||||||
const previous = this.extensionProviders.get(providerId);
|
const previous = this.extensionProviders.get(providerId);
|
||||||
@@ -559,6 +579,7 @@ export class ModelRuntime implements Models {
|
|||||||
|
|
||||||
unregisterProvider(providerId: string): void {
|
unregisterProvider(providerId: string): void {
|
||||||
this.extensionProviders.delete(providerId);
|
this.extensionProviders.delete(providerId);
|
||||||
|
this.nativeExtensionProviders.delete(providerId);
|
||||||
this.recomposeProvider(providerId);
|
this.recomposeProvider(providerId);
|
||||||
this.updateModelSnapshot();
|
this.updateModelSnapshot();
|
||||||
void this.refresh({ allowNetwork: false });
|
void this.refresh({ allowNetwork: false });
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import { existsSync, mkdirSync, rmSync } from "node:fs";
|
import { existsSync, mkdirSync, rmSync } from "node:fs";
|
||||||
import { tmpdir } from "node:os";
|
import { tmpdir } from "node:os";
|
||||||
import { join } from "node:path";
|
import { join } from "node:path";
|
||||||
|
import type { Provider } from "@earendil-works/pi-ai";
|
||||||
import { getModel } from "@earendil-works/pi-ai/compat";
|
import { getModel } from "@earendil-works/pi-ai/compat";
|
||||||
import { afterEach, beforeEach, describe, expect, it } from "vitest";
|
import { afterEach, beforeEach, describe, expect, it } from "vitest";
|
||||||
import { AuthStorage } from "../src/core/auth-storage.ts";
|
import { AuthStorage } from "../src/core/auth-storage.ts";
|
||||||
@@ -11,6 +12,28 @@ import { createAgentSession } from "../src/core/sdk.ts";
|
|||||||
import { SessionManager } from "../src/core/session-manager.ts";
|
import { SessionManager } from "../src/core/session-manager.ts";
|
||||||
import { SettingsManager } from "../src/core/settings-manager.ts";
|
import { SettingsManager } from "../src/core/settings-manager.ts";
|
||||||
|
|
||||||
|
function nativeAnthropicProvider(baseUrl: string): Provider {
|
||||||
|
const model = { ...getModel("anthropic", "claude-sonnet-4-5")!, baseUrl };
|
||||||
|
return {
|
||||||
|
id: "anthropic",
|
||||||
|
name: "Native Anthropic",
|
||||||
|
baseUrl,
|
||||||
|
auth: {
|
||||||
|
apiKey: {
|
||||||
|
name: "Test API key",
|
||||||
|
resolve: async () => ({ auth: { apiKey: "test-key" }, source: "test" }),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
getModels: () => [model],
|
||||||
|
stream: () => {
|
||||||
|
throw new Error("unused");
|
||||||
|
},
|
||||||
|
streamSimple: () => {
|
||||||
|
throw new Error("unused");
|
||||||
|
},
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
describe("AgentSession dynamic provider registration", () => {
|
describe("AgentSession dynamic provider registration", () => {
|
||||||
let tempDir: string;
|
let tempDir: string;
|
||||||
let agentDir: string;
|
let agentDir: string;
|
||||||
@@ -99,6 +122,19 @@ describe("AgentSession dynamic provider registration", () => {
|
|||||||
session.dispose();
|
session.dispose();
|
||||||
});
|
});
|
||||||
|
|
||||||
|
it("registers native pi-ai providers during extension loading", async () => {
|
||||||
|
const session = await createSession([
|
||||||
|
(pi) => {
|
||||||
|
pi.registerProvider(nativeAnthropicProvider("http://localhost:8080/native-top-level"));
|
||||||
|
},
|
||||||
|
]);
|
||||||
|
|
||||||
|
expect(session.model?.baseUrl).toBe("http://localhost:8080/native-top-level");
|
||||||
|
expect(await capturePromptBaseUrl(session)).toBe("http://localhost:8080/native-top-level");
|
||||||
|
|
||||||
|
session.dispose();
|
||||||
|
});
|
||||||
|
|
||||||
it("applies command-time registerProvider overrides without reload", async () => {
|
it("applies command-time registerProvider overrides without reload", async () => {
|
||||||
const session = await createSession([
|
const session = await createSession([
|
||||||
(pi) => {
|
(pi) => {
|
||||||
@@ -119,4 +155,25 @@ describe("AgentSession dynamic provider registration", () => {
|
|||||||
|
|
||||||
session.dispose();
|
session.dispose();
|
||||||
});
|
});
|
||||||
|
|
||||||
|
it("registers native pi-ai providers at command time", async () => {
|
||||||
|
const session = await createSession([
|
||||||
|
(pi) => {
|
||||||
|
pi.registerCommand("use-native", {
|
||||||
|
description: "Use native provider",
|
||||||
|
handler: async () => {
|
||||||
|
pi.registerProvider(nativeAnthropicProvider("http://localhost:8080/native-command"));
|
||||||
|
},
|
||||||
|
});
|
||||||
|
},
|
||||||
|
]);
|
||||||
|
|
||||||
|
await session.bindExtensions({});
|
||||||
|
await session.prompt("/use-native");
|
||||||
|
|
||||||
|
expect(session.model?.baseUrl).toBe("http://localhost:8080/native-command");
|
||||||
|
expect(await capturePromptBaseUrl(session)).toBe("http://localhost:8080/native-command");
|
||||||
|
|
||||||
|
session.dispose();
|
||||||
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -1,6 +1,10 @@
|
|||||||
import { InMemoryModelsStore, type Model } from "@earendil-works/pi-ai";
|
import { mkdtempSync, rmSync, writeFileSync } from "node:fs";
|
||||||
|
import { tmpdir } from "node:os";
|
||||||
|
import { join } from "node:path";
|
||||||
|
import { InMemoryModelsStore, type Model, type Provider } from "@earendil-works/pi-ai";
|
||||||
import { describe, expect, it } from "vitest";
|
import { describe, expect, it } from "vitest";
|
||||||
import { AuthStorage } from "../src/core/auth-storage.ts";
|
import { AuthStorage } from "../src/core/auth-storage.ts";
|
||||||
|
import { ModelRegistry } from "../src/core/model-registry.ts";
|
||||||
import { ModelRuntime } from "../src/core/model-runtime.ts";
|
import { ModelRuntime } from "../src/core/model-runtime.ts";
|
||||||
|
|
||||||
function model(id: string): Model<"openai-completions"> {
|
function model(id: string): Model<"openai-completions"> {
|
||||||
@@ -19,6 +23,118 @@ function model(id: string): Model<"openai-completions"> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
describe("extension provider model lifecycle", () => {
|
describe("extension provider model lifecycle", () => {
|
||||||
|
it("registers native pi-ai providers with their auth implementation", async () => {
|
||||||
|
const runtime = await ModelRuntime.create({
|
||||||
|
credentials: AuthStorage.inMemory(),
|
||||||
|
modelsStore: new InMemoryModelsStore(),
|
||||||
|
modelsPath: null,
|
||||||
|
allowModelNetwork: false,
|
||||||
|
});
|
||||||
|
const nativeModel = {
|
||||||
|
...model("native"),
|
||||||
|
provider: "extension-native",
|
||||||
|
baseUrl: "https://fallback.test/v1",
|
||||||
|
};
|
||||||
|
const provider: Provider = {
|
||||||
|
id: "extension-native",
|
||||||
|
name: "Extension Native",
|
||||||
|
auth: {
|
||||||
|
apiKey: {
|
||||||
|
name: "Native setup",
|
||||||
|
login: async (interaction) => ({
|
||||||
|
type: "api_key",
|
||||||
|
key: await interaction.prompt({ type: "secret", message: "API key" }),
|
||||||
|
}),
|
||||||
|
check: async ({ credential }) =>
|
||||||
|
credential?.key ? { type: "api_key", source: "stored native key" } : undefined,
|
||||||
|
resolve: async ({ credential }) =>
|
||||||
|
credential?.key
|
||||||
|
? {
|
||||||
|
auth: { apiKey: credential.key, baseUrl: "https://resolved.test/v1" },
|
||||||
|
source: "stored native key",
|
||||||
|
}
|
||||||
|
: undefined,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
getModels: () => [nativeModel],
|
||||||
|
stream: () => {
|
||||||
|
throw new Error("unused");
|
||||||
|
},
|
||||||
|
streamSimple: () => {
|
||||||
|
throw new Error("unused");
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
runtime.registerNativeProvider(provider);
|
||||||
|
const registry = new ModelRegistry(runtime);
|
||||||
|
expect(registry.getProvider("extension-native")).toBe(provider);
|
||||||
|
expect(registry.getRegisteredNativeProvider("extension-native")).toBe(provider);
|
||||||
|
expect(registry.getRegisteredProviderIds()).toContain("extension-native");
|
||||||
|
expect(registry.find("extension-native", "native")).toBeDefined();
|
||||||
|
|
||||||
|
await runtime.login("extension-native", "api_key", {
|
||||||
|
prompt: async () => "secret",
|
||||||
|
notify: () => {},
|
||||||
|
});
|
||||||
|
expect(await registry.getProviderAuth("extension-native")).toMatchObject({
|
||||||
|
auth: { apiKey: "secret", baseUrl: "https://resolved.test/v1" },
|
||||||
|
});
|
||||||
|
|
||||||
|
registry.unregisterProvider("extension-native");
|
||||||
|
expect(registry.getProvider("extension-native")).toBeUndefined();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("applies models.json overrides above native providers", async () => {
|
||||||
|
const tempDir = mkdtempSync(join(tmpdir(), "pi-native-provider-"));
|
||||||
|
const modelsPath = join(tempDir, "models.json");
|
||||||
|
writeFileSync(
|
||||||
|
modelsPath,
|
||||||
|
JSON.stringify({
|
||||||
|
providers: {
|
||||||
|
"extension-native": {
|
||||||
|
modelOverrides: {
|
||||||
|
native: { contextWindow: 4242 },
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
try {
|
||||||
|
const runtime = await ModelRuntime.create({
|
||||||
|
credentials: AuthStorage.inMemory(),
|
||||||
|
modelsStore: new InMemoryModelsStore(),
|
||||||
|
modelsPath,
|
||||||
|
allowModelNetwork: false,
|
||||||
|
});
|
||||||
|
const nativeModel = {
|
||||||
|
...model("native"),
|
||||||
|
provider: "extension-native",
|
||||||
|
baseUrl: "https://native.test/v1",
|
||||||
|
};
|
||||||
|
runtime.registerNativeProvider({
|
||||||
|
id: "extension-native",
|
||||||
|
name: "Extension Native",
|
||||||
|
auth: {
|
||||||
|
apiKey: {
|
||||||
|
name: "Native key",
|
||||||
|
resolve: async () => ({ auth: { apiKey: "key" }, source: "native" }),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
getModels: () => [nativeModel],
|
||||||
|
stream: () => {
|
||||||
|
throw new Error("unused");
|
||||||
|
},
|
||||||
|
streamSimple: () => {
|
||||||
|
throw new Error("unused");
|
||||||
|
},
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(runtime.getModel("extension-native", "native")?.contextWindow).toBe(4242);
|
||||||
|
} finally {
|
||||||
|
rmSync(tempDir, { recursive: true, force: true });
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
it("publishes refreshModels results without forcing ModelsStore persistence", async () => {
|
it("publishes refreshModels results without forcing ModelsStore persistence", async () => {
|
||||||
const modelsStore = new InMemoryModelsStore();
|
const modelsStore = new InMemoryModelsStore();
|
||||||
const runtime = await ModelRuntime.create({
|
const runtime = await ModelRuntime.create({
|
||||||
|
|||||||
Reference in New Issue
Block a user