diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index b20520f55..375477f76 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -2,9 +2,26 @@ ## [Unreleased] +### Breaking Changes + +- Replaced Atomic's legacy extension OAuth registration bridge with provider-owned authentication from `@earendil-works/pi-ai`. Extensions must declare OAuth or API-key authentication on their provider registration. The package root no longer exports the bridge functions `registerOAuthProvider`, `resetOAuthProviders`, `getOAuthApiKey`, `getOAuthProvider`, or `getOAuthProviders`, nor the bridge and credential/status types `LegacyOAuthProvider`, `OAuthProviderDescriptor`, `ApiKeyCredential`, `AuthCredential`, `AuthStatus`, or `OAuthCredential`; import current credential types from `@earendil-works/pi-ai` and use provider-owned authentication instead. The internal legacy registration/refresh machinery and custom API-key login hooks beyond pi's provider contract were also removed. +- Model configuration now follows pi's single-file `ModelConfig` contract. Atomic reads one `models.json` from the active Atomic agent directory (`~/.atomic/agent/models.json`, or the directory selected by `ATOMIC_CODING_AGENT_DIR`/`PI_CODING_AGENT_DIR`); it no longer reads project-scoped `.atomic/models.json`, falls back to `~/.pi/agent/models.json`, or layers and merges `.pi` and `.atomic` model configuration files. Move project-scoped custom providers and models into the active agent-directory file. The legacy `.pi` read fallback remains only for configuration surfaces that explicitly use layered paths, such as `auth.json`. +- Model-catalog refreshes now use pi's exact timeout semantics. `modelRefreshTimeoutMs` applies only to the initial runtime creation refresh, and post-login catalog refreshes are unbounded; the interactive `/model` selector owns its 15-second refresh timeout like pi, rendering cached models immediately, aborting a slow refresh, and reporting "Model refresh timed out; showing cached models." Callers of `ModelRuntime.refresh()` that require cancellation must provide their own abort signal. +- Extension `streamSimple` implementations are now scoped to their registered provider and `ModelRuntime`, matching pi's provider composer. Unregistering a newer provider no longer restores an older global API-owner registration; extensions that replace a provider stream must keep that provider registered for as long as the stream should remain active. +- Remote model catalogs now publish refreshed models in memory before persisting the catalog, matching pi's ordering. If the catalog-store write fails, refresh reports the storage error but the newly refreshed in-memory catalog remains active for the current process. +- `AuthStorage` now implements pi's asynchronous `CredentialStore` contract directly. Its synchronous compatibility methods (`get`, `set`, `remove`, `has`, `hasAuth`, `getAll`, `getAuthStatus`, `getLoadError`, `drainErrors`, and synchronous `list`), `asCredentialStore()` adapter, fallback/runtime-key helpers (`setFallbackResolver`, `getRuntimeApiKey`, `setRuntimeApiKey`, and `removeRuntimeApiKey`), and provider-auth methods (`login`, `logout`, `logoutAsync`, `getModelAuth`, `getApiKey`, and `getOAuthProviders`) were removed. Use and `await` `read`, `list`, `modify`, and `delete` for persisted credentials, call `reload` when an existing instance must reread its backing store, and use `ModelRuntime` for authentication resolution, login/logout, and runtime API-key overrides. +- `ModelRegistry` is now constructed from a `ModelRuntime` and exposes only pi's thin compatibility-facade surface. The `create` and `inMemory` factories, public `authStorage` member, and the `canRestoreUnknownModel`, `checkAuth`, `getAuth`, `getCustomApiKeyAuth`, `getCustomApiKeyAuthProviders`, `getProviders`, `hasProvider`, `hasRegisteredStreamSimpleForApi`, `login`, and `logoutProvider` methods were removed. Instantiate and consume `ModelRuntime` for runtime-owned model/provider discovery and authentication, and pass it to `new ModelRegistry(runtime)` only where the synchronous extension-facing facade is required. +- `CreateAgentSessionOptions` no longer accepts the public `authStorage` and `modelRegistry` overrides. Construct or supply a `ModelRuntime` through `modelRuntime` instead; when omitted, `createAgentSession()` creates the runtime from the active agent directory's `auth.json` and `models.json`. + +### Changed + +- Adopted pi's provider-owned `ModelRuntime` architecture for model composition, credentials, streaming, and catalog refresh. `ModelRegistry` is now the thin synchronous extension-compatibility facade used by pi, while coding-agent, SDK, RPC, isolated-engine, workflow, and MCP internals consume `ModelRuntime` directly. + ### Fixed - Fixed OAuth logins being destroyed by an unrelated model-catalog refresh, matching upstream pi's behavior. After a successful `/login`, a timed-out (aborted) or partially failed catalog refresh threw `Model refresh aborted after OAuth login` and rolled the freshly acquired tokens back to the previous credential. Because providers rotate refresh tokens, that rollback could permanently strand a server-side-invalidated credential — every send then failed with `invalid_grant` ("Refresh token not found or invalid") and every re-login was rolled back again, typically on machines with slow routes to catalog endpoints. Freshly persisted OAuth credentials now always survive the post-login refresh: per-provider refresh errors and refresh timeouts no longer fail the login in either the direct interactive or isolated-engine path, and models fall back to the cached snapshot. +- Fixed the `/model` selector reporting `Could not refresh llama.cpp; showing cached models.` for users who never configured a llama.cpp server. The bundled llama.cpp extension now uses pi's provider-owned registration: the provider stays dormant — no refresh attempt, no error — until a server is configured through `LLAMA_BASE_URL` or a stored login, and `/login` prompts for the server URL plus optional API key exactly like pi. +- Fixed RPC `save_provider_credential` writes disappearing after process restart. Saved API-key and OAuth credentials now persist to `auth.json` through `ModelRuntime.saveCredential()` instead of being stored as non-persistent runtime API-key overrides; the RPC command again accepts the full credential union, awaits a model-catalog refresh, and returns the refreshed catalog. ## [0.9.11-alpha.7] - 2026-07-28 diff --git a/packages/coding-agent/README.md b/packages/coding-agent/README.md index 728e625ef..932023205 100644 --- a/packages/coding-agent/README.md +++ b/packages/coding-agent/README.md @@ -429,14 +429,12 @@ See [docs/packages.md](docs/packages.md). ### SDK ```typescript -import { AuthStorage, createAgentSession, ModelRegistry, SessionManager } from "@bastani/atomic"; +import { createAgentSession, SessionManager } from "@bastani/atomic"; -const authStorage = AuthStorage.create(); -const modelRegistry = ModelRegistry.create(authStorage); +// By default, createAgentSession builds a ModelRuntime from the active agent +// directory's auth.json and models.json. const { session } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage, - modelRegistry, }); await session.prompt("What files are in the current directory?"); diff --git a/packages/coding-agent/docs/changelog.mdx b/packages/coding-agent/docs/changelog.mdx index 996719550..f17ce3eb2 100644 --- a/packages/coding-agent/docs/changelog.mdx +++ b/packages/coding-agent/docs/changelog.mdx @@ -30,7 +30,7 @@ See [SDK](/sdk) and [Custom providers](/custom-provider). ## Pi 0.80.6 compatibility - **Synced through upstream Pi 0.80.6.** Atomic and its bundled extensions now use the 0.80.6 Pi runtime packages, including signed empty Anthropic thinking preservation, request-wide GPT-5.4/5.5 long-context pricing, and corrected GPT-5.6 catalog/backend metadata. -- **Model controls are current.** Atomic accepts `max` wherever the active model advertises it, applies `models.json` overrides to matching extension-registered models and built-ins, and loads legacy `.pi` plus primary `.atomic` layers during normal CLI startup. Disjoint override IDs survive, exact primary entries replace complete legacy entries, and removed catalog IDs fall back instead of being synthesized on resume. See [Models](/models). +- **Model controls are current.** Atomic accepts `max` wherever the active model advertises it, applies the active agent directory's single `models.json` overrides to matching extension-registered models and built-ins, and falls back instead of synthesizing removed catalog IDs on resume. See [Models](/models). - **Custom pricing tiers are preserved.** Custom `models.json` entries and extension providers can declare complete request-wide `cost.tiers`; matching `modelOverrides` can update scalar rates without losing inherited tiers or replace/clear the tier array explicitly. See [Models](/models#request-wide-cost-tiers) and [Custom Providers](/custom-provider#usage-and-cost). - **Safer session and tool behavior.** Invalid explicit bash timeouts fail instead of being clamped; missing exact session IDs warn before creation; lax null or omitted message content is normalized; auth writes surface persistence failures; Windows context traversal terminates at drive roots; and duplicate fork selections are ignored. See [Tools](/tools) and [Usage](/usage). - **More reliable streaming and binaries.** Visible custom messages stay before the live assistant row, standalone Linux clipboard reads can fall back to xclip with correctly packaged native bindings, caller-relative `TMPDIR` paths work from external directories, and `--skip-deps` tolerates a missing optional clipboard wrapper. diff --git a/packages/coding-agent/docs/custom-provider.md b/packages/coding-agent/docs/custom-provider.md index ca7ed0f7f..a9e55dd36 100644 --- a/packages/coding-agent/docs/custom-provider.md +++ b/packages/coding-agent/docs/custom-provider.md @@ -312,7 +312,7 @@ After registration, users can authenticate via `/login corporate-ai`. Existing extension OAuth definitions keep their `login`, `refreshToken`, `getApiKey`, and optional `modifyModels` methods. OAuth refresh is serialized so concurrent requests do not overwrite each other's credentials. -In isolated interactive mode, extension code and executable OAuth methods remain in the engine process. Atomic transports only the JSON-safe provider description (`id`, `name`, `loginLabel`, and `usesCallbackServer`) to the terminal process; it never serializes provider functions or acquired credentials and does not load the extension a second time in the frontend. The engine executes the provider's login closure and correlates browser URLs, device codes, progress/info messages, prompts, selections, and manual-code callbacks with the originating login. +In isolated interactive mode, extension code and executable OAuth methods remain in the engine process. Atomic transports only the JSON-safe provider description (`id`, `name`, `loginLabel`, and `usesCallbackServer`) to the terminal process; it never serializes provider functions or acquired credentials and does not load the extension a second time in the frontend. `loginLabel` replaces the login dialog title, while `usesCallbackServer: true` exposes a redirect-URL paste field that races the browser callback. The engine executes the provider's login closure and correlates browser URLs, device codes, progress/info messages, prompts, selections, and manual-code callbacks with the originating login. After acquisition, the engine owns serialized credential persistence, model/current-model refresh, rollback on failure, and logout. The frontend applies the returned catalog only after the engine transaction succeeds. Escape or Ctrl+C cancels only the matching login and leaves the prior credential/catalog intact. Built-in OAuth and direct, non-isolated extension OAuth use the same persistence and cancellation semantics; later provider registrations continue to override earlier registrations by ID. diff --git a/packages/coding-agent/docs/models.md b/packages/coding-agent/docs/models.md index 94d55bd51..2506f4419 100644 --- a/packages/coding-agent/docs/models.md +++ b/packages/coding-agent/docs/models.md @@ -1,8 +1,6 @@ # Custom Models -Add custom providers and models (Ollama, vLLM, LM Studio, proxies) via `~/.atomic/agent/models.json` (legacy `~/.pi/agent/models.json` is also read). - -When both files exist, Atomic reads the legacy `.pi` file first and the primary `.atomic` file second. For `modelOverrides`, entries are layered by provider and model ID: disjoint legacy entries remain available, while an exact primary provider/model entry replaces the complete legacy override entry. Atomic does not field-merge one override entry across files; use `{}` in the primary file to restore the built-in model values for that exact entry. +Add custom providers and models (Ollama, vLLM, LM Studio, proxies) via the single `models.json` in the active Atomic agent directory, normally `~/.atomic/agent/models.json`, or the directory selected by `ATOMIC_CODING_AGENT_DIR`/`PI_CODING_AGENT_DIR`. Atomic reads only that file: it does not read project-scoped `.atomic/models.json`, fall back to `~/.pi/agent/models.json`, or merge `.pi` and `.atomic` model configuration files. The legacy `.pi` read fallback remains available for configuration surfaces that explicitly use layered config paths, such as `auth.json`; it does not apply to `models.json`. A complete `defaultProvider`/`defaultModel` pair in `settings.json` is resolved after built-in, configured, and extension providers register. If the provider remains unsupported, interactive mode reports a generic saved-configuration warning and leaves model selection open instead of routing the session to a different provider. Print and JSON modes write that diagnostic to stderr and exit nonzero before prompting, keeping JSON stdout JSONL-clean. RPC rejects `prompt` with the same correlated diagnostic until an explicit successful `set_model` selects an available model or an explicit model cycle returns a different available model. A null or unchanged cycle result does not clear the condition. If the provider is supported but the model is unknown or lacks authentication, normal automatic selection of an available authenticated model continues. Valid custom- and extension-provider defaults resolve once their provider registration is available. See [Settings](/settings#model--thinking). @@ -95,7 +93,7 @@ Override defaults when you need specific values: } ``` -Atomic reloads every configured `models.json` layer each time you open `/model`: legacy global/project `.pi` sources first, then primary global/project `.atomic` sources. Provider definitions, complete per-model overrides, dynamic catalogs, and isolated-engine model state are rebuilt from that fresh layered view, so edits take effect without restarting. Invalid edits report an error and do not silently reuse a different layer. +Atomic reloads the active agent directory's single `models.json` each time you open `/model`. Provider definitions, per-model overrides, dynamic catalogs, and isolated-engine model state are rebuilt from that fresh configuration, so edits take effect without restarting. Invalid edits report an error. ## Google AI Studio Example @@ -418,13 +416,13 @@ Use `modelOverrides` to customize specific models without replacing the provider `modelOverrides` supports these fields per model: `name`, `reasoning`, `thinkingLevelMap`, `input`, `cost` (partial scalar rates plus optional full tier-array replacement), `contextWindow`, `maxTokens`, `headers`, `compat`. -When both `~/.pi/agent/models.json` and `~/.atomic/agent/models.json` define `modelOverrides`, Atomic merges their nested provider/model maps in that order. Different model IDs survive from both files. For the same provider and model ID, the primary `.atomic` entry replaces the entire legacy `.pi` override entry rather than deep-merging individual fields. This complete-entry rule includes `headers`: a primary exact override without headers removes headers that came from the legacy override, but does not erase a surviving custom model definition's own headers. An empty primary override (`{}`) therefore restores the model's built-in values for that entry. +Atomic reads one `models.json` from the active agent directory. It does not layer model overrides from `.pi` and `.atomic` files. Within a single file, custom model definitions replace matching built-in entries after built-in overrides are applied. `modelOverrides` composes only with built-in and extension-registered models; it does not modify a same-ID custom model definition. Behavior notes: - Atomic retains the parsed override map even when an extension registers the matching provider/model after `models.json` is loaded. -- Layered primary/legacy compatibility merges override maps by provider and model ID; disjoint entries survive, while a primary exact entry replaces the complete legacy entry without cross-file field-level merging. +- Model overrides come from the active agent directory's single `models.json`; no cross-file layering or merging is performed. - For matching built-in and extension-registered models, the model definition is the base and `modelOverrides` wins configured fields. Extension-registered model headers are shallow-merged with override headers, with override headers winning duplicate names. A same-ID custom model replaces the built-in override result, including its complete header record. - A scalar-only `cost` override preserves inherited tiers. Supplying `cost.tiers` replaces the complete tier array, including `[]` to clear it; omitted scalar cost fields remain inherited. - Provider-level request headers remain a separate provider layer and are combined at request time. diff --git a/packages/coding-agent/docs/sdk.md b/packages/coding-agent/docs/sdk.md index 1bafed4f8..86d138efb 100644 --- a/packages/coding-agent/docs/sdk.md +++ b/packages/coding-agent/docs/sdk.md @@ -16,16 +16,13 @@ See [examples/sdk/](https://github.com/bastani-inc/atomic/tree/main/packages/cod ## Quick Start ```typescript -import { AuthStorage, createAgentSession, ModelRegistry, SessionManager } from "@bastani/atomic"; +import { createAgentSession, ModelRuntime, SessionManager } from "@bastani/atomic"; -// Set up credential storage and model registry -const authStorage = AuthStorage.create(); -const modelRegistry = ModelRegistry.create(authStorage); +const modelRuntime = await ModelRuntime.create(); const { session } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage, - modelRegistry, + modelRuntime, }); session.subscribe((event) => { @@ -423,21 +420,19 @@ When you pass a custom `ResourceLoader`, `cwd` and `agentDir` no longer control ```typescript import { getModel } from "@earendil-works/pi-ai/compat"; -import { AuthStorage, ModelRegistry } from "@bastani/atomic"; +import { ModelRuntime } from "@bastani/atomic"; -const authStorage = AuthStorage.create(); -const modelRegistry = ModelRegistry.create(authStorage); +const modelRuntime = await ModelRuntime.create(); -// Find specific built-in model (doesn't check if API key exists) +// Find specific built-in model (doesn't check if credentials exist) const opus = getModel("anthropic", "claude-opus-4-5"); if (!opus) throw new Error("Model not found"); // Find any model by provider/id, including custom models from models.json -// (doesn't check if API key exists) -const customModel = modelRegistry.find("my-provider", "my-model"); +const customModel = modelRuntime.getModel("my-provider", "my-model"); -// Get only models that have valid API keys configured -const available = await modelRegistry.getAvailable(); +// Get only models whose providers have configured authentication +const available = await modelRuntime.getAvailable(); const { session } = await createAgentSession({ model: opus, @@ -449,12 +444,11 @@ const { session } = await createAgentSession({ { model: haiku, thinkingLevel: "off" }, ], - authStorage, - modelRegistry, + modelRuntime, }); ``` -`ModelRegistry` keeps synchronous reads for SDK and extension compatibility, while catalog refresh is asynchronous. Await `modelRegistry.refresh()` before reading `getAll()`, `find()`, or `getAvailable()` when a provider may update its catalog. The refresh result reports `aborted` and per-provider `errors`; successful providers publish their new catalogs even if another provider fails, and failed or timed-out providers retain their last-known models. +`ModelRegistry` keeps synchronous reads for extension compatibility, while catalog refresh is asynchronous. Extensions should await `modelRegistry.refresh()` before synchronous `getAll()`, `find()`, or `getAvailable()` reads when a provider may update its catalog. New SDK integrations use `ModelRuntime`; `await modelRuntime.refresh()` reports `aborted` and per-provider `errors`, and failed providers retain their last-known models. If no model is provided: 1. Tries to restore from session (if continuing) @@ -465,48 +459,40 @@ If no model is provided: ### API Keys and OAuth -`AuthStorage` and `ModelRegistry` remain synchronous public SDK entry points. Token refresh is serialized under Atomic's credential-store lock, and the `authStorage` and `modelRegistry` session options remain available. +`ModelRuntime` is the asynchronous SDK engine for provider composition, credentials, model catalogs, and requests. `ModelRegistry` remains a thin synchronous compatibility facade for extensions; new SDK integrations should pass `modelRuntime` to `createAgentSession`. -API key resolution priority (handled by AuthStorage): -1. Runtime overrides (via `setRuntimeApiKey`, not persisted) -2. Stored credentials in `auth.json` (API keys or OAuth tokens) -3. Environment variables (`ANTHROPIC_API_KEY`, `OPENAI_API_KEY`, etc.) -4. Fallback resolver (for custom provider keys from `models.json`) - -OAuth credentials may provide an API key, request headers, and a credential-specific `baseUrl`. Atomic applies all three to the request. In particular, GitHub Copilot enterprise and token-specific endpoints replace the static model URL without dropping retries, attribution headers, fast mode, or extension request hooks. +Credential resolution combines runtime API-key overrides, stored `auth.json` credentials, environment variables, and the active `models.json` provider configuration. OAuth acquisition is provider-owned and runs through `ModelRuntime.login()`. ```typescript -import { AuthStorage, ModelRegistry } from "@bastani/atomic"; +import { AuthStorage, ModelRuntime } from "@bastani/atomic"; -// Default: uses ~/.atomic/agent/auth.json and ~/.atomic/agent/models.json, -// with legacy ~/.pi/agent/* compatibility reads when available. const authStorage = AuthStorage.create(); -const modelRegistry = ModelRegistry.create(authStorage); +const modelRuntime = await ModelRuntime.create({ credentials: authStorage }); const { session } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage, - modelRegistry, + modelRuntime, }); // Runtime API key override (not persisted to disk) -authStorage.setRuntimeApiKey("anthropic", "sk-my-temp-key"); +await modelRuntime.setRuntimeApiKey("anthropic", "sk-my-temp-key"); -// Custom auth storage location -const customAuth = AuthStorage.create("/my/app/auth.json"); -const customRegistry = ModelRegistry.create(customAuth, "/my/app/models.json"); +// Custom credential and model configuration locations +const customRuntime = await ModelRuntime.create({ + authPath: "/my/app/auth.json", + modelsPath: "/my/app/models.json", +}); -const { session } = await createAgentSession({ +const customSession = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage: customAuth, - modelRegistry: customRegistry, + modelRuntime: customRuntime, }); -// No custom models.json (built-in models only) -const simpleRegistry = ModelRegistry.inMemory(authStorage); +// Disable models.json while retaining built-in providers +const builtinsOnly = await ModelRuntime.create({ modelsPath: null }); ``` -> See [examples/sdk/09-api-keys-and-oauth.ts](https://github.com/bastani-inc/atomic/blob/main/packages/coding-agent/examples/sdk/09-api-keys-and-oauth.ts) +> See the complete [`ModelRuntime` credential and model configuration example](https://github.com/bastani-inc/atomic/blob/main/packages/coding-agent/examples/sdk/09-api-keys-and-oauth.ts). ### System Prompt @@ -1014,22 +1000,20 @@ import { createAgentSession, DefaultResourceLoader, defineTool, - ModelRegistry, + ModelRuntime, SessionManager, SettingsManager, } from "@bastani/atomic"; -// Set up auth storage (custom location) +// Create a runtime with custom credential storage and no models.json. const authStorage = AuthStorage.create("/custom/agent/auth.json"); +const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); // Runtime API key override (not persisted) if (process.env.MY_KEY) { - authStorage.setRuntimeApiKey("anthropic", process.env.MY_KEY); + await modelRuntime.setRuntimeApiKey("anthropic", process.env.MY_KEY); } -// Model registry (no custom models.json) -const modelRegistry = ModelRegistry.create(authStorage); - // Inline tool const statusTool = defineTool({ name: "status", @@ -1065,8 +1049,7 @@ const { session } = await createAgentSession({ model, thinkingLevel: "off", - authStorage, - modelRegistry, + modelRuntime, tools: ["read", "bash", "status"], customTools: [statusTool], diff --git a/packages/coding-agent/docs/workflows.md b/packages/coding-agent/docs/workflows.md index f46349194..b63411abe 100644 --- a/packages/coding-agent/docs/workflows.md +++ b/packages/coding-agent/docs/workflows.md @@ -2333,8 +2333,7 @@ Select the stage working directory and agent configuration directory. Worktree-e ```typescript // Runtime StageOptions forwards non-workflow CreateAgentSessionOptions, // including these advanced host integration fields: -readonly authStorage?: CreateAgentSessionOptions["authStorage"]; -readonly modelRegistry?: CreateAgentSessionOptions["modelRegistry"]; +readonly modelRuntime?: CreateAgentSessionOptions["modelRuntime"]; readonly resourceLoader?: CreateAgentSessionOptions["resourceLoader"]; readonly sessionManager?: SessionManager; readonly settingsManager?: SettingsManager; diff --git a/packages/coding-agent/examples/sdk/02-custom-model.ts b/packages/coding-agent/examples/sdk/02-custom-model.ts index d84c34b0f..734517ad5 100644 --- a/packages/coding-agent/examples/sdk/02-custom-model.ts +++ b/packages/coding-agent/examples/sdk/02-custom-model.ts @@ -5,11 +5,10 @@ */ import { getModel } from "@earendil-works/pi-ai/compat"; -import { AuthStorage, createAgentSession, ModelRegistry } from "@bastani/atomic"; +import { createAgentSession, ModelRuntime } from "@bastani/atomic"; -// Set up auth storage and model registry -const authStorage = AuthStorage.create(); -const modelRegistry = ModelRegistry.create(authStorage); +// ModelRuntime owns credential resolution and built-in/custom model discovery. +const modelRuntime = await ModelRuntime.create(); // Option 1: Find a specific built-in model by provider/id const opus = getModel("anthropic", "claude-opus-4-5"); @@ -17,14 +16,14 @@ if (opus) { console.log(`Found model: ${opus.provider}/${opus.id}`); } -// Option 2: Find model via registry (includes custom models from models.json) -const customModel = modelRegistry.find("my-provider", "my-model"); +// Option 2: Find a model from the runtime (includes custom models from models.json) +const customModel = modelRuntime.getModel("my-provider", "my-model"); if (customModel) { console.log(`Found custom model: ${customModel.provider}/${customModel.id}`); } -// Option 3: Pick from available models (have valid API keys) -const available = await modelRegistry.getAvailable(); +// Option 3: Pick from available models (have valid credentials) +const available = await modelRuntime.getAvailable(); console.log( "Available models:", available.map((m) => `${m.provider}/${m.id}`), @@ -34,8 +33,7 @@ if (available.length > 0) { const { session } = await createAgentSession({ model: available[0], thinkingLevel: "medium", // off, low, medium, high - authStorage, - modelRegistry, + modelRuntime, }); try { diff --git a/packages/coding-agent/examples/sdk/09-api-keys-and-oauth.ts b/packages/coding-agent/examples/sdk/09-api-keys-and-oauth.ts index a3ad2004c..c8bf9c9fd 100644 --- a/packages/coding-agent/examples/sdk/09-api-keys-and-oauth.ts +++ b/packages/coding-agent/examples/sdk/09-api-keys-and-oauth.ts @@ -1,52 +1,46 @@ /** * API Keys and OAuth * - * Configure API key resolution via AuthStorage and ModelRegistry. + * Configure credential and model resolution through ModelRuntime. */ -import { AuthStorage, createAgentSession, ModelRegistry, SessionManager } from "@bastani/atomic"; - -// Default: AuthStorage uses ~/.atomic/agent/auth.json (legacy ~/.pi/agent/auth.json also works) -// ModelRegistry loads built-in + custom models from ~/.atomic/agent/models.json (legacy ~/.pi/agent/models.json also works) -const authStorage = AuthStorage.create(); -const modelRegistry = ModelRegistry.create(authStorage); +import { createAgentSession, ModelRuntime, SessionManager } from "@bastani/atomic"; +// Default: createAgentSession builds a runtime from the active agent directory's +// auth.json and models.json. const { session: defaultAuthSession } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage, - modelRegistry, }); -console.log("Session with default auth storage and model registry"); +console.log("Session with default credential and model configuration"); defaultAuthSession.dispose(); -// Custom auth storage location -const customAuthStorage = AuthStorage.create("/tmp/my-app/auth.json"); -const customModelRegistry = ModelRegistry.create(customAuthStorage, "/tmp/my-app/models.json"); - +// Custom credential and model configuration locations +const customRuntime = await ModelRuntime.create({ + authPath: "/tmp/my-app/auth.json", + modelsPath: "/tmp/my-app/models.json", +}); const { session: customAuthSession } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage: customAuthStorage, - modelRegistry: customModelRegistry, + modelRuntime: customRuntime, }); -console.log("Session with custom auth storage location"); +console.log("Session with custom credential and model configuration"); customAuthSession.dispose(); // Runtime API key override (not persisted to disk) -authStorage.setRuntimeApiKey("anthropic", "sk-my-temp-key"); +const runtimeKeyRuntime = await ModelRuntime.create(); +await runtimeKeyRuntime.setRuntimeApiKey("anthropic", "sk-my-temp-key"); const { session: runtimeKeySession } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage, - modelRegistry, + modelRuntime: runtimeKeyRuntime, }); console.log("Session with runtime API key override"); runtimeKeySession.dispose(); // No models.json - only built-in models -const simpleRegistry = ModelRegistry.inMemory(authStorage); +const builtinsOnlyRuntime = await ModelRuntime.create({ modelsPath: null }); const { session: builtInModelsSession } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage, - modelRegistry: simpleRegistry, + modelRuntime: builtinsOnlyRuntime, }); console.log("Session with only built-in models"); builtInModelsSession.dispose(); diff --git a/packages/coding-agent/examples/sdk/12-full-control.ts b/packages/coding-agent/examples/sdk/12-full-control.ts index 5a5808965..a8bfe1553 100644 --- a/packages/coding-agent/examples/sdk/12-full-control.ts +++ b/packages/coding-agent/examples/sdk/12-full-control.ts @@ -6,26 +6,25 @@ import { getModel } from "@earendil-works/pi-ai/compat"; import { - AuthStorage, createAgentSession, createExtensionRuntime, - ModelRegistry, + ModelRuntime, type ResourceLoader, SessionManager, SettingsManager, } from "@bastani/atomic"; -// Custom auth storage location -const authStorage = AuthStorage.create("/tmp/my-agent/auth.json"); +// Custom credential location with no custom models.json +const modelRuntime = await ModelRuntime.create({ + authPath: "/tmp/my-agent/auth.json", + modelsPath: null, +}); // Runtime API key override (not persisted) if (process.env.MY_ANTHROPIC_KEY) { - authStorage.setRuntimeApiKey("anthropic", process.env.MY_ANTHROPIC_KEY); + await modelRuntime.setRuntimeApiKey("anthropic", process.env.MY_ANTHROPIC_KEY); } -// Model registry with no custom models.json -const modelRegistry = ModelRegistry.inMemory(authStorage); - const model = getModel("anthropic", "claude-sonnet-4-5"); if (!model) throw new Error("Model not found"); @@ -46,7 +45,7 @@ const resourceLoader: ResourceLoader = { getSystemPrompt: () => `You are a minimal assistant. Available: read, bash. Be concise.`, getAppendSystemPrompt: () => [], - extendResources: () => {}, + extendResources: async () => {}, reload: async () => {}, }; @@ -55,8 +54,7 @@ const { session } = await createAgentSession({ agentDir: "/tmp/my-agent", model, thinkingLevel: "off", - authStorage, - modelRegistry, + modelRuntime, resourceLoader, tools: ["read", "bash"], sessionManager: SessionManager.inMemory(cwd), diff --git a/packages/coding-agent/examples/sdk/README.md b/packages/coding-agent/examples/sdk/README.md index 75571e01a..0ed990774 100644 --- a/packages/coding-agent/examples/sdk/README.md +++ b/packages/coding-agent/examples/sdk/README.md @@ -34,50 +34,48 @@ bun examples/sdk/01-minimal.ts ```typescript import { getModel } from "@earendil-works/pi-ai/compat"; import { - AuthStorage, createAgentSession, DefaultResourceLoader, - ModelRegistry, + ModelRuntime, SessionManager, SettingsManager, } from "@bastani/atomic"; -// Auth and models setup -const authStorage = AuthStorage.create(); -const modelRegistry = ModelRegistry.create(authStorage); +// Credential and model setup +const modelRuntime = await ModelRuntime.create(); -// Minimal -const { session } = await createAgentSession({ authStorage, modelRegistry }); +// Minimal (omitting modelRuntime uses the active agent directory) +const { session } = await createAgentSession(); // Custom model const model = getModel("anthropic", "claude-opus-4-5"); -const { session } = await createAgentSession({ model, thinkingLevel: "high", authStorage, modelRegistry }); +const { session } = await createAgentSession({ model, thinkingLevel: "high", modelRuntime }); // Modify prompt const loader = new DefaultResourceLoader({ systemPromptOverride: (base) => `${base}\n\nBe concise.`, }); await loader.reload(); -const { session } = await createAgentSession({ resourceLoader: loader, authStorage, modelRegistry }); +const { session } = await createAgentSession({ resourceLoader: loader, modelRuntime }); // Read-only -const { session } = await createAgentSession({ tools: ["read", "search", "find", "ls"], authStorage, modelRegistry }); +const { session } = await createAgentSession({ tools: ["read", "search", "find", "ls"], modelRuntime }); // Defaults minus one tool -const { session } = await createAgentSession({ excludedTools: ["ask_user_question"], authStorage, modelRegistry }); +const { session } = await createAgentSession({ excludedTools: ["ask_user_question"], modelRuntime }); // In-memory const { session } = await createAgentSession({ sessionManager: SessionManager.inMemory(), - authStorage, - modelRegistry, + modelRuntime, }); // Full control -const customAuth = AuthStorage.create("/my/app/auth.json"); -customAuth.setRuntimeApiKey("anthropic", process.env.MY_KEY!); -const customRegistry = ModelRegistry.create(customAuth); - +const customRuntime = await ModelRuntime.create({ + authPath: "/my/app/auth.json", + modelsPath: "/my/app/models.json", +}); +await customRuntime.setRuntimeApiKey("anthropic", process.env.MY_KEY!); const resourceLoader = new DefaultResourceLoader({ systemPromptOverride: () => "You are helpful.", extensionFactories: [myExtension], @@ -89,8 +87,7 @@ await resourceLoader.reload(); const { session } = await createAgentSession({ model, - authStorage: customAuth, - modelRegistry: customRegistry, + modelRuntime: customRuntime, resourceLoader, tools: ["read", "bash", "my_tool"], customTools: [myTool], @@ -111,8 +108,7 @@ await session.prompt("Hello"); | Option | Default | Description | |--------|---------|-------------| -| `authStorage` | `AuthStorage.create()` | Credential storage | -| `modelRegistry` | `ModelRegistry.create(authStorage)` | Model registry | +| `modelRuntime` | Runtime created from the active agent directory | Credential and model runtime override | | `cwd` | `process.cwd()` | Working directory | | `agentDir` | `~/.atomic/agent` (legacy `~/.pi/agent` also works) | Config directory | | `model` | From settings/first available | Model to use | diff --git a/packages/coding-agent/src/cli/list-models.ts b/packages/coding-agent/src/cli/list-models.ts index 659ffabbe..5b2d44e62 100644 --- a/packages/coding-agent/src/cli/list-models.ts +++ b/packages/coding-agent/src/cli/list-models.ts @@ -6,7 +6,7 @@ import type { Api, Model } from "@earendil-works/pi-ai/compat"; import { fuzzyFilter } from "@earendil-works/pi-tui"; import chalk from "chalk"; import { formatNoModelsAvailableMessage } from "../core/auth-guidance.ts"; -import type { ModelRegistry } from "../core/model-registry.ts"; +import type { ModelRuntime } from "../core/model-runtime.ts"; /** * Format a number as human-readable (e.g., 200000 -> "200K", 1000000 -> "1M") @@ -26,13 +26,13 @@ function formatTokenCount(count: number): string { /** * List available models, optionally filtered by search pattern */ -export async function listModels(modelRegistry: ModelRegistry, searchPattern?: string): Promise { - const loadError = modelRegistry.getError(); +export async function listModels(modelRuntime: ModelRuntime, searchPattern?: string): Promise { + const loadError = modelRuntime.getError(); if (loadError) { console.error(chalk.yellow(`Warning: errors loading models.json:\n${loadError}`)); } - const models = modelRegistry.getAvailable(); + const models = modelRuntime.getAvailableSnapshot(); if (models.length === 0) { console.log(formatNoModelsAvailableMessage()); @@ -40,9 +40,9 @@ export async function listModels(modelRegistry: ModelRegistry, searchPattern?: s } // Apply fuzzy filter if search pattern provided - let filteredModels: Model[] = models; + let filteredModels: Model[] = [...models]; if (searchPattern) { - filteredModels = fuzzyFilter(models, searchPattern, (m) => `${m.provider} ${m.id}`); + filteredModels = fuzzyFilter([...models], searchPattern, (m) => `${m.provider} ${m.id}`); } if (filteredModels.length === 0) { diff --git a/packages/coding-agent/src/core/agent-session-accessors.ts b/packages/coding-agent/src/core/agent-session-accessors.ts index e7eafc71a..f5ea48be4 100644 --- a/packages/coding-agent/src/core/agent-session-accessors.ts +++ b/packages/coding-agent/src/core/agent-session-accessors.ts @@ -3,7 +3,7 @@ import type { AgentSessionInternalSurface as AgentSession } from "./agent-sessio export function installAgentSessionAccessors(prototype: AgentSession): void { Object.defineProperties(prototype, { orchestrationContext: { get() { return this._orchestrationContext; } }, - modelRegistry: { get() { return this._modelRegistry; } }, + modelRuntime: { get() { return this._modelRuntime; } }, state: { get() { return this.agent.state; } }, model: { get() { return this.agent.state.model; } }, thinkingLevel: { get() { return this.agent.state.thinkingLevel; } }, diff --git a/packages/coding-agent/src/core/agent-session-auto-compaction.ts b/packages/coding-agent/src/core/agent-session-auto-compaction.ts index d1dbda327..49b4096c9 100644 --- a/packages/coding-agent/src/core/agent-session-auto-compaction.ts +++ b/packages/coding-agent/src/core/agent-session-auto-compaction.ts @@ -366,11 +366,8 @@ export async function _runAutoCompaction(this: AgentSession, reason: "overflow" // before persistence or continuation, matching other provider-call failures. const result = await this._applyVerbatimCompaction({ resolvePlannerAuth: async (candidate) => { - const authResult = await this._modelRegistry.getApiKeyAndHeaders(candidate); - if (!authResult.ok || (!authResult.apiKey && !authResult.headers)) { - return undefined; - } - return { apiKey: authResult.apiKey, headers: authResult.headers, baseUrl: authResult.baseUrl }; + const authResult = await this._getRequiredRequestAuth(candidate); + return authResult.apiKey || authResult.headers ? authResult : undefined; }, abortController: this._autoCompactionAbortController, backupLabel: reason === "overflow" ? "overflow-auto-compact" : "auto-compact", diff --git a/packages/coding-agent/src/core/agent-session-compaction.ts b/packages/coding-agent/src/core/agent-session-compaction.ts index c04325275..bb909d058 100644 --- a/packages/coding-agent/src/core/agent-session-compaction.ts +++ b/packages/coding-agent/src/core/agent-session-compaction.ts @@ -134,7 +134,7 @@ export async function _applyVerbatimCompaction( // touches session model state or the main chat's attempted-key set. const fallback: FallbackPlannerContext = { fallbackModels: this._fallbackModels, - registry: this._modelRegistry, + registry: this._modelRuntime, preferredProvider: model.provider, sessionThinkingLevel: this.thinkingLevel, }; diff --git a/packages/coding-agent/src/core/agent-session-extension-bindings.ts b/packages/coding-agent/src/core/agent-session-extension-bindings.ts index 844c274b0..00ce0465a 100644 --- a/packages/coding-agent/src/core/agent-session-extension-bindings.ts +++ b/packages/coding-agent/src/core/agent-session-extension-bindings.ts @@ -119,7 +119,7 @@ export function _refreshCurrentModelFromRegistry(this: AgentSession): void { return; } - const refreshedModel = this._modelRegistry.find(currentModel.provider, currentModel.id); + const refreshedModel = this._modelRuntime.getModel(currentModel.provider, currentModel.id); if (!refreshedModel || refreshedModel === currentModel) { return; } @@ -212,7 +212,7 @@ export function _bindExtensionCore(this: AgentSession, runner: ExtensionRunner): refreshTools: () => this._refreshToolRegistry(), getCommands, setModel: async (model) => { - if (!this.modelRegistry.hasConfiguredAuth(model)) return false; + if (!this.modelRuntime.hasConfiguredAuth(model.provider)) return false; await this.setModel(model); return true; }, @@ -251,12 +251,12 @@ export function _bindExtensionCore(this: AgentSession, runner: ExtensionRunner): }, { registerProvider: (providerOrName, config) => { - if (typeof providerOrName === "string") this._modelRegistry.registerProvider(providerOrName, config!); - else this._modelRegistry.registerProvider(providerOrName); + if (typeof providerOrName === "string") this._modelRuntime.registerProvider(providerOrName, config!); + else this._modelRuntime.registerNativeProvider(providerOrName); this.refreshCurrentModelFromRegistry(); }, unregisterProvider: (name) => { - this._modelRegistry.unregisterProvider(name); + this._modelRuntime.unregisterProvider(name); this.refreshCurrentModelFromRegistry(); }, }, diff --git a/packages/coding-agent/src/core/agent-session-methods.ts b/packages/coding-agent/src/core/agent-session-methods.ts index 7dbbf8998..156351e97 100644 --- a/packages/coding-agent/src/core/agent-session-methods.ts +++ b/packages/coding-agent/src/core/agent-session-methods.ts @@ -19,7 +19,7 @@ import type { } from "./extensions/index.ts"; import type { PathMetadata } from "./package-manager.ts"; import type { BashExecutionMessage, CustomMessage } from "./messages.ts"; -import type { ModelRegistry } from "./model-registry.ts"; +import type { ModelRuntime } from "./model-runtime.ts"; import type { PromptTemplate } from "./prompt-templates.ts"; import type { ResourceLoader } from "./resource-loader.ts"; import type { BranchSummaryEntry, SessionManager } from "./session-manager.js"; @@ -89,7 +89,7 @@ export interface AgentSessionQueuePauseControl { export interface AgentSessionMethodSurface extends AgentSessionQueuePauseControl { readonly orchestrationContext: import("./extensions/index.ts").OrchestrationContext | undefined; - readonly modelRegistry: ModelRegistry; + readonly modelRuntime: ModelRuntime; readonly state: AgentState; readonly model: Model | undefined; readonly thinkingLevel: ThinkingLevel; @@ -260,7 +260,7 @@ export interface AgentSessionMethodSurface extends AgentSessionQueuePauseControl export interface AgentSessionPublicSurface extends Pick void; _extensionErrorListener?: ExtensionErrorListener; _extensionErrorUnsubscriber?: () => void; - _modelRegistry: ModelRegistry; + _modelRuntime: ModelRuntime; _toolRegistry: Map; _toolDefinitions: Map; _toolPromptSnippets: Map; diff --git a/packages/coding-agent/src/core/agent-session-models.ts b/packages/coding-agent/src/core/agent-session-models.ts index b77db4819..217fa40c9 100644 --- a/packages/coding-agent/src/core/agent-session-models.ts +++ b/packages/coding-agent/src/core/agent-session-models.ts @@ -10,26 +10,26 @@ import { THINKING_LEVELS, type ModelCycleResult } from "./agent-session-types.ts export async function _getRequiredRequestAuth(this: AgentSession, model: Model): Promise<{ apiKey?: string; headers?: ProviderHeaders; - baseUrl?: string; + env?: Record; }> { - const result = await this._modelRegistry.getApiKeyAndHeaders(model); - if (!result.ok) { - if (result.error.startsWith("No API key found")) { + let result; + try { + result = await this._modelRuntime.getAuth(model); + } catch (error) { + const cause = error instanceof Error ? error.cause : undefined; + if (cause instanceof Error && cause.message === "authHeader requires a resolved API key") { throw new Error(formatNoApiKeyFoundMessage(model.provider)); } - throw new Error(result.error); + throw error; } - if (result.apiKey || result.headers) { - return { apiKey: result.apiKey, headers: result.headers, baseUrl: result.baseUrl }; + if (result && (result.auth.apiKey || result.auth.headers)) { + const headers = result.auth.headers + ? Object.fromEntries(Object.entries(result.auth.headers).filter((entry): entry is [string, string] => entry[1] !== null)) + : undefined; + return { apiKey: result.auth.apiKey, headers, env: result.env }; } - - const isOAuth = this._modelRegistry.isUsingOAuth(model); - if (isOAuth) { - throw new Error( - `Authentication failed for "${model.provider}". ` + - `Credentials may have expired or network is unavailable. ` + - `Run '/login ${model.provider}' to re-authenticate.`, - ); + if (this._modelRuntime.isUsingOAuth(model.provider)) { + throw new Error(`Authentication failed for "${model.provider}". Credentials may have expired or network is unavailable. Run '/login ${model.provider}' to re-authenticate.`); } throw new Error(formatNoApiKeyFoundMessage(model.provider)); } @@ -79,7 +79,7 @@ export async function _emitModelSelect(this: AgentSession, */ export async function setModel(this: AgentSession, model: Model): Promise { - if (!this._modelRegistry.hasConfiguredAuth(model)) { + if (!this._modelRuntime.hasConfiguredAuth(model.provider)) { throw new Error(`No API key for ${model.provider}/${model.id}`); } @@ -114,7 +114,7 @@ export async function cycleModel(this: AgentSession, direction: "forward" | "bac export async function _cycleScopedModel(this: AgentSession, direction: "forward" | "backward"): Promise { - const scopedModels = this._scopedModels.filter((scoped) => this._modelRegistry.hasConfiguredAuth(scoped.model)); + const scopedModels = this._scopedModels.filter((scoped) => this._modelRuntime.hasConfiguredAuth(scoped.model.provider)); if (scopedModels.length <= 1) return undefined; const currentModel = this.model; @@ -147,7 +147,7 @@ export async function _cycleScopedModel(this: AgentSession, direction: "forward" export async function _cycleAvailableModel(this: AgentSession, direction: "forward" | "backward"): Promise { - const availableModels = await this._modelRegistry.getAvailable(); + const availableModels = await this._modelRuntime.getAvailableSnapshot(); if (availableModels.length <= 1) return undefined; const currentModel = this.model; diff --git a/packages/coding-agent/src/core/agent-session-post-tool-compaction.ts b/packages/coding-agent/src/core/agent-session-post-tool-compaction.ts index 22b8fcab4..337f9838c 100644 --- a/packages/coding-agent/src/core/agent-session-post-tool-compaction.ts +++ b/packages/coding-agent/src/core/agent-session-post-tool-compaction.ts @@ -49,10 +49,8 @@ export async function _preflightPostToolContext( try { const result = await this._applyVerbatimCompaction({ resolvePlannerAuth: async (candidate) => { - const auth = await this._modelRegistry.getApiKeyAndHeaders(candidate); - return auth.ok && (auth.apiKey || auth.headers) - ? { apiKey: auth.apiKey, headers: auth.headers, baseUrl: auth.baseUrl } - : undefined; + const auth = await this._getRequiredRequestAuth(candidate); + return auth.apiKey || auth.headers ? auth : undefined; }, abortController, backupLabel: "auto-compact", diff --git a/packages/coding-agent/src/core/agent-session-prompt.ts b/packages/coding-agent/src/core/agent-session-prompt.ts index 330d3cc2e..b0ae79eca 100644 --- a/packages/coding-agent/src/core/agent-session-prompt.ts +++ b/packages/coding-agent/src/core/agent-session-prompt.ts @@ -3,7 +3,7 @@ import type { AgentMessage } from "@earendil-works/pi-agent-core"; import type { ImageContent, TextContent } from "@earendil-works/pi-ai/compat"; import { runCallback } from "./callback-activity.ts"; import { ATOMIC_GUIDE_COMMAND_NAME, ATOMIC_GUIDE_HELP_CHOICES, atomicGuideModeForChoice, getAtomicGuideMessage, isAtomicGuideHelpChoice, normalizeAtomicGuideMode } from "./atomic-guide-command.ts"; -import { formatAuthStorageLoadFailedMessage, formatNoApiKeyFoundMessage, formatNoModelSelectedMessage, formatUnresolvedModelMessage } from "./auth-guidance.ts"; +import { formatNoApiKeyFoundMessage, formatNoModelSelectedMessage, formatUnresolvedModelMessage } from "./auth-guidance.ts"; import { expandPromptTemplate } from "./prompt-templates.ts"; import { stripFrontmatter } from "../utils/frontmatter.ts"; import type { AgentSessionInternalSurface as AgentSession } from "./agent-session-methods.ts"; @@ -126,21 +126,14 @@ export async function prompt(this: AgentSession, text: string, options?: PromptO throw new Error(formatUnresolvedModelMessage(this.model)); } - if (!this._modelRegistry.hasConfiguredAuth(this.model)) { + if (!this._modelRuntime.hasConfiguredAuth(this.model.provider)) { // A failed credential-store load (for example auth.json briefly locked // by a concurrent process, or invalid JSON) leaves an empty in-memory // credential set. That would otherwise be misreported here as // "No API key found" even though the credentials exist on disk. Surface // the real load failure instead so configured providers are not falsely // reported as unauthenticated (issue #1431). - const authLoadError = this._modelRegistry.authStorage.getLoadError(); - if (authLoadError) { - throw new Error( - formatAuthStorageLoadFailedMessage(this.model.provider, authLoadError), - { cause: authLoadError }, - ); - } - const isOAuth = this._modelRegistry.isUsingOAuth(this.model); + const isOAuth = this._modelRuntime.isUsingOAuth(this.model.provider); if (isOAuth) { throw new Error( `Authentication failed for "${this.model.provider}". ` + diff --git a/packages/coding-agent/src/core/agent-session-retry.ts b/packages/coding-agent/src/core/agent-session-retry.ts index d09c74904..ec0f37026 100644 --- a/packages/coding-agent/src/core/agent-session-retry.ts +++ b/packages/coding-agent/src/core/agent-session-retry.ts @@ -16,7 +16,7 @@ function modelLabel(model: Model | undefined): string { function resolveFallbackModel(this: AgentSession, value: string): { model: Model; thinkingLevel?: ThinkingLevel } | undefined { return resolveConfiguredFallbackModel( value, - this._modelRegistry, + this._modelRuntime, this.model?.provider ?? this.settingsManager.getDefaultProvider(), ); } diff --git a/packages/coding-agent/src/core/agent-session-runtime-auth.ts b/packages/coding-agent/src/core/agent-session-runtime-auth.ts index 8b45e14fa..f979a2c9f 100644 --- a/packages/coding-agent/src/core/agent-session-runtime-auth.ts +++ b/packages/coding-agent/src/core/agent-session-runtime-auth.ts @@ -2,36 +2,26 @@ import { ModelsError } from "@earendil-works/pi-ai"; import type { AgentSession } from "./agent-session.ts"; import { createAuthInteraction, - getLegacyOAuthProvider, - loginOAuthProvider, normalizeOAuthLoginError, OAuthLoginTransactionError, type AtomicOAuthLoginCallbacks, -} from "./oauth-provider-bridge.ts"; -export type { AtomicOAuthLoginCallbacks } from "./oauth-provider-bridge.ts"; +} from "./oauth-login.ts"; +export type { AtomicOAuthLoginCallbacks } from "./oauth-login.ts"; -/** Authenticate through provider-owned metadata while preserving extension OAuth. */ +/** Authenticate through provider-owned OAuth metadata. */ export async function loginRuntimeOAuthProvider( session: AgentSession, provider: string, callbacks: AtomicOAuthLoginCallbacks, ): Promise { - const registry = session.modelRegistry; - if (getLegacyOAuthProvider(provider)) { - const credential = await loginOAuthProvider(provider, callbacks); - try { - await registry.authStorage.asCredentialStore().modify(provider, async () => credential); - } catch (error) { - throw new OAuthLoginTransactionError(error); - } - return; - } + const runtime = session.modelRuntime; try { - await registry.login(provider, "oauth", createAuthInteraction(callbacks)); + await runtime.login(provider, "oauth", createAuthInteraction(callbacks)); } catch (error) { if (error instanceof ModelsError && error.code === "auth" && error.message.startsWith("Credential store modify failed")) { throw new OAuthLoginTransactionError(error); } throw normalizeOAuthLoginError(error, callbacks.signal, { includeActiveSignal: false }); } + } diff --git a/packages/coding-agent/src/core/agent-session-runtime.ts b/packages/coding-agent/src/core/agent-session-runtime.ts index d5825cc56..db1a5bee1 100644 --- a/packages/coding-agent/src/core/agent-session-runtime.ts +++ b/packages/coding-agent/src/core/agent-session-runtime.ts @@ -17,7 +17,7 @@ import { emitSessionShutdownEvent } from "./extensions/runner.ts"; import type { CreateAgentSessionResult } from "./sdk.ts"; import { assertSessionCwdExists } from "./session-cwd.ts"; import { SessionManager } from "./session-manager.ts"; -import type { AuthStatus } from "./auth-storage.ts"; +import type { AuthStatus } from "./provider-composer.ts"; import { loginRuntimeOAuthProvider, type AtomicOAuthLoginCallbacks } from "./agent-session-runtime-auth.ts"; import type { ModelFallbackReason } from "./model-resolver-types.ts"; @@ -158,14 +158,14 @@ export class AgentSessionRuntime { } async logoutProvider(provider: string): Promise { - const registry = this.session.modelRegistry; - await registry.authStorage.logoutAsync(provider); + const registry = this.session.modelRuntime; + await registry.logout(provider); await registry.refresh({ allowNetwork: false }); this.session.refreshCurrentModelFromRegistry(); return { provider, authStatus: registry.getProviderAuthStatus(provider), - models: registry.getAvailable(), + models: [...registry.getAvailableSnapshot()], scopedModels: [...this.session.scopedModels], }; } diff --git a/packages/coding-agent/src/core/agent-session-services.ts b/packages/coding-agent/src/core/agent-session-services.ts index 4a2158c19..324e746e2 100644 --- a/packages/coding-agent/src/core/agent-session-services.ts +++ b/packages/coding-agent/src/core/agent-session-services.ts @@ -1,12 +1,10 @@ import { join } from "node:path"; import type { ThinkingLevel } from "@earendil-works/pi-agent-core"; import type { Api, Model } from "@earendil-works/pi-ai/compat"; -import { getAgentDir, getModelsConfigPaths } from "../config.ts"; +import { getAgentDir } from "../config.ts"; import { resolvePath } from "../utils/paths.ts"; -import { AuthStorage } from "./auth-storage.ts"; import type { SessionStartEvent, ToolDefinition } from "./extensions/index.ts"; -import { ModelRegistry } from "./model-registry.ts"; -import type { ModelRuntime } from "./model-runtime.ts"; +import { ModelRuntime } from "./model-runtime.ts"; import { DefaultResourceLoader, type DefaultResourceLoaderOptions, @@ -40,9 +38,7 @@ export interface AgentSessionRuntimeDiagnostic { export interface CreateAgentSessionServicesOptions { cwd: string; agentDir?: string; - authStorage?: AuthStorage; settingsManager?: SettingsManager; - modelRegistry?: ModelRegistry; modelRuntime?: ModelRuntime; extensionFlagValues?: Map; resourceLoaderOptions?: Omit; @@ -63,7 +59,7 @@ export interface CreateAgentSessionFromServicesOptions { thinkingLevel?: ThinkingLevel; fallbackModels?: CreateAgentSessionOptions["fallbackModels"]; scopedModels?: Array<{ model: Model; thinkingLevel?: ThinkingLevel }>; - tools?: CreateAgentSessionOptions["tools"]; + tools?: string[]; excludedTools?: CreateAgentSessionOptions["excludedTools"]; noTools?: CreateAgentSessionOptions["noTools"]; customTools?: ToolDefinition[]; @@ -78,9 +74,8 @@ export interface CreateAgentSessionFromServicesOptions { export interface AgentSessionServices { cwd: string; agentDir: string; - authStorage: AuthStorage; + modelRuntime: ModelRuntime; settingsManager: SettingsManager; - modelRegistry: ModelRegistry; resourceLoader: ResourceLoader; diagnostics: AgentSessionRuntimeDiagnostic[]; } @@ -145,9 +140,14 @@ export async function createAgentSessionServices( ): Promise { const cwd = resolvePath(options.cwd); const agentDir = options.agentDir ? resolvePath(options.agentDir) : getAgentDir(); - const authStorageSpan = startTimingSpan("createAgentSessionServices.authStorage"); - const authStorage = options.modelRuntime?.authStorage ?? options.authStorage ?? AuthStorage.create(join(agentDir, "auth.json")); - endTimingSpan(authStorageSpan); + const modelRuntimeSpan = startTimingSpan("createAgentSessionServices.modelRuntime"); + const modelRuntime = + options.modelRuntime ?? + (await ModelRuntime.create({ + authPath: join(agentDir, "auth.json"), + modelsPath: join(agentDir, "models.json"), + })); + endTimingSpan(modelRuntimeSpan); const settingsSpan = startTimingSpan("createAgentSessionServices.settingsManager"); const settingsManager = options.settingsManager ?? SettingsManager.create(cwd, agentDir); endTimingSpan(settingsSpan); @@ -160,44 +160,31 @@ export async function createAgentSessionServices( const reloadSpan = startTimingSpan("createAgentSessionServices.resourceLoader.reload"); await resourceLoader.reload(options.resourceLoaderReloadOptions); endTimingSpan(reloadSpan); - const modelRegistrySpan = startTimingSpan("createAgentSessionServices.modelRegistry"); - const modelRegistry = options.modelRuntime?.modelRegistry ?? options.modelRegistry ?? ModelRegistry.create( - authStorage, - getModelsConfigPaths(cwd, agentDir, settingsManager.isProjectTrusted()), - agentDir, - ); - endTimingSpan(modelRegistrySpan); const diagnostics: AgentSessionRuntimeDiagnostic[] = []; const providerSpan = startTimingSpan("createAgentSessionServices.providerRegistrations"); const extensionsResult = resourceLoader.getExtensions(); for (const registration of extensionsResult.runtime.pendingProviderRegistrations) { try { - if ("provider" in registration) modelRegistry.registerProvider(registration.provider); - else modelRegistry.registerProvider(registration.name, registration.config); + if ("provider" in registration) modelRuntime.registerNativeProvider(registration.provider); + else modelRuntime.registerProvider(registration.name, registration.config); } catch (error) { const message = error instanceof Error ? error.message : String(error); - diagnostics.push({ - type: "error", - message: `Extension "${registration.extensionPath}" error: ${message}`, - }); + diagnostics.push({ type: "error", message: `Extension "${registration.extensionPath}" error: ${message}` }); } } extensionsResult.runtime.pendingProviderRegistrations = []; endTimingSpan(providerSpan); const catalogRestoreSpan = startTimingSpan("createAgentSessionServices.restoreModelCatalogs"); - await modelRegistry.refresh({ allowNetwork: false }); + await modelRuntime.refresh({ allowNetwork: false }); endTimingSpan(catalogRestoreSpan); - const flagSpan = startTimingSpan("createAgentSessionServices.extensionFlagValidation"); diagnostics.push(...applyExtensionFlagValues(resourceLoader, options.extensionFlagValues)); - endTimingSpan(flagSpan); return { cwd, agentDir, - authStorage, + modelRuntime, settingsManager, - modelRegistry, resourceLoader, diagnostics, }; @@ -216,16 +203,15 @@ export async function createAgentSessionFromServices( return createAgentSession({ cwd: options.services.cwd, agentDir: options.services.agentDir, - authStorage: options.services.authStorage, + modelRuntime: options.services.modelRuntime, settingsManager: options.services.settingsManager, - modelRegistry: options.services.modelRegistry, resourceLoader: options.services.resourceLoader, sessionManager: options.sessionManager, model: options.model, thinkingLevel: options.thinkingLevel, - fallbackModels: options.fallbackModels, scopedModels: options.scopedModels, tools: options.tools, + fallbackModels: options.fallbackModels, excludedTools: options.excludedTools, noTools: options.noTools, customTools: options.customTools, diff --git a/packages/coding-agent/src/core/agent-session-tool-registry.ts b/packages/coding-agent/src/core/agent-session-tool-registry.ts index 6562bc00d..8b7274a0b 100644 --- a/packages/coding-agent/src/core/agent-session-tool-registry.ts +++ b/packages/coding-agent/src/core/agent-session-tool-registry.ts @@ -3,6 +3,7 @@ import { ExtensionRunner, wrapRegisteredTools, type ToolDefinition } from "./ext import { createSyntheticSourceInfo } from "./source-info.ts"; import { createAllToolDefinitions, defaultToolNames } from "./tools/index.ts"; import { createToolDefinitionFromAgentTool } from "./tools/tool-definition-wrapper.ts"; +import { ModelRegistry } from "./model-registry.ts"; import type { AgentSessionInternalSurface as AgentSession } from "./agent-session-methods.ts"; import type { ToolDefinitionEntry } from "./agent-session-types.ts"; import { createSessionAsyncDeliveryHandler } from "./async/session-manager.js"; @@ -172,7 +173,7 @@ export function _buildRuntime(this: AgentSession, options: { extensionsResult.runtime, this._cwd, this.sessionManager, - this._modelRegistry, + new ModelRegistry(this._modelRuntime), this._orchestrationContext, ); if (this._extensionRunnerRef) { diff --git a/packages/coding-agent/src/core/agent-session-types.ts b/packages/coding-agent/src/core/agent-session-types.ts index b2281ca66..a036e4006 100644 --- a/packages/coding-agent/src/core/agent-session-types.ts +++ b/packages/coding-agent/src/core/agent-session-types.ts @@ -19,7 +19,7 @@ import type { ToolDefinition, } from "./extensions/index.ts"; import type { CustomMessage } from "./messages.ts"; -import type { ModelRegistry } from "./model-registry.ts"; +import type { ModelRuntime } from "./model-runtime.ts"; import type { ResourceLoader } from "./resource-loader.ts"; import type { SessionManager } from "./session-manager.ts"; import type { SettingsManager } from "./settings-manager.ts"; @@ -148,7 +148,7 @@ export interface AgentSessionConfig { fallbackModels?: string[]; resourceLoader: ResourceLoader; customTools?: ToolDefinition[]; - modelRegistry: ModelRegistry; + modelRuntime: ModelRuntime; initialActiveToolNames?: string[]; allowedToolNames?: string[]; excludedToolNames?: string[]; diff --git a/packages/coding-agent/src/core/agent-session.ts b/packages/coding-agent/src/core/agent-session.ts index a208a4478..ccccc6ec1 100644 --- a/packages/coding-agent/src/core/agent-session.ts +++ b/packages/coding-agent/src/core/agent-session.ts @@ -14,7 +14,7 @@ import type { import type { Api, AssistantMessage, Model } from "@earendil-works/pi-ai/compat"; import type { VerbatimCompactionResult } from "./compaction/index.ts"; import type { BashExecutionMessage, CustomMessage } from "./messages.ts"; -import type { ModelRegistry } from "./model-registry.ts"; +import type { ModelRuntime } from "./model-runtime.ts"; import type { ResourceLoader } from "./resource-loader.ts"; import type { SessionManager } from "./session-manager.js"; import type { SettingsManager } from "./settings-manager.ts"; @@ -133,7 +133,7 @@ export class AgentSession { protected _extensionShutdownHandler?: () => void; protected _extensionErrorListener?: ExtensionErrorListener; protected _extensionErrorUnsubscriber?: () => void; - protected _modelRegistry: ModelRegistry; + protected _modelRuntime: ModelRuntime; protected _toolRegistry: Map = new Map(); protected _toolDefinitions: Map = new Map(); protected _toolPromptSnippets: Map = new Map(); @@ -154,7 +154,7 @@ export class AgentSession { this._resourceLoader = config.resourceLoader; this._customTools = config.customTools ?? []; this._cwd = config.cwd; - this._modelRegistry = config.modelRegistry; + this._modelRuntime = config.modelRuntime; this._extensionRunnerRef = config.extensionRunnerRef; this._initialActiveToolNames = config.initialActiveToolNames; this._allowedToolNames = config.allowedToolNames ? new Set(config.allowedToolNames) : undefined; diff --git a/packages/coding-agent/src/core/auth-storage-backends.ts b/packages/coding-agent/src/core/auth-storage-backends.ts index dc6dd0675..a1fd0c53c 100644 --- a/packages/coding-agent/src/core/auth-storage-backends.ts +++ b/packages/coding-agent/src/core/auth-storage-backends.ts @@ -251,6 +251,7 @@ export class FileAuthStorageBackend implements AuthStorageBackend { export class InMemoryAuthStorageBackend implements AuthStorageBackend { private value: string | undefined; + private pendingWrite: Promise = Promise.resolve(); read(): string | undefined { return this.value; @@ -265,11 +266,21 @@ export class InMemoryAuthStorageBackend implements AuthStorageBackend { } async withLockAsync(fn: (current: string | undefined) => Promise>): Promise { - const { result, next } = await fn(this.value); - if (next !== undefined) { - this.value = next; + const previous = this.pendingWrite; + let release!: () => void; + this.pendingWrite = new Promise((resolve) => { + release = resolve; + }); + await previous; + try { + const { result, next } = await fn(this.value); + if (next !== undefined) { + this.value = next; + } + return result; + } finally { + release(); } - return result; } } diff --git a/packages/coding-agent/src/core/auth-storage.ts b/packages/coding-agent/src/core/auth-storage.ts index 5e95aa993..00b5d840b 100644 --- a/packages/coding-agent/src/core/auth-storage.ts +++ b/packages/coding-agent/src/core/auth-storage.ts @@ -1,73 +1,22 @@ /** - * Credential storage for API keys and OAuth tokens. - * Handles loading, saving, and refreshing credentials from auth.json. - * - * Uses file locking to prevent race conditions when multiple pi instances - * try to refresh tokens simultaneously. + * CredentialStore implementation backed by auth.json. + * Provider auth orchestration belongs to ModelRuntime and pi-ai Models. */ -import { - type Credential, - type CredentialInfo, - type CredentialStore, - type ModelAuth, - type OAuthCredential as PiOAuthCredential, - type OAuthCredentials, -} from "@earendil-works/pi-ai"; -import { findEnvKeys, getEnvApiKey } from "@earendil-works/pi-ai/compat"; +import type { Credential, CredentialInfo, CredentialStore } from "@earendil-works/pi-ai"; import { join } from "path"; import { getAgentConfigPaths, getAgentDir } from "../config.ts"; import { FileAuthStorageBackend, InMemoryAuthStorageBackend, type AuthStorageBackend } from "./auth-storage-backends.ts"; -import { - type AtomicOAuthLoginCallbacks, - getOAuthProviderDescriptors, - loginOAuthProvider, - oauthCredentialToAuth, - refreshOAuthProvider, -} from "./oauth-provider-bridge.ts"; import { resolveConfigValue } from "./resolve-config-value.ts"; -export type ApiKeyCredential = { - type: "api_key"; - key?: string; - /** Provider-scoped configuration persisted alongside the credential. */ - env?: Record; -}; - -export type OAuthCredential = { - type: "oauth"; -} & OAuthCredentials; - -export type AuthCredential = ApiKeyCredential | OAuthCredential; - -export type AuthStorageData = Record; - -export type AuthStatus = { - configured: boolean; - source?: "stored" | "runtime" | "environment" | "fallback" | "models_json_key" | "models_json_command"; - label?: string; -}; - export { FileAuthStorageBackend, InMemoryAuthStorageBackend, type AuthStorageBackend } from "./auth-storage-backends.ts"; -/** Read one persisted provider credential using Atomic's layered config paths. */ -export function readStoredCredential(providerId: string, authPath?: string | string[]): Credential | undefined { - return AuthStorage.create(authPath).get(providerId) as Credential | undefined; -} +export type AuthStorageData = Record; -/** - * Credential storage backed by a JSON file. - */ -export class AuthStorage { - private data: AuthStorageData = {}; - private runtimeOverrides: Map = new Map(); - private fallbackResolver?: (provider: string) => string | undefined; - private loadError: Error | null = null; - private errors: Error[] = []; - private credentialMutationTail: Promise = Promise.resolve(); - private credentialVersions = new Map(); - declare private storage: AuthStorageBackend; +export class AuthStorage implements CredentialStore { + private data: AuthStorageData = {}; + private storage: AuthStorageBackend; private constructor(storage: AuthStorageBackend) { this.storage = storage; @@ -91,39 +40,6 @@ export class AuthStorage { return AuthStorage.fromStorage(storage); } - /** - * Set a runtime API key override (not persisted to disk). - * Used for CLI --api-key flag. - */ - setRuntimeApiKey(provider: string, apiKey: string): void { - this.runtimeOverrides.set(provider, apiKey); - } - /** Read a non-persisted request override without resolving stored credentials. */ - getRuntimeApiKey(provider: string): string | undefined { - return this.runtimeOverrides.get(provider); - } - - - /** - * Remove a runtime API key override. - */ - removeRuntimeApiKey(provider: string): void { - this.runtimeOverrides.delete(provider); - } - - /** - * Set a fallback resolver for API keys not found in auth.json or env vars. - * Used for custom provider keys from models.json. - */ - setFallbackResolver(resolver: (provider: string) => string | undefined): void { - this.fallbackResolver = resolver; - } - - private recordError(error: unknown): void { - const normalizedError = error instanceof Error ? error : new Error(String(error)); - this.errors.push(normalizedError); - } - private parseStorageData(content: string | undefined): AuthStorageData { if (!content) { return {}; @@ -136,363 +52,86 @@ export class AuthStorage { */ reload(): void { try { - // Pure read: never take the exclusive write lock. Writers replace - // auth.json atomically, so a lock-free read always sees a complete - // snapshot. This keeps many concurrent sessions from starving each other - // on the lock and misreporting configured providers as unreadable under - // contention (issue #1431). - const content = this.readSnapshot(); + let content: string | undefined; + if (this.storage.read) { + content = this.storage.read(); + } else { + this.storage.withLock((current) => { + content = current; + return { result: undefined }; + }); + } this.data = this.parseStorageData(content); - this.loadError = null; - } catch (error) { - this.loadError = error as Error; - this.recordError(error); - } - } - - /** - * Read the credential snapshot, preferring the backend's lock-free `read()`. - * Falls back to a `withLock`-based read for custom backends that predate - * `read()` so the released `AuthStorageBackend` interface stays compatible. - */ - private readSnapshot(): string | undefined { - if (this.storage.read) { - return this.storage.read(); - } - let content: string | undefined; - this.storage.withLock((current) => { - content = current; - return { result: undefined }; - }); - return content; - } - - private bumpCredentialVersion(provider: string): void { - this.credentialVersions.set(provider, (this.credentialVersions.get(provider) ?? 0) + 1); - } - - private persistProviderChange(provider: string, credential: AuthCredential | undefined): void { - if (this.loadError) { - this.reload(); - if (this.loadError) throw this.loadError; - } - - try { - const persistedData = this.storage.withLock((current) => { - const currentData = this.parseStorageData(current); - const merged: AuthStorageData = { ...currentData }; - if (credential) { - merged[provider] = credential; - } else { - delete merged[provider]; - } - return { result: merged, next: JSON.stringify(merged, null, 2) }; - }); - this.data = persistedData; - this.bumpCredentialVersion(provider); - this.loadError = null; - } catch (error) { - this.recordError(error); - throw error instanceof Error ? error : new Error(String(error)); + } catch { + // Preserve the last valid in-memory snapshot. } } - /** - * Get credential for a provider. - */ - get(provider: string): AuthCredential | undefined { - return this.data[provider] ?? undefined; - } + async read(provider: string): Promise { + const credential = this.data[provider]; + if (credential?.type !== "api_key") return credential; + if (credential.key === undefined) return credential; + return { ...credential, key: resolveConfigValue(credential.key, credential.env) }; + } + + async modify( + provider: string, + fn: (current: Credential | undefined) => Promise, + ): Promise { + let persistedData: AuthStorageData | undefined; + const result = await this.storage.withLockAsync(async (content) => { + const currentData = this.parseStorageData(content); + const next = await fn(currentData[provider]); + if (next === undefined) { + persistedData = currentData; + return { result: currentData[provider] }; + } - /** - * Set credential for a provider. - */ - set(provider: string, credential: AuthCredential): void { - this.persistProviderChange(provider, credential); + const merged: AuthStorageData = { ...currentData, [provider]: next }; + persistedData = merged; + return { result: next, next: JSON.stringify(merged, null, 2) }; + }); + if (persistedData) this.data = persistedData; + return result; } - /** - * Remove credential for a provider. - */ - remove(provider: string): void { - if (!this.storage.deleteProvider) { - this.persistProviderChange(provider, undefined); + async delete(provider: string): Promise { + if (this.storage.deleteProviderAsync) { + const content = await this.storage.deleteProviderAsync(provider); + this.data = this.parseStorageData(content); return; } - try { - this.data = this.parseStorageData(this.storage.deleteProvider(provider)); - this.bumpCredentialVersion(provider); - this.loadError = null; - } catch (error) { - this.recordError(error); - throw error instanceof Error ? error : new Error(String(error)); - } - } - - /** - * List all providers with credentials. - */ - list(): string[] { - return Object.keys(this.data); - } - - /** - * Check if credentials exist for a provider in auth.json. - */ - has(provider: string): boolean { - return provider in this.data; - } - - /** - * Check if any form of auth is configured for a provider. - * Unlike getApiKey(), this doesn't refresh OAuth tokens. - */ - hasAuth(provider: string): boolean { - if (this.runtimeOverrides.has(provider)) return true; - if (this.data[provider]) return true; - if (findEnvKeys(provider)?.[0]) return true; - if (this.fallbackResolver?.(provider)) return true; - return false; - } - - /** - * Return auth status without exposing credential values or refreshing tokens. - */ - getAuthStatus(provider: string): AuthStatus { - if (this.data[provider]) { - return { configured: true, source: "stored" }; - } - - if (this.runtimeOverrides.has(provider)) { - return { configured: false, source: "runtime", label: "--api-key" }; - } - - const envKeys = findEnvKeys(provider); - if (envKeys?.[0]) { - return { configured: true, source: "environment", label: envKeys[0] }; - } - - if (this.fallbackResolver?.(provider)) { - return { configured: false, source: "fallback", label: "custom provider config" }; - } - - return { configured: false }; - } - - /** Return a copy of all persisted credentials. */ - getAll(): AuthStorageData { - return { ...this.data }; - } - - drainErrors(): Error[] { - const drained = [...this.errors]; - this.errors = []; - return drained; - } - - /** - * Returns the error from the most recent failed credential load, or null when - * the last reload succeeded. - * - * A non-null value means stored credentials could NOT be read — e.g. the auth - * file was temporarily locked by another process (ELOCKED) or contained - * invalid JSON — so an empty/absent credential set is NOT authoritative. - * Callers that would otherwise report "No API key found" should surface this - * load failure instead of treating the provider as unauthenticated - * (issue #1431). - */ - getLoadError(): Error | null { - return this.loadError; - } - - /** - * Login to an OAuth provider. - */ - async login(providerId: string, callbacks: AtomicOAuthLoginCallbacks): Promise { - const credential = await loginOAuthProvider(providerId, callbacks); - await this.asCredentialStore().modify(providerId, async () => credential); - } - - /** Logout through the preserved synchronous persistence path. */ - logout(provider: string): void { - this.remove(provider); + let persistedData: AuthStorageData | undefined; + await this.storage.withLockAsync(async (content) => { + const currentData = this.parseStorageData(content); + delete currentData[provider]; + persistedData = currentData; + return { result: undefined, next: JSON.stringify(currentData, null, 2) }; + }); + if (persistedData) this.data = persistedData; } - /** Serialized logout for callers that need persistence completion. */ - async logoutAsync(provider: string): Promise { - await this.asCredentialStore().delete(provider); + /** List credential metadata without resolving configured key values. */ + async list(): Promise { + return Object.entries(this.data).map(([providerId, credential]) => ({ providerId, type: credential.type })); } - /** - * Refresh OAuth token with backend locking to prevent race conditions. - * Multiple pi instances may try to refresh simultaneously when tokens expire. - */ - private async refreshOAuthTokenWithLock( - providerId: string, - ): Promise<{ auth: ModelAuth; newCredentials: OAuthCredentials } | null> { - return this.runCredentialMutation(async () => { - let synchronizedData: AuthStorageData | undefined; - const credentialVersion = this.credentialVersions.get(providerId) ?? 0; - const result = await this.storage.withLockAsync(async (current) => { - const currentData = this.parseStorageData(current); - synchronizedData = currentData; - const cred = currentData[providerId]; - if (cred?.type !== "oauth") return { result: null }; - - const oauthCredential: PiOAuthCredential = { ...cred, type: "oauth" }; - if (Date.now() < cred.expires) { - const auth = await oauthCredentialToAuth(providerId, oauthCredential); - return { result: auth ? { auth, newCredentials: cred } : null }; - } - - const refreshed = await refreshOAuthProvider(providerId, oauthCredential); - if (!refreshed) return { result: null }; - if ((this.credentialVersions.get(providerId) ?? 0) !== credentialVersion) { - synchronizedData = undefined; - return { result: null }; - } - const merged: AuthStorageData = { - ...currentData, - [providerId]: refreshed.credential, - }; - synchronizedData = merged; - return { - result: { auth: refreshed.auth, newCredentials: refreshed.credential }, - next: JSON.stringify(merged, null, 2), - }; - }); +} - if (synchronizedData) this.data = synchronizedData; - this.loadError = null; - return result; +/** Read one persisted provider credential using Atomic's layered auth.json paths. */ +export function readStoredCredential(providerId: string, authPath?: string | string[]): Credential | undefined { + const paths = authPath === undefined + ? getAgentConfigPaths("auth.json") + : Array.isArray(authPath) ? authPath : [authPath]; + const storage = new FileAuthStorageBackend(paths[0] ?? join(getAgentDir(), "auth.json"), paths); + try { + let credential: Credential | undefined; + storage.withLock((content) => { + credential = content ? (JSON.parse(content) as AuthStorageData)[providerId] : undefined; + return { result: undefined }; }); - } - - /** - * Get API key for a provider. - * Priority: - * 1. Runtime override (CLI --api-key) - * 2. API key from auth.json - * 3. OAuth token from auth.json (auto-refreshed with locking) - * 4. Environment variable - * 5. Fallback resolver (models.json custom providers) - */ - async getModelAuth(providerId: string, options?: { includeFallback?: boolean }): Promise { - const runtimeKey = this.runtimeOverrides.get(providerId); - if (runtimeKey) return { apiKey: runtimeKey }; - - const cred = this.data[providerId]; - if (cred?.type === "api_key" && cred.key) return { apiKey: resolveConfigValue(cred.key) }; - - if (cred?.type === "oauth") { - try { - if (Date.now() >= cred.expires) { - return (await this.refreshOAuthTokenWithLock(providerId))?.auth; - } - return await oauthCredentialToAuth(providerId, { ...cred, type: "oauth" }); - } catch (error) { - this.recordError(error); - this.reload(); - const updated = this.data[providerId]; - if (updated?.type === "oauth" && Date.now() < updated.expires) { - return oauthCredentialToAuth(providerId, { ...updated, type: "oauth" }); - } - return undefined; - } - } - - const envKey = getEnvApiKey(providerId); - if (envKey) return { apiKey: envKey }; - if (options?.includeFallback !== false) { - const fallback = this.fallbackResolver?.(providerId); - if (fallback) return { apiKey: fallback }; - } + return credential; + } catch { return undefined; } - - async getApiKey(providerId: string, options?: { includeFallback?: boolean }): Promise { - return (await this.getModelAuth(providerId, options))?.apiKey; - } - - /** - * Get all registered OAuth providers - */ - getOAuthProviders() { - return getOAuthProviderDescriptors(); - } - - private runCredentialMutation(operation: () => Promise): Promise { - const result = this.credentialMutationTail.then(operation, operation); - this.credentialMutationTail = result.then( - () => undefined, - () => undefined, - ); - return result; - } - - /** Private async adapter for pi-ai's provider-owned Models runtime. */ - asCredentialStore(): CredentialStore { - const runtimeCredential = (providerId: string): Credential | undefined => { - const key = this.runtimeOverrides.get(providerId); - return key === undefined ? undefined : { type: "api_key", key }; - }; - return { - read: async (providerId) => runtimeCredential(providerId) ?? (this.get(providerId) as Credential | undefined), - list: async (): Promise => { - const providers = new Set([...this.list(), ...this.runtimeOverrides.keys()]); - return [...providers].map((providerId) => ({ - providerId, - type: runtimeCredential(providerId)?.type ?? this.data[providerId].type, - })); - }, - modify: (providerId, fn) => - this.runCredentialMutation(async () => { - let synchronizedData: AuthStorageData | undefined; - try { - const result = await this.storage.withLockAsync(async (current) => { - const data = this.parseStorageData(current); - const next = await fn(data[providerId] as Credential | undefined); - if (next === undefined) { - synchronizedData = data; - return { result: data[providerId] as Credential | undefined }; - } - const merged = { ...data, [providerId]: next as AuthCredential }; - synchronizedData = merged; - return { result: next, next: JSON.stringify(merged, null, 2) }; - }); - if (synchronizedData) { - this.data = synchronizedData; - this.bumpCredentialVersion(providerId); - } - this.loadError = null; - return result; - } catch (error) { - this.recordError(error); - throw error; - } - }), - delete: (providerId) => - this.runCredentialMutation(async () => { - try { - if (this.storage.deleteProviderAsync) { - this.data = this.parseStorageData(await this.storage.deleteProviderAsync(providerId)); - } else { - let synchronizedData: AuthStorageData | undefined; - await this.storage.withLockAsync(async (current) => { - const data = this.parseStorageData(current); - delete data[providerId]; - synchronizedData = data; - return { result: undefined, next: JSON.stringify(data, null, 2) }; - }); - if (synchronizedData) this.data = synchronizedData; - } - this.bumpCredentialVersion(providerId); - this.loadError = null; - } catch (error) { - this.recordError(error); - throw error; - } - }), - }; - } } diff --git a/packages/coding-agent/src/core/extensions/loader-virtual-modules.ts b/packages/coding-agent/src/core/extensions/loader-virtual-modules.ts index 88b2e5d25..1082f3c82 100644 --- a/packages/coding-agent/src/core/extensions/loader-virtual-modules.ts +++ b/packages/coding-agent/src/core/extensions/loader-virtual-modules.ts @@ -25,7 +25,7 @@ async function loadVirtualModules(): Promise> { // "@earendil-works/pi-ai"`, so we load the compat module here and key it // under the root specifier below to keep every extension working unchanged. import("@earendil-works/pi-ai/compat"), - import("../oauth-compat.ts"), + import("@earendil-works/pi-ai/oauth"), // NOTE: This import works because loader.ts exports are NOT re-exported from index.ts, // avoiding a circular dependency while preserving the package-name extension import path. import("../../index.ts"), @@ -193,8 +193,8 @@ function getAliases(): Record { // upstream layout change moves these files, this join needs updating to // match the package's real dist paths. const piAiEntry = resolveWorkspaceOrImport("ai/dist/compat.js", "@earendil-works/pi-ai"); + const piAiOauthEntry = resolveWorkspaceOrImport("ai/dist/oauth.js", "@earendil-works/pi-ai"); const piAiProvidersEntry = resolveWorkspaceOrImport("ai/dist/providers/all.js", "@earendil-works/pi-ai"); - const piAiOauthEntry = path.resolve(__dirname, "../oauth-compat.js"); _aliases = { "@bastani/atomic": piCodingAgentEntry, diff --git a/packages/coding-agent/src/core/extensions/provider-types.ts b/packages/coding-agent/src/core/extensions/provider-types.ts index 23b031dcf..b3b2e046f 100644 --- a/packages/coding-agent/src/core/extensions/provider-types.ts +++ b/packages/coding-agent/src/core/extensions/provider-types.ts @@ -12,33 +12,6 @@ import type { ExtensionAPI } from "./api-types.ts"; import type { AtomicProviderCompat } from "../model-capabilities.ts"; export type { AtomicProviderCompat } from "../model-capabilities.ts"; -export interface ApiKeyAuthPrompt { - type: "text" | "secret"; - message: string; - placeholder?: string; -} - -export interface ApiKeyAuthInteraction { - prompt(prompt: ApiKeyAuthPrompt): Promise; - signal: AbortSignal; -} - -export interface ProviderApiKeyAuthContext { - env(name: string): Promise; -} - -export interface ProviderApiKeyAuthResult { - auth: { apiKey?: string; baseUrl?: string; headers?: Record }; - env?: Record; - source?: string; -} - -export interface ProviderApiKeyAuth { - name: string; - login(interaction: ApiKeyAuthInteraction): Promise; - check?(input: { ctx: ProviderApiKeyAuthContext; credential?: import("../auth-storage.ts").ApiKeyCredential }): Promise<{ type: "api_key"; source?: string } | undefined>; - resolve?(input: { ctx: ProviderApiKeyAuthContext; credential?: import("../auth-storage.ts").ApiKeyCredential }): Promise; -} /** Configuration for registering a provider via pi.registerProvider(). */ export interface ProviderConfig { /** Display name for the provider in UI. */ @@ -59,8 +32,6 @@ export interface ProviderConfig { models?: ProviderModelConfig[]; /** Refresh this provider's catalog. Successful results replace its extension-provided models. */ refreshModels?(context: RefreshModelsContext): Promise; - /** Optional provider-directed API-key authentication flow. */ - auth?: { apiKey?: ProviderApiKeyAuth }; /** OAuth provider for /login support. The `id` is set automatically from the provider name. */ oauth?: { /** Display name for the provider in login UI. */ diff --git a/packages/coding-agent/src/core/fallback-models.ts b/packages/coding-agent/src/core/fallback-models.ts index ea6a31b42..e1d3d26b6 100644 --- a/packages/coding-agent/src/core/fallback-models.ts +++ b/packages/coding-agent/src/core/fallback-models.ts @@ -13,11 +13,11 @@ import type { Api, Model } from "@earendil-works/pi-ai/compat"; const THINKING_SUFFIXES = ["off", "minimal", "low", "medium", "high", "xhigh", "max"] as const satisfies readonly ThinkingLevel[]; const THINKING_SUFFIX_SET: ReadonlySet = new Set(THINKING_SUFFIXES); -/** Minimal registry surface needed to resolve a configured fallback entry. */ +/** Minimal runtime surface needed to resolve a configured fallback entry. */ export interface FallbackModelLookup { - getAvailable(): Model[]; - find(provider: string, modelId: string): Model | undefined; - hasConfiguredAuth(model: Model): boolean; + getAvailableSnapshot(): readonly Model[]; + getModel(provider: string, modelId: string): Model | undefined; + hasConfiguredAuth(provider: string): boolean; } /** A configured fallback entry resolved to a concrete model. */ @@ -50,15 +50,15 @@ export function resolveFallbackModel( ): FallbackModelCandidate | undefined { const parsed = splitFallbackModel(value); if (!parsed.modelId.includes("/")) { - const available = lookup.getAvailable().filter((model) => model.id === parsed.modelId); + const available = lookup.getAvailableSnapshot().filter((model) => model.id === parsed.modelId); const model = available.find((candidate) => candidate.provider === preferredProvider) ?? (available.length === 1 ? available[0] : undefined); return model ? { model, thinkingLevel: parsed.thinkingLevel } : undefined; } const slash = parsed.modelId.indexOf("/"); const provider = parsed.modelId.slice(0, slash); const modelId = parsed.modelId.slice(slash + 1); - const model = lookup.find(provider, modelId); - if (!model || !lookup.hasConfiguredAuth(model)) return undefined; + const model = lookup.getModel(provider, modelId); + if (!model || !lookup.hasConfiguredAuth(model.provider)) return undefined; return { model, thinkingLevel: parsed.thinkingLevel }; } diff --git a/packages/coding-agent/src/core/model-registry-schemas.ts b/packages/coding-agent/src/core/model-config.ts similarity index 71% rename from packages/coding-agent/src/core/model-registry-schemas.ts rename to packages/coding-agent/src/core/model-config.ts index 0d613b92f..9ee6f4368 100644 --- a/packages/coding-agent/src/core/model-registry-schemas.ts +++ b/packages/coding-agent/src/core/model-config.ts @@ -1,6 +1,11 @@ +/** Immutable, credential-blind models.json snapshot. */ + +import { readFile } from "node:fs/promises"; import { type Static, Type } from "typebox"; import { Compile } from "typebox/compile"; import type { TLocalizedValidationError } from "typebox/error"; +import { stripJsonComments } from "../utils/json.ts"; +import { normalizePath } from "../utils/paths.ts"; const PercentileCutoffsSchema = Type.Object({ p50: Type.Optional(Type.Number()), @@ -56,6 +61,7 @@ const ThinkingLevelMapSchema = Type.Object({ xhigh: Type.Optional(ThinkingLevelMapValueSchema), max: Type.Optional(ThinkingLevelMapValueSchema), }); + const ChatTemplateKwargScalarSchema = Type.Union([Type.String(), Type.Number(), Type.Boolean(), Type.Null()]); const ChatTemplateKwargVariableSchema = Type.Object({ $var: Type.Union([Type.Literal("thinking.enabled"), Type.Literal("thinking.effort")]), @@ -63,28 +69,6 @@ const ChatTemplateKwargVariableSchema = Type.Object({ }); const ChatTemplateKwargSchema = Type.Union([ChatTemplateKwargScalarSchema, ChatTemplateKwargVariableSchema]); -const ModelCostRatesSchema = Type.Object({ - input: Type.Number(), - output: Type.Number(), - cacheRead: Type.Number(), - cacheWrite: Type.Number(), -}); -const ModelCostTierSchema = Type.Intersect([ - ModelCostRatesSchema, - Type.Object({ inputTokensAbove: Type.Number() }), -]); -const ModelCostSchema = Type.Intersect([ - ModelCostRatesSchema, - Type.Object({ tiers: Type.Optional(Type.Array(ModelCostTierSchema)) }), -]); -const ModelCostOverrideSchema = Type.Object({ - input: Type.Optional(Type.Number()), - output: Type.Optional(Type.Number()), - cacheRead: Type.Optional(Type.Number()), - cacheWrite: Type.Optional(Type.Number()), - tiers: Type.Optional(Type.Array(ModelCostTierSchema)), -}); - const OpenAICompletionsCompatSchema = Type.Object({ supportsStore: Type.Optional(Type.Boolean()), supportsDeveloperRole: Type.Optional(Type.Boolean()), @@ -113,10 +97,8 @@ const OpenAICompletionsCompatSchema = Type.Object({ cacheControlFormat: Type.Optional(Type.Literal("anthropic")), openRouterRouting: Type.Optional(OpenRouterRoutingSchema), vercelGatewayRouting: Type.Optional(VercelGatewayRoutingSchema), - supportsStrictMode: Type.Optional(Type.Boolean()), supportsOpenAIGrammarTools: Type.Optional(Type.Boolean()), - /** Atomic compatibility alias for supportsOpenAIGrammarTools. */ - supportsGrammarTools: Type.Optional(Type.Boolean()), + supportsStrictMode: Type.Optional(Type.Boolean()), sendSessionAffinityHeaders: Type.Optional(Type.Boolean()), deferredToolsMode: Type.Optional(Type.Literal("kimi")), sessionAffinityFormat: Type.Optional( @@ -133,8 +115,6 @@ const OpenAIResponsesCompatSchema = Type.Object({ supportsLongCacheRetention: Type.Optional(Type.Boolean()), supportsStrictMode: Type.Optional(Type.Boolean()), supportsOpenAIGrammarTools: Type.Optional(Type.Boolean()), - /** Atomic compatibility alias for supportsOpenAIGrammarTools. */ - supportsGrammarTools: Type.Optional(Type.Boolean()), supportsToolSearch: Type.Optional(Type.Boolean()), }); @@ -156,6 +136,21 @@ const ProviderCompatSchema = Type.Union([ AnthropicMessagesCompatSchema, ]); +const ModelCostRatesSchema = { + input: Type.Number(), + output: Type.Number(), + cacheRead: Type.Number(), + cacheWrite: Type.Number(), +}; +const ModelCostTierSchema = Type.Object({ + inputTokensAbove: Type.Number(), + ...ModelCostRatesSchema, +}); +const ModelCostSchema = Type.Object({ + ...ModelCostRatesSchema, + tiers: Type.Optional(Type.Array(ModelCostTierSchema)), +}); + const ModelDefinitionSchema = Type.Object({ id: Type.String({ minLength: 1 }), name: Type.Optional(Type.String({ minLength: 1 })), @@ -176,15 +171,21 @@ const ModelOverrideSchema = Type.Object({ reasoning: Type.Optional(Type.Boolean()), thinkingLevelMap: Type.Optional(ThinkingLevelMapSchema), input: Type.Optional(Type.Array(Type.Union([Type.Literal("text"), Type.Literal("image")]))), - cost: Type.Optional(ModelCostOverrideSchema), + cost: Type.Optional( + Type.Object({ + input: Type.Optional(Type.Number()), + output: Type.Optional(Type.Number()), + cacheRead: Type.Optional(Type.Number()), + cacheWrite: Type.Optional(Type.Number()), + tiers: Type.Optional(Type.Array(ModelCostTierSchema)), + }), + ), contextWindow: Type.Optional(Type.Number()), maxTokens: Type.Optional(Type.Number()), headers: Type.Optional(Type.Record(Type.String(), Type.String())), compat: Type.Optional(ProviderCompatSchema), }); -export type ModelOverride = Static; - const ProviderConfigSchema = Type.Object({ name: Type.Optional(Type.String({ minLength: 1 })), baseUrl: Type.Optional(Type.String({ minLength: 1 })), @@ -201,12 +202,14 @@ const ProviderConfigSchema = Type.Object({ const ModelsConfigSchema = Type.Object({ providers: Type.Record(Type.String(), ProviderConfigSchema), }); +const validateModelsConfig = Compile(ModelsConfigSchema); -export const validateModelsConfig = Compile(ModelsConfigSchema); +export type ModelsJsonModel = Static; +export type ModelsJsonModelOverride = Static; +export type ModelsJsonProvider = Static; +type ModelsJson = Static; -export type ModelsConfig = Static; - -export function formatValidationPath(error: TLocalizedValidationError): string { +function formatValidationPath(error: TLocalizedValidationError): string { if (error.keyword === "required") { const requiredProperties = (error.params as { requiredProperties?: string[] }).requiredProperties; const requiredProperty = requiredProperties?.[0]; @@ -219,8 +222,72 @@ export function formatValidationPath(error: TLocalizedValidationError): string { return path || "root"; } -export function stripJsonComments(input: string): string { - return input - .replace(/"(?:\\.|[^"\\])*"|\/\/[^\n]*/g, (m) => (m[0] === '"' ? m : "")) - .replace(/"(?:\\.|[^"\\])*"|,(\s*[}\]])/g, (m, tail) => tail ?? (m[0] === '"' ? m : "")); +function deepFreeze(value: T): T { + if (typeof value !== "object" || value === null || Object.isFrozen(value)) return value; + for (const child of Object.values(value)) deepFreeze(child); + return Object.freeze(value); +} + +/** One immutable load of models.json. */ +export class ModelConfig { + private readonly providers: ReadonlyMap; + private readonly error: string | undefined; + + private constructor(providers: ReadonlyMap, error?: string) { + this.providers = providers; + this.error = error; + } + + static async load(modelsJsonPath: string | undefined): Promise { + if (!modelsJsonPath) return new ModelConfig(new Map()); + const path = normalizePath(modelsJsonPath); + let content: string; + try { + content = await readFile(path, "utf-8"); + } catch (error) { + if ((error as NodeJS.ErrnoException).code === "ENOENT") return new ModelConfig(new Map()); + return new ModelConfig( + new Map(), + `Failed to load models.json: ${error instanceof Error ? error.message : error}\n\nFile: ${path}`, + ); + } + + let parsed: unknown; + try { + parsed = JSON.parse(stripJsonComments(content)); + } catch (error) { + return new ModelConfig( + new Map(), + `Failed to parse models.json: ${error instanceof Error ? error.message : error}\n\nFile: ${path}`, + ); + } + + if (!validateModelsConfig.Check(parsed)) { + const errors = + validateModelsConfig + .Errors(parsed) + .map((error) => ` - ${formatValidationPath(error)}: ${error.message}`) + .join("\n") || "Unknown schema error"; + return new ModelConfig(new Map(), `Invalid models.json schema:\n${errors}\n\nFile: ${path}`); + } + + const config = parsed as ModelsJson; + const providers = new Map(); + for (const [providerId, provider] of Object.entries(config.providers)) { + providers.set(providerId, deepFreeze(structuredClone(provider))); + } + return new ModelConfig(providers); + } + + getProvider(providerId: string): ModelsJsonProvider | undefined { + return this.providers.get(providerId); + } + + getProviderIds(): readonly string[] { + return [...this.providers.keys()]; + } + + getError(): string | undefined { + return this.error; + } } diff --git a/packages/coding-agent/src/core/model-registry-auth.ts b/packages/coding-agent/src/core/model-registry-auth.ts deleted file mode 100644 index 530f49e05..000000000 --- a/packages/coding-agent/src/core/model-registry-auth.ts +++ /dev/null @@ -1,149 +0,0 @@ -import type { ModelAuth, ProviderHeaders } from "@earendil-works/pi-ai"; -import type { Api, Model } from "@earendil-works/pi-ai/compat"; -import type { AuthStatus, AuthStorage } from "./auth-storage.ts"; -import type { ProviderRequestConfig, ResolvedRequestAuth } from "./model-registry-types.ts"; -import { - getConfigValueEnvVarNames, - isCommandConfigValue, - isConfigValueConfigured, - resolveConfigValueOrThrow, - resolveConfigValueUncached, - resolveHeadersOrThrow, -} from "./resolve-config-value.ts"; -import type { ProviderApiKeyAuthResult } from "./extensions/provider-types.ts"; - -async function resolveCustomApiKeyAuth( - provider: string, - authStorage: AuthStorage, - providerRequestConfigs: Map, -): Promise { - const resolver = providerRequestConfigs.get(provider)?.auth?.apiKey?.resolve; - if (!resolver) return undefined; - const stored = authStorage.get(provider); - const credential = stored?.type === "api_key" ? stored : undefined; - return resolver({ - credential, - ctx: { env: async (name) => credential?.env?.[name] ?? process.env[name] }, - }); -} - -export async function getProviderResolvedAuth( - provider: string, - authStorage: AuthStorage, - providerRequestConfigs: Map, -): Promise { - return resolveCustomApiKeyAuth(provider, authStorage, providerRequestConfigs); -} -function mergeHeaders( - base: ProviderHeaders | undefined, - override: ProviderHeaders | undefined, -): ProviderHeaders | undefined { - if (!base && !override) return undefined; - const merged: ProviderHeaders = {}; - for (const source of [base, override]) { - for (const [name, value] of Object.entries(source ?? {})) { - for (const existing of Object.keys(merged)) { - if (existing.toLowerCase() === name.toLowerCase()) delete merged[existing]; - } - merged[name] = value; - } - } - return Object.keys(merged).length > 0 ? merged : undefined; -} - -export async function getModelRequestAuth( - model: Model, - authStorage: AuthStorage, - providerRequestConfigs: Map, - modelRequestHeaders: Map>, - providerAuth?: ModelAuth, -): Promise { - try { - const providerConfig = providerRequestConfigs.get(model.provider); - const customAuth = await resolveCustomApiKeyAuth(model.provider, authStorage, providerRequestConfigs); - const storedAuth = await authStorage.getModelAuth(model.provider, { includeFallback: false }); - const apiKey = - providerAuth?.apiKey ?? - customAuth?.auth.apiKey ?? - storedAuth?.apiKey ?? - (providerConfig?.apiKey - ? resolveConfigValueOrThrow(providerConfig.apiKey, `API key for provider "${model.provider}"`) - : undefined); - - const providerHeaders = resolveHeadersOrThrow(providerConfig?.headers, `provider "${model.provider}"`); - const modelHeaders = resolveHeadersOrThrow( - modelRequestHeaders.get(`${model.provider}:${model.id}`), - `model "${model.provider}/${model.id}"`, - ); - - let headers = providerAuth - ? mergeHeaders(model.headers, providerAuth.headers) - : mergeHeaders(storedAuth?.headers, model.headers); - headers = mergeHeaders(headers, customAuth?.auth.headers); - headers = mergeHeaders(headers, providerHeaders); - headers = mergeHeaders(headers, modelHeaders); - - if (providerConfig?.authHeader) { - if (!apiKey) { - return { ok: false, error: `No API key found for "${model.provider}"` }; - } - headers = { ...headers, Authorization: `Bearer ${apiKey}` }; - } - - - return { - ok: true, - apiKey, - headers: headers && Object.keys(headers).length > 0 ? headers : undefined, - baseUrl: providerAuth?.baseUrl ?? customAuth?.auth.baseUrl ?? storedAuth?.baseUrl, - }; - } catch (error) { - return { - ok: false, - error: error instanceof Error ? error.message : String(error), - }; - } -} - -export function getProviderAuthStatusFromConfig( - provider: string, - authStorage: AuthStorage, - providerRequestConfigs: Map, -): AuthStatus { - const authStatus = authStorage.getAuthStatus(provider); - if (authStatus.source) { - return authStatus; - } - - const providerApiKey = providerRequestConfigs.get(provider)?.apiKey; - if (!providerApiKey) { - return authStatus; - } - - if (isCommandConfigValue(providerApiKey)) { - return { configured: true, source: "models_json_command" }; - } - - const envVarNames = getConfigValueEnvVarNames(providerApiKey); - if (envVarNames.length > 0) { - return isConfigValueConfigured(providerApiKey) - ? { configured: true, source: "environment", label: envVarNames.join(", ") } - : { configured: false }; - } - - return { configured: true, source: "models_json_key" }; -} - -export async function getApiKeyForProviderFromConfig( - provider: string, - authStorage: AuthStorage, - providerRequestConfigs: Map, -): Promise { - const apiKey = await authStorage.getApiKey(provider, { includeFallback: false }); - if (apiKey !== undefined) { - return apiKey; - } - - const providerApiKey = providerRequestConfigs.get(provider)?.apiKey; - return providerApiKey ? resolveConfigValueUncached(providerApiKey) : undefined; -} diff --git a/packages/coding-agent/src/core/model-registry-builtins.ts b/packages/coding-agent/src/core/model-registry-builtins.ts deleted file mode 100644 index 387ddad85..000000000 --- a/packages/coding-agent/src/core/model-registry-builtins.ts +++ /dev/null @@ -1,127 +0,0 @@ -import { - type Api, - getModels, - getProviders, - type BuiltinProvider, - type Model, - type OpenAICompletionsCompat, -} from "@earendil-works/pi-ai/compat"; -import { normalizeGrammarToolCapability } from "./model-capabilities.ts"; -import type { ModelOverride } from "./model-registry-schemas.ts"; -import type { ProviderCompat, ProviderOverride } from "./model-registry-types.ts"; - - -export function mergeCompat( - baseCompat: Model["compat"], - overrideCompat: ModelOverride["compat"] | Model["compat"], -): Model["compat"] | undefined { - if (!overrideCompat) return normalizeGrammarToolCapability(baseCompat); - - const base = baseCompat as ProviderCompat | undefined; - const override = overrideCompat as ProviderCompat; - const merged = { ...base, ...override } as ProviderCompat; - - const baseCompletions = base as OpenAICompletionsCompat | undefined; - const overrideCompletions = override as OpenAICompletionsCompat; - const mergedCompletions = merged as OpenAICompletionsCompat; - - if (baseCompletions?.openRouterRouting || overrideCompletions.openRouterRouting) { - mergedCompletions.openRouterRouting = { - ...baseCompletions?.openRouterRouting, - ...overrideCompletions.openRouterRouting, - }; - } - - if (baseCompletions?.vercelGatewayRouting || overrideCompletions.vercelGatewayRouting) { - mergedCompletions.vercelGatewayRouting = { - ...baseCompletions?.vercelGatewayRouting, - ...overrideCompletions.vercelGatewayRouting, - }; - } - - if (baseCompletions?.chatTemplateKwargs || overrideCompletions.chatTemplateKwargs) { - mergedCompletions.chatTemplateKwargs = { - ...baseCompletions?.chatTemplateKwargs, - ...overrideCompletions.chatTemplateKwargs, - }; - } - - return normalizeGrammarToolCapability(merged); -} - -export function applyModelOverride(model: Model, override: ModelOverride): Model { - const result = { ...model }; - - if (override.name !== undefined) result.name = override.name; - if (override.reasoning !== undefined) result.reasoning = override.reasoning; - if (override.thinkingLevelMap !== undefined) { - result.thinkingLevelMap = { ...model.thinkingLevelMap, ...override.thinkingLevelMap }; - } - if (override.input !== undefined) result.input = override.input as ("text" | "image")[]; - if (override.contextWindow !== undefined) result.contextWindow = override.contextWindow; - if (override.maxTokens !== undefined) result.maxTokens = override.maxTokens; - - if (override.cost) { - result.cost = { - input: override.cost.input ?? model.cost.input, - output: override.cost.output ?? model.cost.output, - cacheRead: override.cost.cacheRead ?? model.cost.cacheRead, - cacheWrite: override.cost.cacheWrite ?? model.cost.cacheWrite, - ...(override.cost.tiers !== undefined - ? { tiers: override.cost.tiers } - : model.cost.tiers !== undefined - ? { tiers: model.cost.tiers } - : {}), - }; - } - - result.compat = mergeCompat(model.compat, override.compat); - return result; -} - -export function loadBuiltInModels( - overrides: Map, - modelOverrides: Map>, - baseModels?: readonly Model[], -): Model[] { - const providers = baseModels - ? [...new Set(baseModels.map((model) => model.provider))] - : getProviders(); - return providers.flatMap((provider) => { - const providerModels = baseModels - ? baseModels.filter((model) => model.provider === provider) - : getModels(provider as BuiltinProvider) as Model[]; - const models = [...providerModels]; - const providerOverride = overrides.get(provider); - const perModelOverrides = modelOverrides.get(provider); - - return models.map((candidate) => { - let model: Model = { - ...candidate, - compat: normalizeGrammarToolCapability(candidate.compat), - }; - if (providerOverride) { - model = { - ...model, - baseUrl: providerOverride.baseUrl ?? model.baseUrl, - compat: mergeCompat(model.compat, providerOverride.compat), - }; - } - const modelOverride = perModelOverrides?.get(candidate.id); - return modelOverride ? applyModelOverride(model, modelOverride) : model; - }); - }); -} - -export function mergeCustomModels(builtInModels: Model[], customModels: Model[]): Model[] { - const merged = [...builtInModels]; - for (const customModel of customModels) { - const existingIndex = merged.findIndex((m) => m.provider === customModel.provider && m.id === customModel.id); - if (existingIndex >= 0) { - merged[existingIndex] = customModel; - } else { - merged.push(customModel); - } - } - return merged; -} diff --git a/packages/coding-agent/src/core/model-registry-custom-loader.ts b/packages/coding-agent/src/core/model-registry-custom-loader.ts deleted file mode 100644 index c0dd253e3..000000000 --- a/packages/coding-agent/src/core/model-registry-custom-loader.ts +++ /dev/null @@ -1,328 +0,0 @@ -import { type Api, type BuiltinProvider, getModels, getProviders, type Model } from "@earendil-works/pi-ai/compat"; -import { radiusProvider } from "@earendil-works/pi-ai/providers/all"; -import { existsSync, readFileSync } from "fs"; -import { validateContextWindowValue } from "./model-registry-validation.ts"; -import { mergeCompat, mergeCustomModels } from "./model-registry-builtins.ts"; -import { - formatValidationPath, - type ModelsConfig, - type ModelOverride, - stripJsonComments, - validateModelsConfig, -} from "./model-registry-schemas.ts"; -import type { - CustomModelsResult, - ModelRequestHeaderSource, - ProviderOverride, - ProviderRequestConfig, -} from "./model-registry-types.ts"; - -function modelRequestKey(provider: string, modelId: string): string { - return `${provider}:${modelId}`; -} - -function emptyCustomModelsResult(error?: string): CustomModelsResult { - return { - models: [], - overrides: new Map(), - modelOverrides: new Map(), - providerRequestConfigs: new Map(), - modelRequestHeaders: new Map(), - modelRequestHeaderSources: new Map(), - configuredProviders: new Map(), - error, - }; -} - -function mergeModelOverrides( - base: Map>, - incoming: Map>, -): Map> { - const merged = new Map>( - [...base].map(([providerName, overrides]) => [providerName, new Map(overrides)]), - ); - for (const [providerName, incomingOverrides] of incoming) { - const providerOverrides = new Map(merged.get(providerName) ?? []); - for (const [modelId, modelOverride] of incomingOverrides) { - providerOverrides.set(modelId, modelOverride); - } - merged.set(providerName, providerOverrides); - } - return merged; -} - -interface MergedModelRequestHeaders { - headers: Map>; - sources: Map; -} - -function mergeModelRequestHeaders( - base: CustomModelsResult, - incoming: CustomModelsResult, -): MergedModelRequestHeaders { - const headers = new Map(base.modelRequestHeaders); - const sources = new Map(base.modelRequestHeaderSources); - const invalidate = (key: string, source: ModelRequestHeaderSource): void => { - if (sources.get(key) !== source) return; - headers.delete(key); - sources.delete(key); - }; - - for (const model of incoming.models) { - invalidate(modelRequestKey(model.provider, model.id), "model"); - } - for (const [providerName, overrides] of incoming.modelOverrides) { - for (const modelId of overrides.keys()) { - invalidate(modelRequestKey(providerName, modelId), "modelOverride"); - } - } - for (const [key, value] of incoming.modelRequestHeaders) headers.set(key, value); - for (const [key, source] of incoming.modelRequestHeaderSources) sources.set(key, source); - return { headers, sources }; -} - -function mergeCustomModelResults(base: CustomModelsResult, incoming: CustomModelsResult): CustomModelsResult { - const mergedHeaders = mergeModelRequestHeaders(base, incoming); - return { - models: mergeCustomModels(base.models, incoming.models), - overrides: new Map([...base.overrides, ...incoming.overrides]), - modelOverrides: mergeModelOverrides(base.modelOverrides, incoming.modelOverrides), - providerRequestConfigs: new Map([...base.providerRequestConfigs, ...incoming.providerRequestConfigs]), - modelRequestHeaders: mergedHeaders.headers, - modelRequestHeaderSources: mergedHeaders.sources, - configuredProviders: new Map([...base.configuredProviders, ...incoming.configuredProviders]), - error: undefined, - }; -} - -function collectProviderRequestConfig( - providerName: string, - config: ProviderRequestConfig, - requestConfigs: Map, -): void { - if (!config.apiKey && !config.headers && !config.authHeader) return; - requestConfigs.set(providerName, { - apiKey: config.apiKey, - headers: config.headers, - authHeader: config.authHeader, - }); -} - -function collectModelHeaders( - providerName: string, - modelId: string, - headers: Record | undefined, - modelHeaders: Map>, - modelHeaderSources: Map, - source: ModelRequestHeaderSource, -): void { - if (!headers || Object.keys(headers).length === 0) return; - const key = modelRequestKey(providerName, modelId); - modelHeaders.set(key, headers); - modelHeaderSources.set(key, source); -} - -function validateConfig(config: ModelsConfig): void { - const builtInProviders = new Set(getProviders()); - - for (const [providerName, providerConfig] of Object.entries(config.providers)) { - const isBuiltIn = builtInProviders.has(providerName); - const hasProviderApi = !!providerConfig.api; - const models = providerConfig.models ?? []; - const hasModelOverrides = - providerConfig.modelOverrides && Object.keys(providerConfig.modelOverrides).length > 0; - - if (models.length === 0) { - if (!providerConfig.baseUrl && !providerConfig.headers && !providerConfig.compat && !hasModelOverrides) { - throw new Error( - `Provider ${providerName}: must specify "baseUrl", "headers", "compat", "modelOverrides", or "models".`, - ); - } - } else if (!isBuiltIn) { - if (!providerConfig.baseUrl) { - throw new Error(`Provider ${providerName}: "baseUrl" is required when defining custom models.`); - } - if (!providerConfig.apiKey) { - throw new Error(`Provider ${providerName}: "apiKey" is required when defining custom models.`); - } - } - - for (const modelDef of models) { - const hasModelApi = !!modelDef.api; - - if (!hasProviderApi && !hasModelApi && !isBuiltIn) { - throw new Error( - `Provider ${providerName}, model ${modelDef.id}: no "api" specified. Set at provider or model level.`, - ); - } - - if (!modelDef.id) throw new Error(`Provider ${providerName}: model missing "id"`); - if (modelDef.contextWindow !== undefined && validateContextWindowValue(modelDef.contextWindow)) { - throw new Error(`Provider ${providerName}, model ${modelDef.id}: invalid contextWindow`); - } - if (modelDef.maxTokens !== undefined && modelDef.maxTokens <= 0) { - throw new Error(`Provider ${providerName}, model ${modelDef.id}: invalid maxTokens`); - } - } - - for (const [modelId, modelOverride] of Object.entries(providerConfig.modelOverrides ?? {})) { - if (modelOverride.contextWindow !== undefined && validateContextWindowValue(modelOverride.contextWindow)) { - throw new Error(`Provider ${providerName}, model ${modelId}: invalid contextWindow`); - } - if (modelOverride.maxTokens !== undefined && modelOverride.maxTokens <= 0) { - throw new Error(`Provider ${providerName}, model ${modelId}: invalid maxTokens`); - } - } - } -} - -function parseModels( - config: ModelsConfig, - modelHeaders: Map>, - modelHeaderSources: Map, -): Model[] { - const models: Model[] = []; - const builtInProviders = new Set(getProviders()); - const builtInDefaultsCache = new Map(); - const getBuiltInDefaults = (providerName: string): { api: string; baseUrl: string } | undefined => { - if (!builtInProviders.has(providerName)) return undefined; - if (builtInDefaultsCache.has(providerName)) return builtInDefaultsCache.get(providerName); - const builtIn = getModels(providerName as BuiltinProvider) as Model[]; - if (builtIn.length === 0) return undefined; - const defaults = { api: builtIn[0].api, baseUrl: builtIn[0].baseUrl }; - builtInDefaultsCache.set(providerName, defaults); - return defaults; - }; - - for (const [providerName, providerConfig] of Object.entries(config.providers)) { - const modelDefs = providerConfig.models ?? []; - if (modelDefs.length === 0) continue; - - const builtInDefaults = getBuiltInDefaults(providerName); - - for (const modelDef of modelDefs) { - const api = modelDef.api ?? providerConfig.api ?? builtInDefaults?.api; - if (!api) continue; - - const baseUrl = modelDef.baseUrl ?? providerConfig.baseUrl ?? builtInDefaults?.baseUrl; - if (!baseUrl) continue; - - const compat = mergeCompat(providerConfig.compat, modelDef.compat); - collectModelHeaders(providerName, modelDef.id, modelDef.headers, modelHeaders, modelHeaderSources, "model"); - - const defaultCost = { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }; - const contextWindow = modelDef.contextWindow ?? 128000; - models.push({ - id: modelDef.id, - name: modelDef.name ?? modelDef.id, - api: api as Api, - provider: providerName, - baseUrl, - reasoning: modelDef.reasoning ?? false, - thinkingLevelMap: modelDef.thinkingLevelMap, - input: (modelDef.input ?? ["text"]) as ("text" | "image")[], - cost: modelDef.cost ?? defaultCost, - contextWindow, - maxTokens: modelDef.maxTokens ?? 16384, - headers: undefined, - compat, - } as Model); - } - } - - return models; -} - -function loadCustomModels(modelsJsonPath: string): CustomModelsResult { - if (!existsSync(modelsJsonPath)) { - return emptyCustomModelsResult(); - } - - try { - const content = readFileSync(modelsJsonPath, "utf-8"); - const parsed = JSON.parse(stripJsonComments(content)) as unknown; - - if (!validateModelsConfig.Check(parsed)) { - const errors = - validateModelsConfig - .Errors(parsed) - .map((error) => ` - ${formatValidationPath(error)}: ${error.message}`) - .join("\n") || "Unknown schema error"; - return emptyCustomModelsResult(`Invalid models.json schema:\n${errors}\n\nFile: ${modelsJsonPath}`); - } - - const config = parsed as ModelsConfig; - validateConfig(config); - - const overrides = new Map(); - const modelOverrides = new Map>(); - const providerRequestConfigs = new Map(); - const modelRequestHeaders = new Map>(); - const modelRequestHeaderSources = new Map(); - const configuredProviders = new Map>(); - - for (const [providerName, providerConfig] of Object.entries(config.providers)) { - if (providerConfig.oauth === "radius" && providerConfig.baseUrl) { - configuredProviders.set(providerName, radiusProvider({ - id: providerName, - name: providerConfig.name ?? providerName, - gateway: providerConfig.baseUrl.replace(/\/v1\/?$/u, ""), - })); - } - if ((providerConfig.baseUrl && providerConfig.oauth !== "radius") || providerConfig.compat) { - overrides.set(providerName, { - baseUrl: providerConfig.oauth === "radius" ? undefined : providerConfig.baseUrl, - compat: providerConfig.compat as Model["compat"], - }); - } - - collectProviderRequestConfig(providerName, providerConfig, providerRequestConfigs); - - if (providerConfig.modelOverrides) { - modelOverrides.set(providerName, new Map(Object.entries(providerConfig.modelOverrides))); - for (const [modelId, modelOverride] of Object.entries(providerConfig.modelOverrides)) { - collectModelHeaders( - providerName, - modelId, - modelOverride.headers, - modelRequestHeaders, - modelRequestHeaderSources, - "modelOverride", - ); - } - } - } - - return { - models: parseModels(config, modelRequestHeaders, modelRequestHeaderSources), - overrides, - modelOverrides, - providerRequestConfigs, - modelRequestHeaders, - modelRequestHeaderSources, - configuredProviders, - error: undefined, - }; - } catch (error) { - if (error instanceof SyntaxError) { - return emptyCustomModelsResult(`Failed to parse models.json: ${error.message}\n\nFile: ${modelsJsonPath}`); - } - return emptyCustomModelsResult( - `Failed to load models.json: ${error instanceof Error ? error.message : error}\n\nFile: ${modelsJsonPath}`, - ); - } -} - -export function loadCustomModelsFromPaths(modelsJsonPaths: string[]): CustomModelsResult { - let combined = emptyCustomModelsResult(); - const errors: string[] = []; - for (let i = modelsJsonPaths.length - 1; i >= 0; i--) { - const result = loadCustomModels(modelsJsonPaths[i]!); - if (result.error) { - errors.push(result.error); - continue; - } - combined = mergeCustomModelResults(combined, result); - } - return { ...combined, error: errors.length > 0 ? errors.join("\n\n") : undefined }; -} diff --git a/packages/coding-agent/src/core/model-registry-dynamic.ts b/packages/coding-agent/src/core/model-registry-dynamic.ts deleted file mode 100644 index 56b7a5423..000000000 --- a/packages/coding-agent/src/core/model-registry-dynamic.ts +++ /dev/null @@ -1,235 +0,0 @@ -import { - type Api, - getApiProvider, - type Model, - registerApiProvider, - type SimpleStreamOptions, - unregisterApiProviders, -} from "@earendil-works/pi-ai/compat"; -import { warnDeprecation } from "../utils/deprecation.ts"; -import { validateContextWindowValue } from "./model-registry-validation.ts"; -import { normalizeGrammarToolCapability } from "./model-capabilities.ts"; -import { applyModelOverride } from "./model-registry-builtins.ts"; -import type { DynamicProviderApplyInput, ProviderConfigInput } from "./model-registry-types.ts"; -import { registerLegacyOAuthProvider, unregisterLegacyOAuthProviders } from "./oauth-provider-bridge.ts"; -import { isLegacyEnvVarNameConfigValue } from "./resolve-config-value.ts"; - - -type RuntimeApiProvider = Parameters[0]; -type ActiveApiProvider = NonNullable>; - -interface RuntimeApiRegistration { - sourceId: string; - provider: RuntimeApiProvider; - activeProvider?: ActiveApiProvider; -} - -const runtimeApiRegistrations = new Map(); -const runtimeApiFallbacks = new Map(); - -function unregisterRuntimeApiProvider(sourceId: string): void { - for (const [api, registrations] of runtimeApiRegistrations) { - const index = registrations.findIndex((entry) => entry.sourceId === sourceId); - if (index < 0) continue; - const [removed] = registrations.splice(index, 1); - if (removed && getApiProvider(api) === removed.activeProvider) { - unregisterApiProviders(sourceId); - const previous = registrations.at(-1); - if (previous) { - registerApiProvider(previous.provider, previous.sourceId); - previous.activeProvider = getApiProvider(api); - } else { - const fallback = runtimeApiFallbacks.get(api); - if (fallback) registerApiProvider(fallback, `atomic:restored-api:${api}`); - } - } - if (registrations.length === 0) { - runtimeApiRegistrations.delete(api); - runtimeApiFallbacks.delete(api); - } - } -} - -function registerRuntimeApiProvider(provider: RuntimeApiProvider, sourceId: string): void { - unregisterRuntimeApiProvider(sourceId); - const registrations = runtimeApiRegistrations.get(provider.api) ?? []; - if (registrations.length === 0) { - const fallback = getApiProvider(provider.api); - if (fallback) runtimeApiFallbacks.set(provider.api, fallback); - } - const registration: RuntimeApiRegistration = { sourceId, provider }; - registrations.push(registration); - runtimeApiRegistrations.set(provider.api, registrations); - registerApiProvider(provider, sourceId); - registration.activeProvider = getApiProvider(provider.api); -} - -export function unregisterProviderRuntime(sourceId: string): void { - unregisterRuntimeApiProvider(sourceId); - unregisterLegacyOAuthProviders(sourceId); -} -function migrateLegacyRegisterProviderConfigValue(providerName: string, field: string, value: string): string { - if (!isLegacyEnvVarNameConfigValue(value) || process.env[value] === undefined) return value; - warnDeprecation( - `registerProvider("${providerName}") ${field} value "${value}" is treated as a legacy environment variable reference. This will no longer be detected as an environment variable reference in a future release. Pass "$${value}" instead.`, - ); - return `$${value}`; -} - -function migrateLegacyRegisterProviderHeaders( - providerName: string, - field: string, - headers: Record | undefined, -): Record | undefined { - if (!headers) return undefined; - let migratedHeaders: Record | undefined; - for (const [key, value] of Object.entries(headers)) { - const migratedValue = migrateLegacyRegisterProviderConfigValue(providerName, `${field} header "${key}"`, value); - if (migratedValue === value) continue; - migratedHeaders ??= { ...headers }; - migratedHeaders[key] = migratedValue; - } - return migratedHeaders ?? headers; -} - -export function migrateLegacyRegisterProviderConfigValues( - providerName: string, - config: ProviderConfigInput, -): ProviderConfigInput { - let migratedConfig: ProviderConfigInput | undefined; - - const setMigratedConfigValue = ( - key: TKey, - value: ProviderConfigInput[TKey], - ) => { - migratedConfig ??= { ...config }; - migratedConfig[key] = value; - }; - - if (config.apiKey) { - const apiKey = migrateLegacyRegisterProviderConfigValue(providerName, "apiKey", config.apiKey); - if (apiKey !== config.apiKey) { - setMigratedConfigValue("apiKey", apiKey); - } - } - - const headers = migrateLegacyRegisterProviderHeaders(providerName, "headers", config.headers); - if (headers !== config.headers) { - setMigratedConfigValue("headers", headers); - } - - if (config.models) { - let models: ProviderConfigInput["models"] | undefined; - for (let index = 0; index < config.models.length; index++) { - const model = config.models[index]; - const modelHeaders = migrateLegacyRegisterProviderHeaders( - providerName, - `model "${model.id}" headers`, - model.headers, - ); - if (modelHeaders === model.headers) continue; - models ??= [...config.models]; - models[index] = { ...model, headers: modelHeaders }; - } - if (models) { - setMigratedConfigValue("models", models); - } - } - - return migratedConfig ?? config; -} - -export function validateProviderConfig(providerName: string, config: ProviderConfigInput): void { - if (config.streamSimple && !config.api) { - throw new Error(`Provider ${providerName}: "api" is required when registering streamSimple.`); - } - - if (!config.models || config.models.length === 0) { - return; - } - - if (!config.baseUrl) { - throw new Error(`Provider ${providerName}: "baseUrl" is required when defining models.`); - } - if (!config.apiKey && !config.oauth && !config.auth?.apiKey) { - throw new Error(`Provider ${providerName}: "apiKey", "oauth", or "auth.apiKey" is required when defining models.`); - } - - for (const modelDef of config.models) { - const api = modelDef.api || config.api; - if (!api) { - throw new Error(`Provider ${providerName}, model ${modelDef.id}: no "api" specified.`); - } - if (validateContextWindowValue(modelDef.contextWindow)) { - throw new Error(`Provider ${providerName}, model ${modelDef.id}: invalid contextWindow`); - } - if (modelDef.maxTokens <= 0) { - throw new Error(`Provider ${providerName}, model ${modelDef.id}: invalid maxTokens`); - } - } -} - -export function applyProviderConfigToModels(input: DynamicProviderApplyInput): Model[] { - const { providerName, registrationSource, config, authStorage, modelOverrides, storeProviderRequestConfig, storeModelHeaders } = input; - let models = input.models; - - if (input.registerRuntime && config.oauth) registerLegacyOAuthProvider(providerName, config.oauth, registrationSource); - - if (input.registerRuntime && config.streamSimple) { - const streamSimple = config.streamSimple; - registerRuntimeApiProvider({ - api: config.api!, - stream: (model, context, options) => streamSimple(model, context, options as SimpleStreamOptions), - streamSimple, - }, registrationSource); - } - - storeProviderRequestConfig(providerName, config); - - if (config.models && config.models.length > 0) { - models = models.filter((m) => m.provider !== providerName); - - for (const modelDef of config.models) { - const api = modelDef.api || config.api; - const modelOverride = modelOverrides.get(providerName)?.get(modelDef.id); - storeModelHeaders(providerName, modelDef.id, { - ...modelDef.headers, - ...modelOverride?.headers, - }); - - const model = { - id: modelDef.id, - name: modelDef.name, - api: api as Api, - provider: providerName, - baseUrl: modelDef.baseUrl ?? config.baseUrl!, - reasoning: modelDef.reasoning, - thinkingLevelMap: modelDef.thinkingLevelMap, - input: modelDef.input as ("text" | "image")[], - cost: modelDef.cost, - contextWindow: modelDef.contextWindow, - maxTokens: modelDef.maxTokens, - headers: undefined, - compat: normalizeGrammarToolCapability(modelDef.compat), - } as Model; - models.push(modelOverride ? applyModelOverride(model, modelOverride) : model); - } - - if (config.oauth?.modifyModels) { - const cred = authStorage.get(providerName); - if (cred?.type === "oauth") { - models = config.oauth.modifyModels(models, cred); - } - } - } else if (config.baseUrl || config.headers) { - models = models.map((m) => { - if (m.provider !== providerName) return m; - return { - ...m, - baseUrl: config.baseUrl ?? m.baseUrl, - }; - }); - } - - return models; -} diff --git a/packages/coding-agent/src/core/model-registry-extension-refresh.ts b/packages/coding-agent/src/core/model-registry-extension-refresh.ts deleted file mode 100644 index 0ad81b8a3..000000000 --- a/packages/coding-agent/src/core/model-registry-extension-refresh.ts +++ /dev/null @@ -1,49 +0,0 @@ -import type { Credential, RefreshModelsContext } from "@earendil-works/pi-ai"; -import type { ProviderConfigInput } from "./model-registry-types.ts"; - -type ProviderModels = NonNullable; - -export interface ExtensionRefreshInput { - config: ProviderConfigInput; - credential: Credential | undefined; - store: RefreshModelsContext["store"]; - allowNetwork: boolean; - force?: boolean; - signal: AbortSignal; - isCurrent(): boolean; -} - -export interface ExtensionRefreshResult { - models?: ProviderModels; - error?: Error; -} - -function toError(error: unknown): Error { - return error instanceof Error ? error : new Error(String(error)); -} - -/** - * Refresh an extension provider, falling back to its validated persisted catalog - * when an online refresh fails. The online error remains authoritative. - */ -export async function refreshExtensionProvider(input: ExtensionRefreshInput): Promise { - const refresh = (allowNetwork: boolean, force: boolean | undefined) => input.config.refreshModels!({ - credential: input.credential, - store: input.store, - allowNetwork, - force, - signal: input.signal, - }); - try { - return { models: await refresh(input.allowNetwork, input.force) }; - } catch (error) { - const result: ExtensionRefreshResult = { error: toError(error) }; - if (!input.allowNetwork || !input.isCurrent()) return result; - try { - result.models = await refresh(false, false); - } catch { - // Preserve the original online error; cache recovery is best-effort. - } - return result; - } -} diff --git a/packages/coding-agent/src/core/model-registry-loader.ts b/packages/coding-agent/src/core/model-registry-loader.ts deleted file mode 100644 index 2afab49ee..000000000 --- a/packages/coding-agent/src/core/model-registry-loader.ts +++ /dev/null @@ -1,50 +0,0 @@ -import type { Api, Model } from "@earendil-works/pi-ai/compat"; -import type { AuthStorage } from "./auth-storage.ts"; -import { loadBuiltInModels, mergeCustomModels } from "./model-registry-builtins.ts"; -import { loadCustomModelsFromPaths } from "./model-registry-custom-loader.ts"; -import type { ModelRegistryLoadResult } from "./model-registry-types.ts"; - -const OPENAI_COMPATIBLE_APIS = new Set(["openai-completions", "openai-responses"]); - -export function loadModelRegistryModels( - authStorage: AuthStorage, - modelsJsonPaths: string[], - baseModels?: readonly Model[], -): ModelRegistryLoadResult { - const { - models: customModels, - overrides, - modelOverrides, - providerRequestConfigs, - modelRequestHeaders, - configuredProviders, - error, - } = loadCustomModelsFromPaths(modelsJsonPaths); - - const builtInModels = loadBuiltInModels(overrides, modelOverrides, baseModels); - const builtInProviders = new Set(builtInModels.map((model) => model.provider)); - const customOpenAICompatibleProviders = new Set( - customModels - .filter((model) => !builtInProviders.has(model.provider) && OPENAI_COMPATIBLE_APIS.has(model.api)) - .map((model) => model.provider), - ); - let combined: Model[] = mergeCustomModels(builtInModels, customModels); - - for (const oauthProvider of authStorage.getOAuthProviders()) { - const cred = authStorage.get(oauthProvider.id); - if (cred?.type === "oauth" && oauthProvider.modifyModels) { - combined = oauthProvider.modifyModels(combined, cred); - } - } - - return { - modelOverrides, - models: combined, - providerRequestConfigs, - modelRequestHeaders, - configuredProviders, - builtInProviders, - customOpenAICompatibleProviders, - loadError: error, - }; -} diff --git a/packages/coding-agent/src/core/model-registry-types.ts b/packages/coding-agent/src/core/model-registry-types.ts deleted file mode 100644 index ad7e49b85..000000000 --- a/packages/coding-agent/src/core/model-registry-types.ts +++ /dev/null @@ -1,103 +0,0 @@ -import type { - Api, - AssistantMessageEventStream, - Context, - Model, - SimpleStreamOptions, -} from "@earendil-works/pi-ai/compat"; -import type { Provider, ProviderHeaders } from "@earendil-works/pi-ai"; -import type { LegacyOAuthProvider } from "./oauth-provider-bridge.ts"; -import type { RefreshModelsContext } from "@earendil-works/pi-ai"; -import type { AuthStorage } from "./auth-storage.ts"; -import type { ModelOverride } from "./model-registry-schemas.ts"; -import type { ProviderApiKeyAuth } from "./extensions/provider-types.ts"; -import type { AtomicProviderCompat } from "./model-capabilities.ts"; - -export interface ProviderOverride { - baseUrl?: string; - compat?: AtomicProviderCompat; -} - -export interface ProviderRequestConfig { - apiKey?: string; - headers?: Record; - authHeader?: boolean; - auth?: { apiKey?: ProviderApiKeyAuth }; -} - -export type ResolvedRequestAuth = - | { - ok: true; - apiKey?: string; - headers?: ProviderHeaders; - baseUrl?: string; - } - | { - ok: false; - error: string; - }; - -export type ModelRequestHeaderSource = "model" | "modelOverride"; - -export interface CustomModelsResult { - models: Model[]; - overrides: Map; - modelOverrides: Map>; - providerRequestConfigs: Map; - modelRequestHeaders: Map>; - modelRequestHeaderSources: Map; - configuredProviders: Map; - error: string | undefined; -} - -export interface ModelRegistryLoadResult { - models: Model[]; - modelOverrides: Map>; - providerRequestConfigs: Map; - modelRequestHeaders: Map>; - configuredProviders: Map; - builtInProviders: Set; - customOpenAICompatibleProviders: Set; - loadError: string | undefined; -} - -export interface DynamicProviderApplyInput { - providerName: string; - registrationSource: string; - registerRuntime?: boolean; - config: ProviderConfigInput; - models: Model[]; - modelOverrides: Map>; - authStorage: AuthStorage; - storeProviderRequestConfig: (providerName: string, config: ProviderRequestConfig) => void; - storeModelHeaders: (providerName: string, modelId: string, headers?: Record) => void; -} - -export type ProviderCompat = AtomicProviderCompat; - -export interface ProviderConfigInput { - name?: string; - baseUrl?: string; - apiKey?: string; - api?: Api; - streamSimple?: (model: Model, context: Context, options?: SimpleStreamOptions) => AssistantMessageEventStream; - headers?: Record; - authHeader?: boolean; - refreshModels?(context: RefreshModelsContext): Promise>; - auth?: { apiKey?: ProviderApiKeyAuth }; - oauth?: LegacyOAuthProvider; - models?: Array<{ - id: string; - name: string; - api?: Api; - baseUrl?: string; - reasoning: boolean; - thinkingLevelMap?: Model["thinkingLevelMap"]; - input: ("text" | "image")[]; - cost: Model["cost"]; - contextWindow: number; - maxTokens: number; - headers?: Record; - compat?: AtomicProviderCompat; - }>; -} diff --git a/packages/coding-agent/src/core/model-registry-validation.ts b/packages/coding-agent/src/core/model-registry-validation.ts deleted file mode 100644 index 6a96def18..000000000 --- a/packages/coding-agent/src/core/model-registry-validation.ts +++ /dev/null @@ -1,13 +0,0 @@ -/** Shared value validation for custom/dynamic model-registry definitions. */ - -function isPositiveInteger(value: number): boolean { - return Number.isFinite(value) && Number.isInteger(value) && value > 0; -} - -/** - * Validate a model `contextWindow` token count. Returns an error string when the - * value is not a positive integer, or `undefined` when it is valid. - */ -export function validateContextWindowValue(value: number): string | undefined { - return isPositiveInteger(value) ? undefined : "Context window must be a positive integer token count"; -} diff --git a/packages/coding-agent/src/core/model-registry.ts b/packages/coding-agent/src/core/model-registry.ts index 4410dd3ec..4d8d82f5f 100644 --- a/packages/coding-agent/src/core/model-registry.ts +++ b/packages/coding-agent/src/core/model-registry.ts @@ -1,483 +1,145 @@ -/** - * Model registry - manages built-in and custom models, provides API key resolution. - */ - -import { - createModels, - type AuthInteraction, - type AuthResult, - type AuthType, - type Credential, - type CredentialStore, - type ModelAuth, - type ModelsRefreshResult, - type MutableModels, - type Provider, -} from "@earendil-works/pi-ai"; -import { builtinProviders, getBuiltinModelDataGeneratedAt } from "@earendil-works/pi-ai/providers/all"; -import { type Api, type Model } from "@earendil-works/pi-ai/compat"; -import { dirname, join } from "node:path"; -import { isDeepStrictEqual } from "node:util"; -import { getAgentConfigPaths } from "../config.ts"; -import { normalizePath } from "../utils/paths.ts"; -import type { AuthStatus, AuthStorage } from "./auth-storage.ts"; -import { getModelRequestAuth, getApiKeyForProviderFromConfig, getProviderAuthStatusFromConfig, getProviderResolvedAuth } from "./model-registry-auth.ts"; -import { applyProviderConfigToModels, migrateLegacyRegisterProviderConfigValues, unregisterProviderRuntime, validateProviderConfig } from "./model-registry-dynamic.ts"; -import { loadModelRegistryModels } from "./model-registry-loader.ts"; -import { refreshExtensionProvider } from "./model-registry-extension-refresh.ts"; -import type { ProviderConfigInput, ProviderRequestConfig, ResolvedRequestAuth } from "./model-registry-types.ts"; -import type { ModelOverride } from "./model-registry-schemas.ts"; -import { BUILT_IN_PROVIDER_DISPLAY_NAMES } from "./provider-display-names.ts"; -import { type CodingAgentModelsStore, FileModelsStore, InMemoryCodingAgentModelsStore } from "./models-store.ts"; -import { getLegacyOAuthProvider, oauthCredentialToAuth } from "./oauth-provider-bridge.ts"; -import { withRemoteCatalog } from "./remote-catalog-provider.ts"; -import { clearConfigValueCache, isConfigValueConfigured } from "./resolve-config-value.ts"; -const REMOTE_CATALOG_PROVIDERS = new Set(["github-copilot", "openrouter", "vercel-ai-gateway"]); -const OPENAI_COMPATIBLE_APIS = new Set(["openai-completions", "openai-responses"]); -let nextRegistryRegistrationId = 0; - -function overlayRuntimeApiKey( - runtimeApiKey: string, - resolvedAuth: ModelAuth | undefined, - storedOAuthAuth: ModelAuth | undefined, - runtimeBaseUrl: string | undefined, -): ModelAuth { - const headers = storedOAuthAuth?.headers || resolvedAuth?.headers - ? { ...storedOAuthAuth?.headers, ...resolvedAuth?.headers } - : undefined; - return { - ...storedOAuthAuth, - ...resolvedAuth, - apiKey: runtimeApiKey, - headers, - baseUrl: runtimeBaseUrl ?? storedOAuthAuth?.baseUrl ?? resolvedAuth?.baseUrl, - }; -} - -function createProviderCredentialStore(credentials: CredentialStore): CredentialStore { - return { - ...credentials, - read: async (providerId) => { - const credential = await credentials.read(providerId); - if (credential?.type !== "oauth" || !getLegacyOAuthProvider(providerId)) return credential; - const auth = await oauthCredentialToAuth(providerId, credential); - return auth?.apiKey === undefined ? undefined : { type: "api_key", key: auth.apiKey }; - }, - }; -} +import type { Api, AuthResult, Model, Provider } from "@earendil-works/pi-ai"; +import type { ModelRuntime } from "./model-runtime.ts"; +import type { AuthStatus, ProviderConfigInput } from "./provider-composer.ts"; + +export type { ProviderConfigInput } from "./provider-composer.ts"; +export type ResolvedRequestAuth = + | { + ok: true; + apiKey?: string; + headers?: Record; + env?: Record; + } + | { ok: false; error: string }; +export { clearApiKeyCache } from "./provider-composer.ts"; -export type { ProviderConfigInput, ResolvedRequestAuth } from "./model-registry-types.ts"; -/** Clear the config value command cache. Exported for testing. */ -export const clearApiKeyCache = clearConfigValueCache; /** - * Model registry - loads and manages models, resolves API keys via AuthStorage. + * Synchronous compatibility facade exposed to extensions. + * Coding-agent internals use ModelRuntime directly. */ export class ModelRegistry { - private models: Model[] = []; - private modelOverrides: Map> = new Map(); - private providerRequestConfigs: Map = new Map(); - private modelRequestHeaders: Map> = new Map(); - private registeredProviders: Map = new Map(); - private nativeProviders: Map = new Map(); - private builtInProviders: Set = new Set(); - private customOpenAICompatibleProviders: Set = new Set(); - private loadError: string | undefined = undefined; - private refreshGeneration = 0; - private resolvedCustomAuthProviders = new Set(); - private readonly registrationSource = `atomic:model-registry:${++nextRegistryRegistrationId}`; - declare private readonly modelsStore: CodingAgentModelsStore; - declare private readonly credentialStore: CredentialStore; - declare private readonly providerModels: MutableModels; - private readonly defaultProviders = new Map(); - private configuredProviderIds = new Set(); - - declare readonly authStorage: AuthStorage; - declare private readonly modelsJsonPaths: string[]; - declare private readonly modelDataDir: string | undefined; - - private constructor( - authStorage: AuthStorage, - modelsJsonPaths: string[], - modelDataDir?: string, - ) { - this.authStorage = authStorage; - this.modelsJsonPaths = modelsJsonPaths.map((path) => normalizePath(path)); - this.modelDataDir = modelDataDir ? normalizePath(modelDataDir) : this.modelsJsonPaths[0] ? dirname(this.modelsJsonPaths[0]) : undefined; - this.modelsStore = this.modelDataDir ? new FileModelsStore(join(this.modelDataDir, "models-store.json")) : new InMemoryCodingAgentModelsStore(); - this.credentialStore = authStorage.asCredentialStore(); - this.providerModels = createModels({ - credentials: createProviderCredentialStore(this.credentialStore), - modelsStore: this.modelsStore, - }); - const builtinModelDataGeneratedAt = getBuiltinModelDataGeneratedAt(); - for (const provider of builtinProviders()) { - const configured = provider.id === "radius" - ? provider - : withRemoteCatalog(provider, undefined, builtinModelDataGeneratedAt); - this.defaultProviders.set(provider.id, configured); - this.providerModels.setProvider(configured); - } - this.loadModels(); - } - - static create( - authStorage: AuthStorage, - modelsJsonPath: string | string[] = getAgentConfigPaths("models.json"), - modelDataDir?: string, - ): ModelRegistry { - return new ModelRegistry(authStorage, Array.isArray(modelsJsonPath) ? modelsJsonPath : [modelsJsonPath], modelDataDir); - } + private readonly runtime: ModelRuntime; - static inMemory(authStorage: AuthStorage): ModelRegistry { return new ModelRegistry(authStorage, []); } - - async refresh(options: { signal?: AbortSignal; timeoutMs?: number; force?: boolean; allowNetwork?: boolean } = {}): Promise { - const generation = ++this.refreshGeneration; - this.loadError = undefined; - this.rebuildProviderModels(); - - const controller = new AbortController(); - const abort = () => controller.abort(); - options.signal?.addEventListener("abort", abort, { once: true }); - if (options.signal?.aborted) controller.abort(options.signal.reason); - const timeout = setTimeout(abort, options.timeoutMs ?? 15_000); - const errors = new Map(); - const aborted = new Promise((resolve) => { - if (controller.signal.aborted) resolve(); - else controller.signal.addEventListener("abort", () => resolve(), { once: true }); - }); - try { - const restore = this.providerModels - .refresh({ allowNetwork: false, signal: controller.signal }) - .then((result) => { - if (controller.signal.aborted || generation !== this.refreshGeneration) return; - for (const [provider, error] of result.errors) errors.set(provider, error); - this.publishProviderModels(); - }); - await Promise.race([restore, aborted]); - - if (options.allowNetwork !== false && !controller.signal.aborted) { - const legacyOAuthProviders = new Set([ - ...this.models.map((model) => model.provider), - ...[...this.registeredProviders] - .filter(([, config]) => config.oauth !== undefined) - .map(([providerId]) => providerId), - ].filter((providerId) => getLegacyOAuthProvider(providerId) !== undefined)); - const refreshLegacyOAuth = Promise.all([...legacyOAuthProviders].map(async (providerId) => { - const credential = this.authStorage.get(providerId); - if (credential?.type !== "oauth" || Date.now() < credential.expires) return; - await this.authStorage.getModelAuth(providerId, { includeFallback: false }); - })); - await Promise.race([refreshLegacyOAuth, aborted]); - } - - const extensionRefreshes = controller.signal.aborted ? [] : [...this.registeredProviders].map(async ([providerName, config]) => { - if (!config.refreshModels) return; - const isCurrentExtensionRefresh = () => !controller.signal.aborted - && generation === this.refreshGeneration - && this.registeredProviders.get(providerName) === config; - const store = { - read: () => this.modelsStore.read(providerName), - write: (entry: Parameters[1]) => - this.modelsStore.writeIf(providerName, entry, isCurrentExtensionRefresh), - delete: () => this.modelsStore.deleteIf(providerName, isCurrentExtensionRefresh), - }; - try { - let credential = await this.credentialStore.read(providerName); - if (credential?.type === "api_key") { - const effectiveApiKey = await this.authStorage.getApiKey(providerName, { includeFallback: false }); - credential = { ...credential, key: effectiveApiKey }; - } - if (credential === undefined) { - const customAuth = await getProviderResolvedAuth(providerName, this.authStorage, this.providerRequestConfigs); - if (customAuth) { this.resolvedCustomAuthProviders.add(providerName); credential = { type: "api_key", key: customAuth.auth.apiKey, env: customAuth.env }; } - else this.resolvedCustomAuthProviders.delete(providerName); - const configuredApiKey = customAuth ? undefined : await getApiKeyForProviderFromConfig(providerName, this.authStorage, this.providerRequestConfigs); - if (configuredApiKey !== undefined) credential = { type: "api_key", key: configuredApiKey }; - } - if (!isCurrentExtensionRefresh()) return; - const refreshed = await refreshExtensionProvider({ - config, - credential, - store, - allowNetwork: options.allowNetwork ?? true, - force: options.force, - signal: controller.signal, - isCurrent: isCurrentExtensionRefresh, - }); - if (!isCurrentExtensionRefresh()) return; - if (refreshed.error) errors.set(providerName, refreshed.error); - if (!refreshed.models) return; - this.registeredProviders.set(providerName, { ...config, models: refreshed.models }); - this.rebuildProviderModels(); - } catch (error) { if (isCurrentExtensionRefresh()) errors.set(providerName, error instanceof Error ? error : new Error(String(error))); } - }); - const builtinRefresh = controller.signal.aborted - ? Promise.resolve() - : this.providerModels - .refresh({ allowNetwork: options.allowNetwork ?? true, force: options.force, signal: controller.signal }) - .then((result) => { - if (controller.signal.aborted || generation !== this.refreshGeneration) return; - for (const [provider, error] of result.errors) errors.set(provider, error); - this.publishProviderModels(); - }); - await Promise.race([Promise.all(extensionRefreshes), aborted]); - await Promise.race([builtinRefresh, aborted]); - } finally { - clearTimeout(timeout); - options.signal?.removeEventListener("abort", abort); - } - return { aborted: controller.signal.aborted, errors }; + constructor(runtime: ModelRuntime) { + this.runtime = runtime; } - private publishProviderModels(): void { - this.rebuildProviderModels(); + /** Reload models.json asynchronously. Await before making synchronous registry reads. */ + async refresh(): Promise { + await this.runtime.refresh(); } - private rebuildProviderModels(): void { - this.loadModels(this.providerModels.getModels()); - for (const [providerName, config] of this.registeredProviders) { - this.applyProviderConfig(providerName, config); - } - } - - private providerRegistrationSource(providerName: string): string { return `${this.registrationSource}:${providerName}`; } - getError(): string | undefined { - return this.loadError; + return this.runtime.getError(); } - private loadModels(baseModels?: readonly Model[]): void { - const loaded = loadModelRegistryModels(this.authStorage, this.modelsJsonPaths, baseModels); - for (const providerId of this.configuredProviderIds) { - const fallback = this.defaultProviders.get(providerId); - if (fallback) this.providerModels.setProvider(fallback); - else this.providerModels.deleteProvider(providerId); - } - for (const provider of loaded.configuredProviders.values()) this.providerModels.setProvider(provider); - this.configuredProviderIds = new Set(loaded.configuredProviders.keys()); - this.modelOverrides = loaded.modelOverrides; - const previousModels = new Map(this.models.map((model) => [`${model.provider}\0${model.id}`, model])); - this.models = loaded.models.map((model) => { - const previous = previousModels.get(`${model.provider}\0${model.id}`); - return previous && isDeepStrictEqual(previous, model) ? previous : model; - }); - this.providerRequestConfigs = loaded.providerRequestConfigs; - this.modelRequestHeaders = loaded.modelRequestHeaders; - this.builtInProviders = loaded.builtInProviders; - this.customOpenAICompatibleProviders = loaded.customOpenAICompatibleProviders; - this.loadError = loaded.loadError; + getAll(): Model[] { + return [...this.runtime.getModels()]; } - getAll(): Model[] { return this.models; } - getAvailable(): Model[] { - const configured = this.models.filter((model) => this.hasConfiguredAuth(model)); - const allowedByProvider = new Map>(); - for (const provider of this.providerModels.getProviders()) { - const extension = this.registeredProviders.get(provider.id); - if (!provider.filterModels || extension?.models || extension?.oauth) continue; - const providerModels = configured.filter((model) => model.provider === provider.id); - const runtimeApiKey = this.authStorage.getRuntimeApiKey(provider.id); - const credential: Credential | undefined = runtimeApiKey === undefined - ? this.authStorage.get(provider.id) as Credential | undefined - : { type: "api_key", key: runtimeApiKey }; - allowedByProvider.set( - provider.id, - new Set(provider.filterModels(providerModels, credential).map((model) => model.id)), - ); - } - return configured.filter((model) => allowedByProvider.get(model.provider)?.has(model.id) ?? true); + return [...this.runtime.getAvailableSnapshot()]; } - find(provider: string, modelId: string): Model | undefined { return this.models.find((model) => model.provider === provider && model.id === modelId); } - getProviders(): readonly Provider[] { return this.providerModels.getProviders(); } - getProvider(providerId: string): Provider | undefined { return this.providerModels.getProvider(providerId); } - /** Whether an exact provider id belongs to a built-in, configured, or extension registration. */ - hasProvider(providerId: string): boolean { - return this.registeredProviders.has(providerId) - || this.providerModels.getProvider(providerId) !== undefined - || this.models.some((model) => model.provider === providerId); - } - checkAuth(providerId: string) { return this.providerModels.checkAuth(providerId); } - getAuth(providerId: string, overrides?: { apiKey?: string; env?: Record }): Promise; - getAuth(model: Model, overrides?: { apiKey?: string; env?: Record }): Promise; - getAuth(providerOrModel: string | Model, overrides?: { apiKey?: string; env?: Record }): Promise { - return typeof providerOrModel === "string" ? this.providerModels.getAuth(providerOrModel, overrides) : this.providerModels.getAuth(providerOrModel, overrides); - } - login(providerId: string, type: AuthType, interaction: AuthInteraction): Promise { return this.providerModels.login(providerId, type, interaction); } - async logoutProvider(providerId: string): Promise { await this.providerModels.logout(providerId); this.authStorage.reload(); } - - /** Whether an authenticated provider may reconstruct an absent saved model ID. */ - canRestoreUnknownModel(provider: string): boolean { - if (REMOTE_CATALOG_PROVIDERS.has(provider)) return true; - if (this.customOpenAICompatibleProviders.has(provider)) return true; - if (this.builtInProviders.has(provider)) return false; - - const config = this.registeredProviders.get(provider); - return ( - config?.models?.some((model) => { - const api = model.api ?? config.api; - return api !== undefined && OPENAI_COMPATIBLE_APIS.has(api); - }) === true - ); + find(provider: string, modelId: string): Model | undefined { + return this.runtime.getModel(provider, modelId); } hasConfiguredAuth(model: Model): boolean { - const providerApiKey = this.providerRequestConfigs.get(model.provider)?.apiKey; - return ( - this.authStorage.hasAuth(model.provider) || - (providerApiKey !== undefined && isConfigValueConfigured(providerApiKey)) - ); - } - - private getModelRequestKey(provider: string, modelId: string): string { - return `${provider}:${modelId}`; - } - - private storeProviderRequestConfig(providerName: string, config: ProviderRequestConfig): void { - if (config.apiKey || config.headers || config.authHeader || config.auth?.apiKey) this.providerRequestConfigs.set(providerName, config); - } - - private storeModelHeaders(providerName: string, modelId: string, headers?: Record): void { - const key = this.getModelRequestKey(providerName, modelId); - if (!headers || Object.keys(headers).length === 0) { - this.modelRequestHeaders.delete(key); - return; - } - this.modelRequestHeaders.set(key, headers); + return this.runtime.hasConfiguredAuth(model.provider); } - /** - * Get API key and request headers for a model. - */ async getApiKeyAndHeaders(model: Model): Promise { try { - const runtimeApiKey = this.authStorage.getRuntimeApiKey(model.provider); - const extensionReplacesOAuth = this.registeredProviders.get(model.provider)?.oauth !== undefined - || getLegacyOAuthProvider(model.provider) !== undefined; - const resolvedProviderAuth = this.providerModels.getProvider(model.provider) && !extensionReplacesOAuth - ? (await this.providerModels.getAuth( - model, - runtimeApiKey === undefined ? undefined : { apiKey: runtimeApiKey }, - ))?.auth - : undefined; - const storedCredential = this.authStorage.get(model.provider); - const storedOAuthAuth = runtimeApiKey !== undefined && !extensionReplacesOAuth && storedCredential?.type === "oauth" - ? await oauthCredentialToAuth(model.provider, storedCredential) + const resolution = await this.runtime.getAuth(model); + if (!resolution) { + const compatibility = this.runtime.getCompatibilityRequestConfig(model); + if (compatibility.authHeader) { + return { ok: false, error: `No API key found for "${model.provider}"` }; + } + const headers = compatibility.headers + ? Object.fromEntries( + Object.entries(compatibility.headers).filter( + (entry): entry is [string, string] => entry[1] !== null, + ), + ) + : undefined; + return { ok: true, headers }; + } + const headers = resolution.auth.headers + ? Object.fromEntries( + Object.entries(resolution.auth.headers).filter( + (entry): entry is [string, string] => entry[1] !== null, + ), + ) : undefined; - const providerAuth = runtimeApiKey === undefined - ? resolvedProviderAuth - : overlayRuntimeApiKey(runtimeApiKey, resolvedProviderAuth, storedOAuthAuth, undefined); - return getModelRequestAuth( - model, - this.authStorage, - this.providerRequestConfigs, - this.modelRequestHeaders, - providerAuth, - ); + return { ok: true, apiKey: resolution.auth.apiKey, headers, env: resolution.env }; } catch (error) { - return { ok: false, error: error instanceof Error ? error.message : String(error) }; + const cause = error instanceof Error ? error.cause : undefined; + const message = + cause instanceof Error ? cause.message : error instanceof Error ? error.message : String(error); + return { + ok: false, + error: + message === "authHeader requires a resolved API key" + ? `No API key found for "${model.provider}"` + : message, + }; } } - /** - * Return auth status for a provider, including request auth configured in models.json. - * This intentionally does not execute command-backed config values. - */ getProviderAuthStatus(provider: string): AuthStatus { - const status = getProviderAuthStatusFromConfig(provider, this.authStorage, this.providerRequestConfigs); - return status.source || !this.resolvedCustomAuthProviders.has(provider) ? status : { configured: true, source: "environment" }; + return this.runtime.getProviderAuthStatus(provider); + } + + getProvider(provider: string): Provider | undefined { + return this.runtime.getProvider(provider); } - /** Registered extension providers with a custom API-key login contract. */ - getCustomApiKeyAuthProviders(): Array<{ id: string; name: string }> { return [...this.registeredProviders].flatMap(([id, config]) => config.auth?.apiKey ? [{ id, name: config.auth.apiKey.name || this.getProviderDisplayName(id) }] : []); } - getCustomApiKeyAuth(provider: string) { return this.registeredProviders.get(provider)?.auth?.apiKey; } - async getProviderAuth(provider: string) { return getProviderResolvedAuth(provider, this.authStorage, this.providerRequestConfigs); } - /** - * Get display name for a provider. - */ + getProviderDisplayName(provider: string): string { - const registeredProvider = this.registeredProviders.get(provider); - const oauthProvider = this.authStorage.getOAuthProviders().find((p) => p.id === provider); + return this.runtime.getProvider(provider)?.name ?? provider; + } - return ( - registeredProvider?.name ?? - registeredProvider?.oauth?.name ?? - oauthProvider?.name ?? - BUILT_IN_PROVIDER_DISPLAY_NAMES[provider] ?? - provider - ); + getProviderAuth(provider: string): Promise { + return this.runtime.getAuth(provider); } async getApiKeyForProvider(provider: string): Promise { - return getApiKeyForProviderFromConfig(provider, this.authStorage, this.providerRequestConfigs); + try { + return (await this.runtime.getAuth(provider))?.auth.apiKey; + } catch { + return undefined; + } } isUsingOAuth(model: Model): boolean { - const cred = this.authStorage.get(model.provider); - return cred?.type === "oauth"; + return this.runtime.isUsingOAuth(model.provider); } registerProvider(provider: Provider): void; registerProvider(providerName: string, config: ProviderConfigInput): void; registerProvider(providerOrName: Provider | string, config?: ProviderConfigInput): void { - if (typeof providerOrName !== "string") { - if (!providerOrName.id.trim()) throw new Error("Provider id must not be empty."); - this.registeredProviders.delete(providerOrName.id); - this.nativeProviders.set(providerOrName.id, providerOrName); - this.providerModels.setProvider(providerOrName); - this.rebuildProviderModels(); + if (typeof providerOrName === "string") { + if (!config) throw new Error("Provider config is required when registering by name"); + this.runtime.registerProvider(providerOrName, config); return; } - if (!config) throw new Error("Provider config is required"); - this.nativeProviders.delete(providerOrName); - const migratedConfig = migrateLegacyRegisterProviderConfigValues(providerOrName, config); - validateProviderConfig(providerOrName, migratedConfig); - const mergedConfig = this.upsertRegisteredProvider(providerOrName, migratedConfig); - unregisterProviderRuntime(this.providerRegistrationSource(providerOrName)); - this.rebuildProviderModels(); - this.applyProviderConfig(providerOrName, mergedConfig, true); + this.runtime.registerNativeProvider(providerOrName); } - hasRegisteredStreamSimpleForApi(api: Api): boolean { - for (const config of this.registeredProviders.values()) { - if (config.api === api && config.streamSimple) { - return true; - } - } - return false; + unregisterProvider(providerName: string): void { + this.runtime.unregisterProvider(providerName); } - unregisterProvider(providerName: string): void { - const hadLegacy = this.registeredProviders.delete(providerName); - const hadNative = this.nativeProviders.delete(providerName); - if (!hadLegacy && !hadNative) return; - unregisterProviderRuntime(this.providerRegistrationSource(providerName)); - const fallback = this.defaultProviders.get(providerName); - if (fallback) this.providerModels.setProvider(fallback); - else this.providerModels.deleteProvider(providerName); - this.rebuildProviderModels(); + getRegisteredProviderConfig(providerName: string): ProviderConfigInput | undefined { + return this.runtime.getRegisteredProviderConfig(providerName); } - private upsertRegisteredProvider(providerName: string, config: ProviderConfigInput): ProviderConfigInput { - const definedConfig = Object.fromEntries( - Object.entries(config).filter(([, value]) => value !== undefined), - ) as ProviderConfigInput; - const merged = { ...this.registeredProviders.get(providerName), ...definedConfig }; - this.registeredProviders.set(providerName, merged); - return merged; + getRegisteredNativeProvider(providerName: string): Provider | undefined { + return this.runtime.getRegisteredNativeProvider(providerName); } - private applyProviderConfig(providerName: string, config: ProviderConfigInput, registerRuntime = false): void { - this.models = applyProviderConfigToModels({ - registrationSource: this.providerRegistrationSource(providerName), - registerRuntime, - providerName, - config, - models: this.models, - modelOverrides: this.modelOverrides, - authStorage: this.authStorage, - storeProviderRequestConfig: (name, requestConfig) => this.storeProviderRequestConfig(name, requestConfig), - storeModelHeaders: (name, modelId, headers) => this.storeModelHeaders(name, modelId, headers), - }); + getRegisteredProviderIds(): readonly string[] { + return this.runtime.getRegisteredProviderIds(); } } diff --git a/packages/coding-agent/src/core/model-resolver-cli.ts b/packages/coding-agent/src/core/model-resolver-cli.ts index 13324ebd9..c53e071b0 100644 --- a/packages/coding-agent/src/core/model-resolver-cli.ts +++ b/packages/coding-agent/src/core/model-resolver-cli.ts @@ -2,7 +2,7 @@ import type { Api, Model } from "@earendil-works/pi-ai/compat"; import { isValidThinkingLevel } from "../cli/args.ts"; import { buildFallbackModel, parseModelPattern } from "./model-resolver-patterns.ts"; import type { ResolveCliModelResult } from "./model-resolver-types.ts"; -import type { ModelRegistry } from "./model-registry.ts"; +import type { ModelRuntime } from "./model-runtime.ts"; function buildProviderMap(availableModels: Model[]): Map { const providerMap = new Map(); @@ -49,15 +49,15 @@ function splitCustomModelThinkingSuffix(pattern: string): { export function resolveCliModel(options: { cliProvider?: string; cliModel?: string; - modelRegistry: ModelRegistry; + modelRuntime: ModelRuntime; }): ResolveCliModelResult { - const { cliProvider, cliModel, modelRegistry } = options; + const { cliProvider, cliModel, modelRuntime } = options; if (!cliModel) { return { model: undefined, warning: undefined, error: undefined }; } - const availableModels = modelRegistry.getAll(); + const availableModels = [...modelRuntime.getModels()]; if (availableModels.length === 0) { return { model: undefined, diff --git a/packages/coding-agent/src/core/model-resolver-initial.ts b/packages/coding-agent/src/core/model-resolver-initial.ts index 4e0d552dc..eff4dad09 100644 --- a/packages/coding-agent/src/core/model-resolver-initial.ts +++ b/packages/coding-agent/src/core/model-resolver-initial.ts @@ -2,7 +2,7 @@ import type { ThinkingLevel } from "@earendil-works/pi-agent-core"; import type { Api, Model } from "@earendil-works/pi-ai/compat"; import chalk from "chalk"; import { DEFAULT_THINKING_LEVEL } from "./defaults.ts"; -import type { ModelRegistry } from "./model-registry.ts"; +import type { ModelRuntime } from "./model-runtime.ts"; import { findPreferredAvailableModel } from "./model-resolver-defaults.ts"; import { buildFallbackModel } from "./model-resolver-patterns.ts"; import { resolveCliModel } from "./model-resolver-cli.ts"; @@ -14,9 +14,9 @@ const CONFIGURED_DEFAULT_MODEL_UNAVAILABLE_MESSAGE = async function buildConfiguredProviderFallbackModel( provider: string, modelId: string, - modelRegistry: ModelRegistry, + modelRuntime: ModelRuntime, ): Promise | undefined> { - return buildFallbackModel(provider, modelId, await modelRegistry.getAvailable()); + return buildFallbackModel(provider, modelId, await [...modelRuntime.getAvailableSnapshot()]); } /** @@ -27,12 +27,12 @@ async function buildConfiguredProviderFallbackModel( export async function resolveRestoredModelReference( provider: string, modelId: string, - modelRegistry: ModelRegistry, + modelRuntime: ModelRuntime, ): Promise | undefined> { - const found = modelRegistry.find(provider, modelId); - if (found) return modelRegistry.hasConfiguredAuth(found) ? found : undefined; - if (!modelRegistry.canRestoreUnknownModel(provider)) return undefined; - return buildConfiguredProviderFallbackModel(provider, modelId, modelRegistry); + const found = modelRuntime.getModel(provider, modelId); + if (found) return modelRuntime.hasConfiguredAuth(found.provider) ? found : undefined; + if (!modelRuntime.canRestoreUnknownModel(provider)) return undefined; + return buildConfiguredProviderFallbackModel(provider, modelId, modelRuntime); } /** @@ -51,7 +51,7 @@ export async function findInitialModel(options: { defaultProvider?: string; defaultModelId?: string; defaultThinkingLevel?: ThinkingLevel; - modelRegistry: ModelRegistry; + modelRuntime: ModelRuntime; }): Promise { const { cliProvider, @@ -61,7 +61,7 @@ export async function findInitialModel(options: { defaultProvider, defaultModelId, defaultThinkingLevel, - modelRegistry, + modelRuntime, } = options; let model: Model | undefined; @@ -71,7 +71,7 @@ export async function findInitialModel(options: { const resolved = resolveCliModel({ cliProvider, cliModel, - modelRegistry, + modelRuntime, }); if (resolved.error) { console.error(chalk.red(resolved.error)); @@ -95,15 +95,15 @@ export async function findInitialModel(options: { } if (defaultProvider && defaultModelId) { - const found = modelRegistry.find(defaultProvider, defaultModelId); - if (found && modelRegistry.hasConfiguredAuth(found)) { + const found = modelRuntime.getModel(defaultProvider, defaultModelId); + if (found && modelRuntime.hasConfiguredAuth(found.provider)) { model = found; if (defaultThinkingLevel) { thinkingLevel = defaultThinkingLevel; } return { model, thinkingLevel, fallbackMessage: undefined }; } - if (!modelRegistry.hasProvider(defaultProvider)) { + if (!modelRuntime.getProvider(defaultProvider)) { return { model: undefined, thinkingLevel: DEFAULT_THINKING_LEVEL, @@ -113,7 +113,7 @@ export async function findInitialModel(options: { } } - const availableModels = await modelRegistry.getAvailable(); + const availableModels = await [...modelRuntime.getAvailableSnapshot()]; if (availableModels.length > 0) { return { model: findPreferredAvailableModel(availableModels), @@ -137,13 +137,13 @@ export async function restoreModelFromSession( savedModelId: string, currentModel: Model | undefined, shouldPrintMessages: boolean, - modelRegistry: ModelRegistry, + modelRuntime: ModelRuntime, ): Promise<{ model: Model | undefined; fallbackMessage: string | undefined; }> { - const exactRestoredModel = modelRegistry.find(savedProvider, savedModelId); - const restoredModel = await resolveRestoredModelReference(savedProvider, savedModelId, modelRegistry); + const exactRestoredModel = modelRuntime.getModel(savedProvider, savedModelId); + const restoredModel = await resolveRestoredModelReference(savedProvider, savedModelId, modelRuntime); if (restoredModel) { if (shouldPrintMessages) { @@ -168,7 +168,7 @@ export async function restoreModelFromSession( }; } - const availableModels = await modelRegistry.getAvailable(); + const availableModels = await [...modelRuntime.getAvailableSnapshot()]; const fallbackModel = findPreferredAvailableModel(availableModels); if (fallbackModel) { if (shouldPrintMessages) { diff --git a/packages/coding-agent/src/core/model-resolver-scope.ts b/packages/coding-agent/src/core/model-resolver-scope.ts index 9832bd5cb..fd6051d25 100644 --- a/packages/coding-agent/src/core/model-resolver-scope.ts +++ b/packages/coding-agent/src/core/model-resolver-scope.ts @@ -3,7 +3,7 @@ import { modelsAreEqual } from "@earendil-works/pi-ai/compat"; import chalk from "chalk"; import { minimatch } from "minimatch"; import { isValidThinkingLevel } from "../cli/args.ts"; -import type { ModelRegistry } from "./model-registry.ts"; +import type { ModelRuntime } from "./model-runtime.ts"; import { findExactModelReferenceMatch, parseModelPattern } from "./model-resolver-patterns.ts"; import type { ScopedModel } from "./model-resolver-types.ts"; @@ -50,9 +50,9 @@ function parseGlobThinkingLevel(pattern: string): { globPattern: string; thinkin */ export async function resolveModelScopeWithDiagnostics( patterns: string[], - modelRegistry: ModelRegistry, + modelRuntime: ModelRuntime, ): Promise { - const availableModels = await modelRegistry.getAvailable(); + const availableModels = await [...modelRuntime.getAvailableSnapshot()]; const scopedModels: ScopedModel[] = []; const diagnostics: ModelScopeDiagnostic[] = []; for (const pattern of patterns) { @@ -110,8 +110,8 @@ export async function resolveModelScopeWithDiagnostics( return { scopedModels, diagnostics }; } -export async function resolveModelScope(patterns: string[], modelRegistry: ModelRegistry): Promise { - const { scopedModels, diagnostics } = await resolveModelScopeWithDiagnostics(patterns, modelRegistry); +export async function resolveModelScope(patterns: string[], modelRuntime: ModelRuntime): Promise { + const { scopedModels, diagnostics } = await resolveModelScopeWithDiagnostics(patterns, modelRuntime); for (const diagnostic of diagnostics) { console.warn(chalk.yellow(`Warning: ${diagnostic.message}`)); } diff --git a/packages/coding-agent/src/core/model-runtime-auth.ts b/packages/coding-agent/src/core/model-runtime-auth.ts new file mode 100644 index 000000000..80431cd7c --- /dev/null +++ b/packages/coding-agent/src/core/model-runtime-auth.ts @@ -0,0 +1,29 @@ +import type { Api, AuthResult, Model } from "@earendil-works/pi-ai"; +import type { ModelConfig } from "./model-config.ts"; +import type { ModelRuntimeAuthOverrides } from "./model-runtime-types.ts"; +import type { ProviderConfigInput } from "./provider-composer.ts"; +import { resolveConfiguredModelHeaders } from "./provider-composer.ts"; +import { mergeHeaders } from "./model-runtime-streaming.ts"; + +/** Apply models.json and extension headers after pi-ai resolves provider credentials. */ +export function mergeConfiguredAuthHeaders( + resolution: AuthResult, + model: Model, + config: ModelConfig, + extension: ProviderConfigInput | undefined, + overrides: ModelRuntimeAuthOverrides, +): AuthResult { + const configuredHeaders = resolveConfiguredModelHeaders( + model, + config.getProvider(model.provider), + extension, + { ...(resolution.env ?? {}), ...(overrides.env ?? {}) }, + ); + return { + ...resolution, + auth: { + ...resolution.auth, + headers: mergeHeaders(resolution.auth.headers, configuredHeaders), + }, + }; +} diff --git a/packages/coding-agent/src/core/model-runtime-providers.ts b/packages/coding-agent/src/core/model-runtime-providers.ts new file mode 100644 index 000000000..09458ca80 --- /dev/null +++ b/packages/coding-agent/src/core/model-runtime-providers.ts @@ -0,0 +1,25 @@ +import type { Provider } from "@earendil-works/pi-ai"; +import * as builtinProviderCatalog from "@earendil-works/pi-ai/providers/all"; +import type { ModelConfig } from "./model-config.ts"; + +/** Rebuild the builtin layer, including configured Radius gateways. */ +export function configureBuiltinProviders( + target: Map, + defaults: ReadonlyMap, + config: ModelConfig, +): void { + target.clear(); + for (const [providerId, provider] of defaults) target.set(providerId, provider); + for (const providerId of config.getProviderIds()) { + const providerConfig = config.getProvider(providerId); + if (providerConfig?.oauth !== "radius" || !providerConfig.baseUrl) continue; + target.set( + providerId, + builtinProviderCatalog.radiusProvider({ + id: providerId, + name: providerConfig.name ?? providerId, + gateway: providerConfig.baseUrl.replace(/\/v1\/?$/u, ""), + }), + ); + } +} diff --git a/packages/coding-agent/src/core/model-runtime-restoration.ts b/packages/coding-agent/src/core/model-runtime-restoration.ts new file mode 100644 index 000000000..203c6ef38 --- /dev/null +++ b/packages/coding-agent/src/core/model-runtime-restoration.ts @@ -0,0 +1,35 @@ +import type { Api, Provider } from "@earendil-works/pi-ai"; +import type { ModelsJsonProvider } from "./model-config.ts"; +import type { ProviderConfigInput } from "./provider-composer.ts"; + +const REMOTE_CATALOG_PROVIDERS = new Set(["github-copilot", "openrouter", "vercel-ai-gateway"]); +const OPENAI_COMPATIBLE_APIS = new Set(["openai-completions", "openai-responses"]); + +type ConfiguredModel = { + api?: Api; +}; + +function hasOpenAICompatibleModel( + models: readonly ConfiguredModel[] | undefined, + defaultApi: Api | undefined, +): boolean { + return models?.some((model) => { + const api = model.api ?? defaultApi; + return api !== undefined && OPENAI_COMPATIBLE_APIS.has(api); + }) === true; +} + +/** Whether an authenticated provider may reconstruct an absent saved model ID. */ +export function canRestoreUnknownModel( + providerId: string, + builtin: Provider | undefined, + modelsConfig: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, + nativeExtension: Provider | undefined, +): boolean { + if (REMOTE_CATALOG_PROVIDERS.has(providerId)) return true; + if (builtin !== undefined) return false; + if (hasOpenAICompatibleModel(modelsConfig?.models, modelsConfig?.api)) return true; + if (hasOpenAICompatibleModel(extension?.models, extension?.api)) return true; + return hasOpenAICompatibleModel(nativeExtension?.getModels(), undefined); +} diff --git a/packages/coding-agent/src/core/model-runtime-snapshot.ts b/packages/coding-agent/src/core/model-runtime-snapshot.ts new file mode 100644 index 000000000..2bd3e0345 --- /dev/null +++ b/packages/coding-agent/src/core/model-runtime-snapshot.ts @@ -0,0 +1,82 @@ +import type { Api, AuthCheck, CredentialInfo, Model } from "@earendil-works/pi-ai"; +import type { AuthStatus } from "./provider-composer.ts"; + +export interface ModelRuntimeSnapshot { + all: readonly Model[]; + available: readonly Model[]; + configuredProviders: ReadonlySet; + storedProviders: ReadonlySet; + storedCredentialTypes: ReadonlyMap; + auth: ReadonlyMap; +} + +export function createEmptyModelRuntimeSnapshot(): ModelRuntimeSnapshot { + return { + all: [], + available: [], + configuredProviders: new Set(), + storedProviders: new Set(), + storedCredentialTypes: new Map(), + auth: new Map(), + }; +} + +export function createModelRuntimeSnapshot( + all: readonly Model[], + available: readonly Model[], + checks: readonly (readonly [string, AuthCheck | undefined])[], + credentials: readonly CredentialInfo[], +): ModelRuntimeSnapshot { + const auth = new Map(checks); + const configuredProviders = new Set( + checks + .filter((entry): entry is readonly [string, AuthCheck] => entry[1] !== undefined) + .map(([providerId]) => providerId), + ); + return { + all, + available, + configuredProviders, + storedProviders: new Set(credentials.map((entry) => entry.providerId)), + storedCredentialTypes: new Map(credentials.map((entry) => [entry.providerId, entry.type])), + auth, + }; +} + +export function updateSnapshotModels( + snapshot: ModelRuntimeSnapshot, + all: readonly Model[], +): ModelRuntimeSnapshot { + return { + ...snapshot, + all, + available: all.filter((model) => snapshot.configuredProviders.has(model.provider)), + }; +} + +export function addRuntimeApiKeyProvider( + snapshot: ModelRuntimeSnapshot, + providerId: string, +): ModelRuntimeSnapshot { + const configuredProviders = new Set(snapshot.configuredProviders).add(providerId); + return { + ...snapshot, + auth: new Map(snapshot.auth).set(providerId, { type: "api_key", source: "runtime API key" }), + configuredProviders, + storedProviders: new Set(snapshot.storedProviders).add(providerId), + available: snapshot.all.filter((model) => configuredProviders.has(model.provider)), + }; +} + +export function getSnapshotProviderAuthStatus( + snapshot: ModelRuntimeSnapshot, + providerId: string, + hasRuntimeApiKey: boolean, + configured: AuthStatus | undefined, +): AuthStatus { + if (hasRuntimeApiKey) return { configured: true, source: "runtime" }; + if (snapshot.storedProviders.has(providerId)) return { configured: true, source: "stored" }; + if (configured) return configured; + const check = snapshot.auth.get(providerId); + return check ? { configured: true, source: "environment", label: check.source } : { configured: false }; +} diff --git a/packages/coding-agent/src/core/model-runtime-streaming.ts b/packages/coding-agent/src/core/model-runtime-streaming.ts new file mode 100644 index 000000000..0ac26dd01 --- /dev/null +++ b/packages/coding-agent/src/core/model-runtime-streaming.ts @@ -0,0 +1,116 @@ +import { + type Api, + type ApiStreamOptions, + type AssistantMessage, + type AssistantMessageEventStream, + type AuthResult, + type Context, + lazyStream, + type Model, + type ModelsApiStreamOptions, + ModelsError, + type ModelsSimpleStreamOptions, + type ModelsStreamTransforms, + type MutableModels, + type Provider, + type ProviderHeaders, + type SimpleStreamOptions, + type StreamOptions, +} from "@earendil-works/pi-ai"; +import type { ModelRuntimeAuthOverrides } from "./model-runtime-types.ts"; + +export function mergeHeaders( + base: ProviderHeaders | undefined, + override: ProviderHeaders | undefined, +): ProviderHeaders | undefined { + if (!base && !override) return undefined; + const merged = { ...base }; + for (const [name, value] of Object.entries(override ?? {})) { + const lowerName = name.toLowerCase(); + for (const existingName of Object.keys(merged)) { + if (existingName.toLowerCase() === lowerName) delete merged[existingName]; + } + merged[name] = value; + } + return merged; +} + +type ResolveAuth = ( + model: Model, + overrides?: ModelRuntimeAuthOverrides, +) => Promise; + +/** Streaming request preparation split from ModelRuntime solely for Atomic's 500-line source gate. */ +export class ModelRuntimeStreaming { + private readonly models: MutableModels; + private readonly resolveAuth: ResolveAuth; + constructor(models: MutableModels, resolveAuth: ResolveAuth) { + this.models = models; + this.resolveAuth = resolveAuth; + } + + private async prepareRequest( + model: Model, + options: (StreamOptions & ModelsStreamTransforms) | undefined, + ): Promise<{ provider: Provider; model: Model; options: StreamOptions }> { + const provider = this.models.getProvider(model.provider); + if (!provider) throw new ModelsError("provider", `Unknown provider: ${model.provider}`); + const resolution = await this.resolveAuth(model, { apiKey: options?.apiKey, env: options?.env }); + if (!resolution) throw new ModelsError("auth", `Provider is not configured: ${model.provider}`); + + const { transformHeaders, ...providerOptions } = options ?? {}; + let headers = mergeHeaders(resolution.auth.headers, providerOptions.headers); + if (transformHeaders) headers = await transformHeaders(headers ?? {}); + const env = + resolution.env || providerOptions.env + ? { ...(resolution.env ?? {}), ...(providerOptions.env ?? {}) } + : undefined; + return { + provider, + model: resolution.auth.baseUrl ? { ...model, baseUrl: resolution.auth.baseUrl } : model, + options: { + ...providerOptions, + apiKey: providerOptions.apiKey ?? resolution.auth.apiKey, + headers, + env, + }, + }; + } + + stream( + model: Model, + context: Context, + options?: ModelsApiStreamOptions, + ): AssistantMessageEventStream { + return lazyStream(model, async () => { + const prepared = await this.prepareRequest( + model, + options as (StreamOptions & ModelsStreamTransforms) | undefined, + ); + return prepared.provider.stream( + prepared.model as Model, + context, + prepared.options as ApiStreamOptions, + ); + }); + } + + complete( + model: Model, + context: Context, + options?: ModelsApiStreamOptions, + ): Promise { + return this.stream(model, context, options).result(); + } + + streamSimple(model: Model, context: Context, options?: ModelsSimpleStreamOptions): AssistantMessageEventStream { + return lazyStream(model, async () => { + const prepared = await this.prepareRequest(model, options); + return prepared.provider.streamSimple(prepared.model, context, prepared.options as SimpleStreamOptions); + }); + } + + completeSimple(model: Model, context: Context, options?: ModelsSimpleStreamOptions): Promise { + return this.streamSimple(model, context, options).result(); + } +} diff --git a/packages/coding-agent/src/core/model-runtime-types.ts b/packages/coding-agent/src/core/model-runtime-types.ts new file mode 100644 index 000000000..2f8e3554e --- /dev/null +++ b/packages/coding-agent/src/core/model-runtime-types.ts @@ -0,0 +1,22 @@ +import type { CredentialStore, ModelsStore } from "@earendil-works/pi-ai"; + +export interface CreateModelRuntimeOptions { + /** Credential storage. Defaults to the file at authPath. */ + credentials?: CredentialStore; + authPath?: string; + modelsPath?: string | null; + modelsStore?: ModelsStore; + modelsStorePath?: string; + /** Allow create() to refresh model catalogs over the network. Defaults to false. */ + allowModelNetwork?: boolean; + /** Timeout for the create-time network model refresh. */ + modelRefreshTimeoutMs?: number; + catalogBaseUrl?: string; +} + +export interface ModelRuntimeAuthOverrides { + apiKey?: string; + env?: Record; + /** Require this much remaining OAuth-token validity; defaults to five minutes. */ + minOAuthValidityMs?: number; +} diff --git a/packages/coding-agent/src/core/model-runtime.ts b/packages/coding-agent/src/core/model-runtime.ts index e214d602e..8ebb0e018 100644 --- a/packages/coding-agent/src/core/model-runtime.ts +++ b/packages/coding-agent/src/core/model-runtime.ts @@ -1,164 +1,480 @@ -import type { - ApiStreamOptions, - AssistantMessage, - AssistantMessageEventStream, - AuthInteraction, - AuthResult, - AuthType, - Context, - Credential, - CredentialInfo, - Models, - ModelsApiStreamOptions, - ModelsRefreshOptions, - ModelsRefreshResult, - ModelsSimpleStreamOptions, - Provider, - ProviderHeaders, - StreamOptions, +import { dirname, join } from "node:path"; +import { + type Api, + type AssistantMessage, + type AssistantMessageEventStream, + type AuthCheck, + type AuthInteraction, + type AuthResult, + type AuthType, + type Context, + type Credential, + type CredentialInfo, + createModels, + type Model, + type Models, + type ModelsApiStreamOptions, + type ModelsRefreshOptions, + type ModelsRefreshResult, + type ModelsSimpleStreamOptions, + type ModelsStore, + type MutableModels, + type Provider, } from "@earendil-works/pi-ai"; -import { lazyStream } from "@earendil-works/pi-ai"; -import type { Api, Model } from "@earendil-works/pi-ai/compat"; -import { getAgentConfigPaths } from "../config.ts"; -import { AuthStorage } from "./auth-storage.ts"; -import type { AuthStatus } from "./auth-storage.ts"; -import { ModelRegistry } from "./model-registry.ts"; - -export interface CreateModelRuntimeOptions { - authStorage?: AuthStorage; - authPath?: string | string[]; - modelRegistry?: ModelRegistry; - modelsPath?: string | string[] | null; - allowModelNetwork?: boolean; - modelRefreshTimeoutMs?: number; -} +import * as builtinProviderCatalog from "@earendil-works/pi-ai/providers/all"; +import { getAgentDir } from "../config.ts"; +import { AuthStorage as DefaultAuthStorage } from "./auth-storage.ts"; +import { ModelConfig } from "./model-config.ts"; +import { FileModelsStore, InMemoryCodingAgentModelsStore } from "./models-store.ts"; +import { + type AuthStatus, + composeModelProvider, + configuredRequestAuthStatus, + type CompatibilityRequestConfig, + type ProviderConfigInput, + resolveCompatibilityRequestConfig, + validateExtensionProvider, +} from "./provider-composer.ts"; +import { withRemoteCatalog } from "./remote-catalog-provider.ts"; +import { RuntimeCredentials } from "./runtime-credentials.ts"; +import { isOfflineModeEnabled } from "./package-manager-env.ts"; +import { collectOAuthProviderMetadata } from "./oauth-provider-metadata.ts"; +import { OAuthLoginTransactionError } from "./oauth-login.ts"; +import { + addRuntimeApiKeyProvider, + createEmptyModelRuntimeSnapshot, + createModelRuntimeSnapshot, + getSnapshotProviderAuthStatus, + type ModelRuntimeSnapshot, + updateSnapshotModels, +} from "./model-runtime-snapshot.ts"; +export type { CreateModelRuntimeOptions, ModelRuntimeAuthOverrides } from "./model-runtime-types.ts"; +import type { CreateModelRuntimeOptions, ModelRuntimeAuthOverrides } from "./model-runtime-types.ts"; +import { ModelRuntimeStreaming } from "./model-runtime-streaming.ts"; +import { canRestoreUnknownModel as canRestoreUnknownModelProvider } from "./model-runtime-restoration.ts"; +import { mergeConfiguredAuthHeaders } from "./model-runtime-auth.ts"; +import { configureBuiltinProviders } from "./model-runtime-providers.ts"; +/** Configured pi-ai Models collection used by coding-agent and SDK consumers. */ +export class ModelRuntime implements Models { + private readonly models: MutableModels; + private readonly credentials: RuntimeCredentials; + private readonly streaming: ModelRuntimeStreaming; + private readonly defaultBuiltins: ReadonlyMap; + private readonly builtins = new Map(); + private readonly nativeExtensionProviders = new Map(); + private readonly extensionProviders = new Map(); + private readonly compositionErrors = new Map(); + private readonly modelsPath: string | undefined; + private readonly modelNetworkEnabled: boolean; + private config: ModelConfig; + private snapshot: ModelRuntimeSnapshot = createEmptyModelRuntimeSnapshot(); + private availabilityRefresh: Promise | undefined; + private availabilityError: string | undefined; + private constructor( + credentials: RuntimeCredentials, + config: ModelConfig, + modelsPath: string | undefined, + modelsStore: ModelsStore, + providers: readonly Provider[], + modelNetworkEnabled: boolean, + ) { + this.credentials = credentials; + this.config = config; + this.modelsPath = modelsPath; + this.modelNetworkEnabled = modelNetworkEnabled; + this.defaultBuiltins = new Map(providers.map((provider) => [provider.id, provider])); + for (const [providerId, provider] of this.defaultBuiltins) this.builtins.set(providerId, provider); + this.models = createModels({ credentials, modelsStore }); + this.streaming = new ModelRuntimeStreaming(this.models, (model, overrides) => this.getAuth(model, overrides)); + this.rebuildProviders(); + } + static async create(options: CreateModelRuntimeOptions = {}): Promise { + const credentials = new RuntimeCredentials(options.credentials ?? DefaultAuthStorage.create(options.authPath)); + const modelsPath = + options.modelsPath === null ? undefined : (options.modelsPath ?? join(getAgentDir(), "models.json")); + const config = await ModelConfig.load(modelsPath); + const modelsStore = + options.modelsStore ?? + (modelsPath + ? new FileModelsStore(options.modelsStorePath ?? join(dirname(modelsPath), "models-store.json")) + : new InMemoryCodingAgentModelsStore()); + const builtinModelDataGeneratedAt = builtinProviderCatalog.getBuiltinModelDataGeneratedAt(); + const providers = builtinProviderCatalog + .builtinProviders() + .map((provider) => + provider.id === "radius" + ? provider + : withRemoteCatalog(provider, options.catalogBaseUrl, builtinModelDataGeneratedAt), + ); + const runtime = new ModelRuntime( + credentials, + config, + modelsPath, + modelsStore, + providers, + !isOfflineModeEnabled(), + ); + runtime.configureRadiusProviders(); + runtime.rebuildProviders(); + const refreshFromNetwork = runtime.modelNetworkEnabled && options.allowModelNetwork === true; + const controller = refreshFromNetwork ? new AbortController() : undefined; + const timeout = controller + ? setTimeout(() => controller.abort(), options.modelRefreshTimeoutMs ?? 15_000) + : undefined; + try { + await runtime.refresh({ allowNetwork: refreshFromNetwork, signal: controller?.signal }); + } finally { + if (timeout) clearTimeout(timeout); + } + return runtime; + } + private configureRadiusProviders(): void { + configureBuiltinProviders(this.builtins, this.defaultBuiltins, this.config); + } + private providerIds(): Set { + return new Set([ + ...this.builtins.keys(), + ...this.nativeExtensionProviders.keys(), + ...this.config.getProviderIds(), + ...this.extensionProviders.keys(), + ]); + } + private recomposeProvider(providerId: string): void { + const base = this.nativeExtensionProviders.get(providerId) ?? this.builtins.get(providerId); + const extension = this.extensionProviders.get(providerId); + if (!base && !this.config.getProvider(providerId) && !extension) { + this.models.deleteProvider(providerId); + this.compositionErrors.delete(providerId); + return; + } + if (base && !this.config.getProvider(providerId) && !extension) { + // No overlays: use the builtin untouched so its auth/login/stream behavior is exact. + this.models.setProvider(base); + this.compositionErrors.delete(providerId); + return; + } + try { + this.models.setProvider(composeModelProvider(providerId, base, this.config, extension)); + this.compositionErrors.delete(providerId); + } catch (error) { + this.compositionErrors.set(providerId, error instanceof Error ? error.message : String(error)); + if (base) this.models.setProvider(base); + else this.models.deleteProvider(providerId); + } + } + private rebuildProviders(): void { + this.models.clearProviders(); + this.compositionErrors.clear(); + for (const providerId of this.providerIds()) this.recomposeProvider(providerId); + this.updateModelSnapshot(); + } + private updateModelSnapshot(): void { + this.snapshot = updateSnapshotModels(this.snapshot, [...this.models.getModels()]); + } + private async runAvailabilityRefresh(): Promise { + const providers = this.models.getProviders(); + const [available, checks, credentials] = await Promise.all([ + this.models.getAvailable(), + Promise.all( + providers.map( + async (provider): Promise<[string, AuthCheck | undefined]> => [ + provider.id, + await this.models.checkAuth(provider.id), + ], + ), + ), + this.credentials.list(), + ]); + this.snapshot = createModelRuntimeSnapshot( + [...this.models.getModels()], + [...available], + checks, + credentials, + ); + this.availabilityError = undefined; + } + private queueAvailabilityRefresh(after: Promise | undefined): Promise { + const refresh = (after ?? Promise.resolve()).catch(() => {}).then(() => this.runAvailabilityRefresh()); + const recorded = refresh.catch((error) => { + this.availabilityError = error instanceof Error ? error.message : String(error); + throw error; + }); + const tracked = recorded.finally(() => { + if (this.availabilityRefresh === tracked) this.availabilityRefresh = undefined; + }); + this.availabilityRefresh = tracked; + return tracked; + } + /** Coalesce concurrent readers onto the pending refresh. */ + private refreshAvailability(): Promise { + return this.availabilityRefresh ?? this.queueAvailabilityRefresh(undefined); + } + /** Mutations must not observe an in-flight refresh started before them. */ + private forceRefreshAvailability(): Promise { + return this.queueAvailabilityRefresh(this.availabilityRefresh); + } + getProviders(): readonly Provider[] { + return this.models.getProviders(); + } + getProvider(providerId: string): Provider | undefined { + return this.models.getProvider(providerId); + } + /** Whether an authenticated provider may reconstruct an absent saved model ID. */ + canRestoreUnknownModel(providerId: string): boolean { + return canRestoreUnknownModelProvider( + providerId, + this.defaultBuiltins.get(providerId), + this.config.getProvider(providerId), + this.extensionProviders.get(providerId), + this.nativeExtensionProviders.get(providerId), + ); + } + getModels(providerId?: string): readonly Model[] { + return this.models.getModels(providerId); + } + getModel(providerId: string, modelId: string): Model | undefined { + return this.models.getModel(providerId, modelId); + } + async checkAuth(providerId: string): Promise { + return this.models.checkAuth(providerId); + } + async getAvailable(providerId?: string): Promise[]> { + if (providerId) { + if (this.availabilityRefresh) { + await this.availabilityRefresh; + return this.snapshot.available.filter((model) => model.provider === providerId); + } + try { + return await this.models.getAvailable(providerId); + } catch (error) { + this.availabilityError = error instanceof Error ? error.message : String(error); + throw error; + } + } + await this.refreshAvailability(); + return this.snapshot.available; + } + getAvailableSnapshot(): readonly Model[] { + return this.snapshot.available; + } + getError(): string | undefined { + const errors: string[] = []; + const configError = this.config.getError(); + if (configError) errors.push(configError); + for (const [providerId, error] of this.compositionErrors) { + errors.push(`Provider "${providerId}": ${error}`); + } + if (this.availabilityError) errors.push(`Availability refresh: ${this.availabilityError}`); + return errors.length > 0 ? errors.join("\n\n") : undefined; + } -export interface ModelRuntimeAuthOverrides { - apiKey?: string; - env?: Record; -} + getRegisteredProviderConfig(providerId: string): ProviderConfigInput | undefined { + return this.extensionProviders.get(providerId); + } -function mergeHeaders(base?: ProviderHeaders, override?: ProviderHeaders): ProviderHeaders | undefined { - if (!base && !override) return undefined; - const merged = { ...base }; - for (const [name, value] of Object.entries(override ?? {})) { - for (const existing of Object.keys(merged)) { - if (existing.toLowerCase() === name.toLowerCase()) delete merged[existing]; - } - merged[name] = value; - } - return merged; -} + getRegisteredProviderIds(): readonly string[] { + return [...new Set([...this.extensionProviders.keys(), ...this.nativeExtensionProviders.keys()])]; + } -/** Canonical model/auth facade for SDK consumers. */ -export class ModelRuntime implements Models { - readonly modelRegistry: ModelRegistry; - readonly authStorage: AuthStorage; - - constructor(modelRegistry: ModelRegistry, authStorage: AuthStorage = modelRegistry.authStorage) { - this.modelRegistry = modelRegistry; - this.authStorage = authStorage; - } - - static async create(options: CreateModelRuntimeOptions = {}): Promise { - const authStorage = options.authStorage ?? AuthStorage.create(options.authPath); - const paths = options.modelsPath === null ? [] : (options.modelsPath ?? getAgentConfigPaths("models.json")); - const modelRegistry = options.modelRegistry ?? ( - Array.isArray(paths) && paths.length === 0 - ? ModelRegistry.inMemory(authStorage) - : ModelRegistry.create(authStorage, paths) - ); - const runtime = new ModelRuntime(modelRegistry, authStorage); - await modelRegistry.refresh({ - allowNetwork: options.allowModelNetwork ?? false, - timeoutMs: options.modelRefreshTimeoutMs, - }); - return runtime; - } - - getProviders(): readonly Provider[] { return this.modelRegistry.getProviders(); } - getProvider(providerId: string): Provider | undefined { return this.modelRegistry.getProvider(providerId); } - getModels(providerId?: string): readonly Model[] { - return providerId ? this.modelRegistry.getAll().filter((model) => model.provider === providerId) : this.modelRegistry.getAll(); - } - getModel(providerId: string, modelId: string): Model | undefined { return this.modelRegistry.find(providerId, modelId); } - checkAuth(providerId: string) { return this.modelRegistry.checkAuth(providerId); } - async getAvailable(providerId?: string): Promise[]> { - const models = this.modelRegistry.getAvailable(); - return providerId ? models.filter((model) => model.provider === providerId) : models; - } - getAvailableSnapshot(): readonly Model[] { return this.modelRegistry.getAvailable(); } - getError(): string | undefined { - const errors = [this.modelRegistry.getError(), this.authStorage.getLoadError()?.message].filter((value): value is string => Boolean(value)); - return errors.length ? errors.join("\n\n") : undefined; - } - getAuth(providerId: string, overrides?: ModelRuntimeAuthOverrides): Promise; - getAuth(model: Model, overrides?: ModelRuntimeAuthOverrides): Promise; - getAuth(providerOrModel: string | Model, overrides?: ModelRuntimeAuthOverrides): Promise { - return typeof providerOrModel === "string" - ? this.modelRegistry.getAuth(providerOrModel, overrides) - : this.modelRegistry.getAuth(providerOrModel, overrides); - } - getProviderAuthStatus(providerId: string): AuthStatus { return this.modelRegistry.getProviderAuthStatus(providerId); } - isUsingOAuth(providerId: string): boolean { return this.authStorage.get(providerId)?.type === "oauth"; } - hasConfiguredAuth(providerId: string): boolean { return this.modelRegistry.getProviderAuthStatus(providerId).configured; } - async setRuntimeApiKey(providerId: string, apiKey: string, options: ModelsRefreshOptions = {}): Promise { - this.authStorage.setRuntimeApiKey(providerId, apiKey); - await this.refresh(options); - } - async removeRuntimeApiKey(providerId: string): Promise { - this.authStorage.removeRuntimeApiKey(providerId); - await this.refresh({ allowNetwork: false }); - } - async listCredentials(): Promise { return this.authStorage.asCredentialStore().list(); } - - private async prepare( - model: Model, - options: (StreamOptions & { transformHeaders?: (headers: ProviderHeaders) => ProviderHeaders | Promise }) | undefined, - ) { - const provider = this.getProvider(model.provider); - if (!provider) throw new Error(`Unknown provider: ${model.provider}`); - const auth = await this.modelRegistry.getApiKeyAndHeaders(model); - if (!auth.ok) throw new Error(auth.error); - const { transformHeaders, ...providerOptions } = options ?? {}; - let headers = mergeHeaders(auth.headers, providerOptions.headers); - if (transformHeaders) headers = await transformHeaders(headers ?? {}); - return { - provider, - model: auth.baseUrl ? { ...model, baseUrl: auth.baseUrl } : model, - options: { ...providerOptions, apiKey: providerOptions.apiKey ?? auth.apiKey, headers }, - }; - } - stream(model: Model, context: Context, options?: ModelsApiStreamOptions): AssistantMessageEventStream { - return lazyStream(model, async () => { - const prepared = await this.prepare(model, options as StreamOptions | undefined); - return prepared.provider.stream(prepared.model as Model, context, prepared.options as ApiStreamOptions); - }); - } - complete(model: Model, context: Context, options?: ModelsApiStreamOptions): Promise { - return this.stream(model, context, options).result(); - } - streamSimple(model: Model, context: Context, options?: ModelsSimpleStreamOptions): AssistantMessageEventStream { - return lazyStream(model, async () => { - const prepared = await this.prepare(model, options); - return prepared.provider.streamSimple(prepared.model, context, prepared.options); - }); - } - completeSimple(model: Model, context: Context, options?: ModelsSimpleStreamOptions): Promise { - return this.streamSimple(model, context, options).result(); - } - async login(providerId: string, type: AuthType, interaction: AuthInteraction): Promise { - const credential = await this.modelRegistry.login(providerId, type, interaction); - this.authStorage.reload(); - await this.refresh({ allowNetwork: false }); - return credential; - } - async logout(providerId: string): Promise { - await this.modelRegistry.logoutProvider(providerId); - await this.refresh({ allowNetwork: false }); - } - async reloadConfig(): Promise { await this.refresh({ allowNetwork: false }); } - refresh(options: ModelsRefreshOptions = {}): Promise { return this.modelRegistry.refresh(options); } - registerProvider(provider: Provider): void { this.modelRegistry.registerProvider(provider); } - unregisterProvider(providerId: string): void { this.modelRegistry.unregisterProvider(providerId); } + getRegisteredNativeProvider(providerId: string): Provider | undefined { + return this.nativeExtensionProviders.get(providerId); + } + getOAuthProviderMetadata() { return collectOAuthProviderMetadata(this.getProviders(), this.extensionProviders); } + /** @internal Compatibility fallback for ModelRegistry when provider auth is unconfigured. */ + getCompatibilityRequestConfig(model: Model): CompatibilityRequestConfig { + return resolveCompatibilityRequestConfig( + model, + this.config.getProvider(model.provider), + this.extensionProviders.get(model.provider), + ); + } + + isUsingOAuth(providerId: string): boolean { + return this.snapshot.auth.get(providerId)?.type === "oauth"; + } + + hasConfiguredAuth(providerId: string): boolean { + return this.snapshot.configuredProviders.has(providerId); + } + + getAuth(providerId: string, overrides?: ModelRuntimeAuthOverrides): Promise; + getAuth(model: Model, overrides?: ModelRuntimeAuthOverrides): Promise; + async getAuth( + providerOrModel: string | Model, + overrides: ModelRuntimeAuthOverrides = {}, + ): Promise { + if (typeof providerOrModel === "string") return this.models.getAuth(providerOrModel, overrides); + const resolution = await this.models.getAuth(providerOrModel, overrides); + if (!resolution) return undefined; + return mergeConfiguredAuthHeaders( + resolution, + providerOrModel, + this.config, + this.extensionProviders.get(providerOrModel.provider), + overrides, + ); + } + /** Reload credentials changed by the authoritative isolated engine and update auth status snapshots. */ + async reloadCredentials(): Promise { + await this.credentials.reload(); + await this.forceRefreshAvailability(); + } + + async saveCredential(providerId: string, credential: Credential): Promise { + await this.credentials.modify(providerId, async () => credential); + await this.refresh(); + } + async setRuntimeApiKey( + providerId: string, + apiKey: string, + refreshOptions: ModelsRefreshOptions = {}, + ): Promise { + this.credentials.setRuntimeApiKey(providerId, apiKey); + this.snapshot = addRuntimeApiKeyProvider(this.snapshot, providerId); + await this.refresh(refreshOptions); + } + + async removeRuntimeApiKey(providerId: string): Promise { + this.credentials.removeRuntimeApiKey(providerId); + await this.refresh({ allowNetwork: this.modelNetworkEnabled }); + } + + listCredentials(): Promise { + return this.credentials.list(); + } + + getStoredCredentialType(providerId: string): CredentialInfo["type"] | undefined { + return this.snapshot.storedCredentialTypes.get(providerId); + } + + getProviderAuthStatus(providerId: string): AuthStatus { + return getSnapshotProviderAuthStatus( + this.snapshot, + providerId, + this.credentials.hasRuntimeApiKey(providerId), + configuredRequestAuthStatus( + this.config.getProvider(providerId), + this.extensionProviders.get(providerId), + ), + ); + } + + stream( + model: Model, + context: Context, + options?: ModelsApiStreamOptions, + ): AssistantMessageEventStream { + return this.streaming.stream(model, context, options); + } + + complete( + model: Model, + context: Context, + options?: ModelsApiStreamOptions, + ): Promise { + return this.streaming.complete(model, context, options); + } + + streamSimple(model: Model, context: Context, options?: ModelsSimpleStreamOptions): AssistantMessageEventStream { + return this.streaming.streamSimple(model, context, options); + } + + completeSimple(model: Model, context: Context, options?: ModelsSimpleStreamOptions): Promise { + return this.streaming.completeSimple(model, context, options); + } + + async login(providerId: string, type: AuthType, interaction: AuthInteraction): Promise { + const credential = await this.models.login(providerId, type, interaction); + try { + await this.refresh({ allowNetwork: this.modelNetworkEnabled }); + } catch (error) { + throw new OAuthLoginTransactionError(error); + } + return credential; + } + + async logout(providerId: string): Promise { + await this.models.logout(providerId); + // Reset credential-dependent compatibility projections before the unconfigured provider is skipped by refresh. + this.recomposeProvider(providerId); + await this.refresh({ allowNetwork: this.modelNetworkEnabled }); + } + + async refresh(options: ModelsRefreshOptions = {}): Promise { + this.config = await ModelConfig.load(this.modelsPath); + this.configureRadiusProviders(); + this.rebuildProviders(); + const refreshOptions = { + ...options, + allowNetwork: options.allowNetwork ?? this.modelNetworkEnabled, + }; + // Published pi-ai builds before ModelsStore returned void and accepted a provider ID. + // The fallback keeps source-mode CLI tests working without rebuilding workspace dependencies. + const result = ((await this.models.refresh(refreshOptions)) as ModelsRefreshResult | undefined) ?? { + aborted: refreshOptions.signal?.aborted ?? false, + errors: new Map(), + }; + this.updateModelSnapshot(); + try { + await this.forceRefreshAvailability(); + } catch { + // Availability errors are recorded by forceRefreshAvailability; refreshed models remain usable. + } + 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 { + // Validate the incoming registration on its own, like the legacy registry: + // a broken re-registration must throw without touching the stored 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 + // preserves undefined ones, matching the legacy ModelRegistry contract. + const previous = this.extensionProviders.get(providerId); + const effective: ProviderConfigInput = { ...previous }; + for (const [key, value] of Object.entries(config)) { + if (value !== undefined) (effective as Record)[key] = value; + } + this.extensionProviders.set(providerId, effective); + this.recomposeProvider(providerId); + this.updateModelSnapshot(); + if ( + this.snapshot.storedProviders.has(providerId) || + configuredRequestAuthStatus(this.config.getProvider(providerId), effective)?.configured + ) { + const configuredProviders = new Set(this.snapshot.configuredProviders).add(providerId); + const auth = new Map(this.snapshot.auth); + // Provisional entry until the async refresh lands; never clobber a real check result. + if (!auth.get(providerId)) { + auth.set(providerId, { + type: effective.oauth && !effective.apiKey ? "oauth" : "api_key", + source: "configured provider", + }); + } + this.snapshot = { + ...this.snapshot, + auth, + configuredProviders, + available: this.snapshot.all.filter((model) => configuredProviders.has(model.provider)), + }; + } + void this.refresh({ allowNetwork: false }); + } + + unregisterProvider(providerId: string): void { + this.extensionProviders.delete(providerId); + this.nativeExtensionProviders.delete(providerId); + this.recomposeProvider(providerId); + this.updateModelSnapshot(); + void this.refresh({ allowNetwork: false }); + } } diff --git a/packages/coding-agent/src/core/models-store.ts b/packages/coding-agent/src/core/models-store.ts index a7af02451..a1e8528e1 100644 --- a/packages/coding-agent/src/core/models-store.ts +++ b/packages/coding-agent/src/core/models-store.ts @@ -1,42 +1,28 @@ -import type { ModelsStore, ModelsStoreEntry } from "@earendil-works/pi-ai"; import { join } from "node:path"; +import type { ModelsStore, ModelsStoreEntry } from "@earendil-works/pi-ai"; import { getAgentDir } from "../config.ts"; -import { FileAuthStorageBackend, type AuthStorageBackend } from "./auth-storage-backends.ts"; +import { type AuthStorageBackend, FileAuthStorageBackend } from "./auth-storage-backends.ts"; type StoredModels = Record; -export interface CodingAgentModelsStore extends ModelsStore { - writeIf(providerId: string, entry: ModelsStoreEntry, predicate: () => boolean): Promise; - deleteIf(providerId: string, predicate: () => boolean): Promise; -} - -export class InMemoryCodingAgentModelsStore implements CodingAgentModelsStore { +export class InMemoryCodingAgentModelsStore implements ModelsStore { private readonly entries = new Map(); async read(providerId: string): Promise { - const entry = this.entries.get(providerId); - return entry === undefined ? undefined : structuredClone(entry); + return this.entries.get(providerId); } async write(providerId: string, entry: ModelsStoreEntry): Promise { - this.entries.set(providerId, structuredClone(entry)); + this.entries.set(providerId, entry); } async delete(providerId: string): Promise { this.entries.delete(providerId); } - - async writeIf(providerId: string, entry: ModelsStoreEntry, predicate: () => boolean): Promise { - if (predicate()) await this.write(providerId, entry); - } - - async deleteIf(providerId: string, predicate: () => boolean): Promise { - if (predicate()) await this.delete(providerId); - } } /** Locked JSON-backed storage for dynamically refreshed provider catalogs. */ -export class FileModelsStore implements CodingAgentModelsStore { +export class FileModelsStore implements ModelsStore { private readonly storage: AuthStorageBackend; constructor(path: string = join(getAgentDir(), "models-store.json")) { @@ -68,22 +54,4 @@ export class FileModelsStore implements CodingAgentModelsStore { return { result: undefined, next: JSON.stringify(current, null, 2) }; }); } - - async writeIf(providerId: string, entry: ModelsStoreEntry, predicate: () => boolean): Promise { - await this.storage.withLockAsync(async (content) => { - if (!predicate()) return { result: undefined }; - const current = this.parse(content); - current[providerId] = structuredClone(entry); - return { result: undefined, next: JSON.stringify(current, null, 2) }; - }); - } - - async deleteIf(providerId: string, predicate: () => boolean): Promise { - await this.storage.withLockAsync(async (content) => { - if (!predicate()) return { result: undefined }; - const current = this.parse(content); - delete current[providerId]; - return { result: undefined, next: JSON.stringify(current, null, 2) }; - }); - } } diff --git a/packages/coding-agent/src/core/oauth-compat.ts b/packages/coding-agent/src/core/oauth-compat.ts deleted file mode 100644 index bd6c0c4d1..000000000 --- a/packages/coding-agent/src/core/oauth-compat.ts +++ /dev/null @@ -1,30 +0,0 @@ -import type { OAuthCredentials } from "@earendil-works/pi-ai/oauth"; -import type { LegacyOAuthProvider, OAuthProviderDescriptor } from "./oauth-provider-bridge.ts"; - -export type { - OAuthAuthInfo, - OAuthCredentials, - OAuthDeviceCodeInfo, - OAuthLoginCallbacks, - OAuthPrompt, - OAuthSelectOption, - OAuthSelectPrompt, -} from "@earendil-works/pi-ai/oauth"; -export { - getOAuthApiKey, - getOAuthProvider, - getOAuthProviders, - registerOAuthProvider, - resetOAuthProviders, -} from "./oauth-provider-bridge.ts"; - -declare module "@earendil-works/pi-ai/oauth" { - export function getOAuthApiKey( - providerId: string, - credentials: Record, - ): Promise<{ newCredentials: OAuthCredentials; apiKey: string } | null>; - export function getOAuthProvider(providerId: string): OAuthProviderDescriptor | undefined; - export function getOAuthProviders(): OAuthProviderDescriptor[]; - export function registerOAuthProvider(provider: LegacyOAuthProvider & { id: string }): void; - export function resetOAuthProviders(): void; -} diff --git a/packages/coding-agent/src/core/oauth-login.ts b/packages/coding-agent/src/core/oauth-login.ts new file mode 100644 index 000000000..d688d42e1 --- /dev/null +++ b/packages/coding-agent/src/core/oauth-login.ts @@ -0,0 +1,95 @@ +import type { AuthInfoLink, AuthInteraction, OAuthLoginCallbacks } from "@earendil-works/pi-ai"; + +export interface AtomicOAuthLoginCallbacks extends OAuthLoginCallbacks { + onManualCodeCancel?(): void; + onInfo?(message: string, links: readonly AuthInfoLink[]): void; +} + +/** JSON-safe provider-owned OAuth metadata transported across isolated-engine boundaries. */ +export interface OAuthProviderMetadata { + id: string; + name: string; + loginLabel?: string; + usesCallbackServer?: boolean; +} + +/** Marks failures after credential acquisition so presentation never mistakes nested AbortErrors for cancellation. */ +export class OAuthLoginTransactionError extends Error { + constructor(cause: unknown) { + super(cause instanceof Error ? cause.message : String(cause), { cause }); + this.name = "OAuthLoginTransactionError"; + } +} + +export function isOAuthLoginCancelled( + error: unknown, + signal?: AbortSignal, + options: { includeActiveSignal?: boolean } = {}, +): boolean { + if (error instanceof OAuthLoginTransactionError) return false; + if (options.includeActiveSignal !== false && signal?.aborted) return true; + const reason = signal?.reason; + const seen = new Set(); + let current: unknown = error; + while (current !== undefined && current !== null) { + if (current === reason && reason !== undefined) return true; + if (current === "Login cancelled") return true; + if (typeof current !== "object") return false; + if (seen.has(current)) return false; + seen.add(current); + const candidate = current as { name?: unknown; message?: unknown; cause?: unknown }; + if (candidate.name === "AbortError" || candidate.message === "Login cancelled") return true; + current = candidate.cause; + } + return false; +} + +export function normalizeOAuthLoginError( + error: unknown, + signal?: AbortSignal, + options?: { includeActiveSignal?: boolean }, +): unknown { + if (error instanceof Error && error.message === "Login cancelled" && error.cause !== undefined) return error; + if (!isOAuthLoginCancelled(error, signal, options)) return error; + return new Error("Login cancelled", { cause: error }); +} + +function abortable(promise: Promise, signal: AbortSignal | undefined, onAbort?: () => void): Promise { + if (!signal) return promise; + if (signal.aborted) { + onAbort?.(); + return Promise.reject(signal.reason ?? new Error("Authentication prompt cancelled")); + } + return new Promise((resolve, reject) => { + const abort = () => { + onAbort?.(); + reject(signal.reason ?? new Error("Authentication prompt cancelled")); + }; + signal.addEventListener("abort", abort, { once: true }); + promise.then(resolve, reject).finally(() => signal.removeEventListener("abort", abort)); + }); +} + +export function createAuthInteraction(callbacks: AtomicOAuthLoginCallbacks): AuthInteraction { + return { + signal: callbacks.signal, + prompt: async (prompt) => { + switch (prompt.type) { + case "select": + return (await callbacks.onSelect({ message: prompt.message, options: prompt.options.map(({ id, label }) => ({ id, label })) })) ?? ""; + case "manual_code": + return abortable(callbacks.onManualCodeInput ? callbacks.onManualCodeInput() : callbacks.onPrompt({ message: prompt.message, placeholder: prompt.placeholder }), prompt.signal, callbacks.onManualCodeCancel); + default: + return callbacks.onPrompt({ message: prompt.message, placeholder: prompt.placeholder }); + } + }, + notify: (event) => { + switch (event.type) { + case "auth_url": callbacks.onAuth({ url: event.url, instructions: event.instructions }); break; + case "device_code": callbacks.onDeviceCode(event); break; + case "progress": callbacks.onProgress?.(event.message); break; + case "info": callbacks.onInfo?.(event.message, event.links ?? []); break; + } + }, + }; +} diff --git a/packages/coding-agent/src/core/oauth-provider-bridge.ts b/packages/coding-agent/src/core/oauth-provider-bridge.ts deleted file mode 100644 index 8e88b9b0a..000000000 --- a/packages/coding-agent/src/core/oauth-provider-bridge.ts +++ /dev/null @@ -1,295 +0,0 @@ -import { - type AuthInfoLink, - type AuthInteraction, - type ModelAuth, - type OAuthAuth, - type OAuthCredential, - type OAuthCredentials, - type OAuthLoginCallbacks, -} from "@earendil-works/pi-ai"; -import { builtinProviders } from "@earendil-works/pi-ai/providers/all"; -import type { Api, Model } from "@earendil-works/pi-ai/compat"; - -export interface LegacyOAuthProvider { - name: string; - loginLabel?: string; - usesCallbackServer?: boolean; - login(callbacks: OAuthLoginCallbacks): Promise; - refreshToken(credentials: OAuthCredentials): Promise; - getApiKey(credentials: OAuthCredentials): string; - modifyModels?(models: Model[], credentials: OAuthCredentials): Model[]; -} - -export interface OAuthProviderDescriptor extends LegacyOAuthProvider { - id: string; - loginLabel?: string; -} - - -/** JSON-safe OAuth metadata transported across isolated-engine boundaries. */ -export interface OAuthProviderMetadata { - id: string; - name: string; - loginLabel?: string; - usesCallbackServer?: boolean; -} -export interface AtomicOAuthLoginCallbacks extends OAuthLoginCallbacks { - onManualCodeCancel?(): void; - onInfo?(message: string, links: readonly AuthInfoLink[]): void; -} - -const GLOBAL_LEGACY_SOURCE = "atomic:legacy-global"; -const legacyProviders = new Map>(); -const CALLBACK_SERVER_PROVIDERS = new Set(["anthropic", "openai-codex"]); -/** Marks failures after credential acquisition so presentation never mistakes nested AbortErrors for cancellation. */ -export class OAuthLoginTransactionError extends Error { - constructor(cause: unknown) { - super(cause instanceof Error ? cause.message : String(cause), { cause }); - this.name = "OAuthLoginTransactionError"; - } -} - -/** Cycle-safe classification shared by direct, provider-owned, and isolated OAuth. */ -export function isOAuthLoginCancelled( - error: unknown, - signal?: AbortSignal, - options: { includeActiveSignal?: boolean } = {}, -): boolean { - if (error instanceof OAuthLoginTransactionError) return false; - if (options.includeActiveSignal !== false && signal?.aborted) return true; - const reason = signal?.reason; - const seen = new Set(); - let current: unknown = error; - while (current !== undefined && current !== null) { - if (current === reason && reason !== undefined) return true; - if (current === "Login cancelled") return true; - if (typeof current !== "object") return false; - if (seen.has(current)) return false; - seen.add(current); - const candidate = current as { name?: unknown; message?: unknown; cause?: unknown }; - if (candidate.name === "AbortError" || candidate.message === "Login cancelled") return true; - current = candidate.cause; - } - return false; -} - -/** Canonicalize only intentional cancellation while retaining the provider error as cause. */ -export function normalizeOAuthLoginError( - error: unknown, - signal?: AbortSignal, - options?: { includeActiveSignal?: boolean }, -): unknown { - if (error instanceof Error && error.message === "Login cancelled" && error.cause !== undefined) return error; - if (!isOAuthLoginCancelled(error, signal, options)) return error; - return new Error("Login cancelled", { cause: error }); -} - -function abortable(promise: Promise, signal: AbortSignal | undefined, onAbort?: () => void): Promise { - if (!signal) return promise; - if (signal.aborted) { - onAbort?.(); - return Promise.reject(signal.reason ?? new Error("Authentication prompt cancelled")); - } - return new Promise((resolve, reject) => { - const abort = () => { - onAbort?.(); - reject(signal.reason ?? new Error("Authentication prompt cancelled")); - }; - signal.addEventListener("abort", abort, { once: true }); - promise.then(resolve, reject).finally(() => signal.removeEventListener("abort", abort)); - }); -} -function builtinOAuth(providerId: string): OAuthAuth | undefined { - return builtinProviders().find((provider) => provider.id === providerId)?.auth.oauth; -} - -export function createAuthInteraction(callbacks: AtomicOAuthLoginCallbacks): AuthInteraction { - return { - signal: callbacks.signal, - prompt: async (prompt) => { - switch (prompt.type) { - case "select": - return ( - (await callbacks.onSelect({ - message: prompt.message, - options: prompt.options.map(({ id, label }) => ({ id, label })), - })) ?? "" - ); - case "manual_code": - return abortable( - callbacks.onManualCodeInput - ? callbacks.onManualCodeInput() - : callbacks.onPrompt({ message: prompt.message, placeholder: prompt.placeholder }), - prompt.signal, - callbacks.onManualCodeCancel, - ); - default: - return callbacks.onPrompt({ message: prompt.message, placeholder: prompt.placeholder }); - } - }, - notify: (event) => { - switch (event.type) { - case "auth_url": - callbacks.onAuth({ url: event.url, instructions: event.instructions }); - break; - case "device_code": - callbacks.onDeviceCode(event); - break; - case "progress": - callbacks.onProgress?.(event.message); - break; - case "info": - callbacks.onInfo?.(event.message, event.links ?? []); - break; - } - }, - }; -} - -function latestLegacyProvider(providerId: string): LegacyOAuthProvider | undefined { - const providers = legacyProviders.get(providerId); - return providers ? [...providers.values()].at(-1) : undefined; -} - -export function registerLegacyOAuthProvider( - providerId: string, - provider: LegacyOAuthProvider, - sourceId: string = GLOBAL_LEGACY_SOURCE, -): void { - const providers = legacyProviders.get(providerId) ?? new Map(); - providers.delete(sourceId); - providers.set(sourceId, provider); - legacyProviders.set(providerId, providers); -} -/** Legacy registry aliases retained for Atomic's extension/test compatibility. */ -export function registerOAuthProvider(provider: LegacyOAuthProvider & { id: string }): void { - registerLegacyOAuthProvider(provider.id, provider); -} - -export function getOAuthProvider(providerId: string): OAuthProviderDescriptor | undefined { - const provider = latestLegacyProvider(providerId); - return provider ? { id: providerId, ...provider } : getOAuthProviderDescriptors().find((entry) => entry.id === providerId); -} - -export function resetLegacyOAuthProviders(): void { - legacyProviders.clear(); -} - -export function unregisterLegacyOAuthProviders(sourceId: string): void { - for (const [providerId, providers] of legacyProviders) { - providers.delete(sourceId); - if (providers.size === 0) legacyProviders.delete(providerId); - } -} - -export function getLegacyOAuthProvider(providerId: string): LegacyOAuthProvider | undefined { - return latestLegacyProvider(providerId); -} - -export function getOAuthProviders(): OAuthProviderDescriptor[] { - return getOAuthProviderDescriptors(); -} - -export const resetOAuthProviders = resetLegacyOAuthProviders; -export function getOAuthProviderDescriptors(): OAuthProviderDescriptor[] { - const descriptors = new Map(); - for (const provider of builtinProviders()) { - const oauth = provider.auth.oauth; - if (!oauth) continue; - descriptors.set(provider.id, { - id: provider.id, - name: oauth.name, - loginLabel: oauth.loginLabel, - usesCallbackServer: CALLBACK_SERVER_PROVIDERS.has(provider.id), - login: async (callbacks) => oauth.login(createAuthInteraction(callbacks)), - refreshToken: async (credentials) => oauth.refresh({ ...credentials, type: "oauth" }), - getApiKey: (credentials) => credentials.access, - }); - } - for (const [id, providers] of legacyProviders) { - const provider = [...providers.values()].at(-1); - if (provider) descriptors.set(id, { id, ...provider }); - } - return [...descriptors.values()]; -} - -export function getOAuthProviderMetadata(): OAuthProviderMetadata[] { - return getOAuthProviderDescriptors().map(({ id, name, loginLabel, usesCallbackServer }) => ({ - id, - name, - ...(loginLabel === undefined ? {} : { loginLabel }), - ...(usesCallbackServer === undefined ? {} : { usesCallbackServer }), - })); -} - -export async function loginOAuthProvider( - providerId: string, - callbacks: AtomicOAuthLoginCallbacks, -): Promise { - const legacy = latestLegacyProvider(providerId); - const oauth = legacy ? undefined : builtinOAuth(providerId); - if (!legacy && !oauth) throw new Error(`Unknown OAuth provider: ${providerId}`); - if (callbacks.signal?.aborted) { - throw normalizeOAuthLoginError(callbacks.signal.reason, callbacks.signal); - } - try { - const credential = legacy - ? { type: "oauth" as const, ...(await legacy.login(callbacks)) } - : await oauth!.login(createAuthInteraction(callbacks)); - if (callbacks.signal?.aborted) { - throw normalizeOAuthLoginError(callbacks.signal.reason, callbacks.signal); - } - return credential; - } catch (error) { - throw normalizeOAuthLoginError(error, callbacks.signal); - } -} - -export async function refreshOAuthProvider( - providerId: string, - credential: OAuthCredential, -): Promise<{ credential: OAuthCredential; auth: ModelAuth } | undefined> { - const legacy = latestLegacyProvider(providerId); - if (legacy) { - const { type: _type, ...credentials } = credential; - const refreshed = { type: "oauth" as const, ...(await legacy.refreshToken(credentials)) }; - return { credential: refreshed, auth: { apiKey: legacy.getApiKey(refreshed) } }; - } - const oauth = builtinOAuth(providerId); - if (!oauth) return undefined; - const refreshed = await oauth.refresh(credential); - return { credential: refreshed, auth: await oauth.toAuth(refreshed) }; -} - -export async function oauthCredentialToAuth( - providerId: string, - credential: OAuthCredential, -): Promise { - const legacy = latestLegacyProvider(providerId); - if (legacy) return { apiKey: legacy.getApiKey(credential) }; - return builtinOAuth(providerId)?.toAuth(credential); -} - -export async function getOAuthApiKey( - providerId: string, - credentials: Record, -): Promise<{ newCredentials: OAuthCredentials; apiKey: string } | null> { - if (!getOAuthProvider(providerId)) throw new Error(`Unknown OAuth provider: ${providerId}`); - const selected = credentials[providerId]; - if (!selected) return null; - const credential: OAuthCredential = { type: "oauth", ...selected }; - let resolved: { credential: OAuthCredential; auth: ModelAuth | undefined } | undefined; - if (Date.now() >= credential.expires) { - try { - resolved = await refreshOAuthProvider(providerId, credential); - } catch (cause) { - throw new Error(`Failed to refresh OAuth token for ${providerId}`, { cause }); - } - } else { - resolved = { credential, auth: await oauthCredentialToAuth(providerId, credential) }; - } - const apiKey = resolved?.auth?.apiKey; - if (!resolved || !apiKey) return null; - const { type: _type, ...newCredentials } = resolved.credential; - return { newCredentials, apiKey }; -} diff --git a/packages/coding-agent/src/core/oauth-provider-metadata.ts b/packages/coding-agent/src/core/oauth-provider-metadata.ts new file mode 100644 index 000000000..618e14c02 --- /dev/null +++ b/packages/coding-agent/src/core/oauth-provider-metadata.ts @@ -0,0 +1,26 @@ +import type { Provider } from "@earendil-works/pi-ai"; +import type { OAuthProviderMetadata } from "./oauth-login.ts"; +import type { ProviderConfigInput } from "./provider-composer.ts"; + +const CALLBACK_SERVER_PROVIDERS = new Set(["anthropic", "openai-codex"]); + +export function collectOAuthProviderMetadata( + providers: readonly Provider[], + extensions: ReadonlyMap, +): OAuthProviderMetadata[] { + return providers.filter((provider) => provider.auth.oauth).map((provider) => { + const providerOAuth = provider.auth.oauth; + const extensionOAuth = extensions.get(provider.id)?.oauth; + const loginLabel = extensionOAuth?.loginLabel ?? providerOAuth?.loginLabel; + const hasCallbackServerMetadata = extensionOAuth?.usesCallbackServer !== undefined + || CALLBACK_SERVER_PROVIDERS.has(provider.id); + const usesCallbackServer = extensionOAuth?.usesCallbackServer + ?? CALLBACK_SERVER_PROVIDERS.has(provider.id); + return { + id: provider.id, + name: provider.name ?? provider.id, + ...(loginLabel ? { loginLabel } : {}), + ...(hasCallbackServerMetadata ? { usesCallbackServer } : {}), + }; + }); +} diff --git a/packages/coding-agent/src/core/provider-composer-internal.ts b/packages/coding-agent/src/core/provider-composer-internal.ts new file mode 100644 index 000000000..0f4b84988 --- /dev/null +++ b/packages/coding-agent/src/core/provider-composer-internal.ts @@ -0,0 +1,397 @@ +import { + type Api, + type ApiKeyAuth, + type AssistantMessageEventStream, + type AuthContext, + type AuthInteraction, + type AuthResult, + type Context, + type Model, + type ModelAuth, + type OAuthAuth, + type OAuthCredentials, + type OAuthLoginCallbacks, + type Provider, + type ProviderHeaders, + type RefreshModelsContext, + type SimpleStreamOptions, +} from "@earendil-works/pi-ai"; +import type { ModelsJsonModel, ModelsJsonModelOverride, ModelsJsonProvider } from "./model-config.ts"; +import { + clearConfigValueCache, + getConfigValueEnvVarNames, + isCommandConfigValue, + resolveConfigValueOrThrow, + resolveHeadersOrThrow, +} from "./resolve-config-value.ts"; + +export interface ExtensionOAuthConfig { + name: string; + loginLabel?: string; + usesCallbackServer?: boolean; + login(callbacks: OAuthLoginCallbacks): Promise; + refreshToken(credentials: OAuthCredentials): Promise; + getApiKey(credentials: OAuthCredentials): string; + modifyModels?(models: Model[], credentials: OAuthCredentials): Model[]; +} + +/** Input type for the extension registerProvider API. */ +export interface ProviderConfigInput { + name?: string; + baseUrl?: string; + apiKey?: string; + api?: Api; + streamSimple?: (model: Model, context: Context, options?: SimpleStreamOptions) => AssistantMessageEventStream; + headers?: Record; + authHeader?: boolean; + oauth?: ExtensionOAuthConfig; + models?: Array<{ + id: string; + name: string; + api?: Api; + baseUrl?: string; + reasoning: boolean; + thinkingLevelMap?: Model["thinkingLevelMap"]; + input: ("text" | "image")[]; + cost: Model["cost"]; + contextWindow: number; + maxTokens: number; + headers?: Record; + compat?: Model["compat"]; + }>; + refreshModels?(context: RefreshModelsContext): Promise>; +} + +export type AuthStatus = { + configured: boolean; + source?: "stored" | "runtime" | "environment" | "fallback" | "models_json_key" | "models_json_command"; + label?: string; +}; + +export const clearApiKeyCache = clearConfigValueCache; + +function mergeCompat( + base: Model["compat"], + override: Model["compat"] | ModelsJsonModelOverride["compat"], +): Model["compat"] { + if (!override) return base; + const merged = { ...base, ...override } as NonNullable["compat"]>; + const baseNested = base as Record | undefined; + const overrideNested = override as Record; + const mergedNested = merged as Record; + for (const key of ["openRouterRouting", "vercelGatewayRouting", "chatTemplateKwargs"] as const) { + const baseValue = baseNested?.[key]; + const overrideValue = overrideNested[key]; + if ( + (typeof baseValue === "object" && baseValue !== null) || + (typeof overrideValue === "object" && overrideValue !== null) + ) { + mergedNested[key] = { ...(baseValue as object | undefined), ...(overrideValue as object | undefined) }; + } + } + return merged; +} + +export function applyModelOverride(model: Model, override: ModelsJsonModelOverride): Model { + return { + ...model, + name: override.name ?? model.name, + reasoning: override.reasoning ?? model.reasoning, + thinkingLevelMap: override.thinkingLevelMap + ? { ...model.thinkingLevelMap, ...override.thinkingLevelMap } + : model.thinkingLevelMap, + input: (override.input as ("text" | "image")[] | undefined) ?? model.input, + cost: override.cost + ? { + input: override.cost.input ?? model.cost.input, + output: override.cost.output ?? model.cost.output, + cacheRead: override.cost.cacheRead ?? model.cost.cacheRead, + cacheWrite: override.cost.cacheWrite ?? model.cost.cacheWrite, + tiers: override.cost.tiers ?? model.cost.tiers, + } + : model.cost, + contextWindow: override.contextWindow ?? model.contextWindow, + maxTokens: override.maxTokens ?? model.maxTokens, + compat: mergeCompat(model.compat, override.compat), + }; +} + +function modelFromJson( + providerId: string, + definition: ModelsJsonModel, + providerConfig: ModelsJsonProvider, + defaults: Model | undefined, +): Model { + const api = definition.api ?? providerConfig.api ?? defaults?.api; + if (!api) { + throw new Error( + `Provider ${providerId}, model ${definition.id}: no "api" specified. Set at provider or model level.`, + ); + } + const baseUrl = definition.baseUrl ?? providerConfig.baseUrl ?? defaults?.baseUrl; + if (!baseUrl) throw new Error(`Provider ${providerId}: "baseUrl" is required when defining custom models.`); + if (definition.contextWindow !== undefined && definition.contextWindow <= 0) { + throw new Error(`Provider ${providerId}, model ${definition.id}: invalid contextWindow`); + } + if (definition.maxTokens !== undefined && definition.maxTokens <= 0) { + throw new Error(`Provider ${providerId}, model ${definition.id}: invalid maxTokens`); + } + return { + id: definition.id, + name: definition.name ?? definition.id, + api: api as Api, + provider: providerId, + baseUrl, + reasoning: definition.reasoning ?? false, + thinkingLevelMap: definition.thinkingLevelMap, + input: (definition.input ?? ["text"]) as ("text" | "image")[], + cost: definition.cost ?? { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: definition.contextWindow ?? 128000, + maxTokens: definition.maxTokens ?? 16384, + headers: undefined, + compat: mergeCompat(providerConfig.compat, definition.compat), + }; +} + +export function applyModelsJson( + providerId: string, + baseModels: readonly Model[], + config: ModelsJsonProvider | undefined, +): Model[] { + if (!config) return [...baseModels]; + if (config.oauth && !config.baseUrl) { + throw new Error(`Provider ${providerId}: "baseUrl" is required when "oauth" is set.`); + } + const hasOverrides = config.modelOverrides && Object.keys(config.modelOverrides).length > 0; + if ( + !config.models?.length && + !config.baseUrl && + !config.headers && + !config.compat && + !hasOverrides && + !config.apiKey && + !config.oauth && + config.authHeader === undefined + ) { + throw new Error( + `Provider ${providerId}: must specify "baseUrl", "headers", "compat", "modelOverrides", or "models".`, + ); + } + + const models: Model[] = baseModels.map((model) => ({ + ...model, + baseUrl: config.oauth === "radius" ? model.baseUrl : (config.baseUrl ?? model.baseUrl), + compat: mergeCompat(model.compat, config.compat), + })); + for (const definition of config.models ?? []) { + const existingIndex = models.findIndex((model) => model.id === definition.id); + const defaults = existingIndex >= 0 ? models[existingIndex] : models[0]; + const model = modelFromJson(providerId, definition, config, defaults); + if (existingIndex >= 0) models[existingIndex] = model; + else models.push(model); + } + return models; +} + +export function applyExtension( + providerId: string, + models: readonly Model[], + config: ProviderConfigInput | undefined, +): Model[] { + if (!config) return [...models]; + if (!config.models) { + return config.baseUrl ? models.map((model) => ({ ...model, baseUrl: config.baseUrl! })) : [...models]; + } + return config.models.map((definition) => { + const defaults = models.find((model) => model.id === definition.id) ?? models[0]; + const api = definition.api ?? config.api ?? defaults?.api; + if (!api) { + throw new Error( + `Provider ${providerId}, model ${definition.id}: no "api" specified. Set at provider or model level.`, + ); + } + const baseUrl = definition.baseUrl ?? config.baseUrl ?? defaults?.baseUrl; + if (!baseUrl) throw new Error(`Provider ${providerId}: "baseUrl" is required when defining custom models.`); + return { + ...definition, + api, + provider: providerId, + baseUrl, + headers: undefined, + }; + }); +} + +function adaptOAuth(config: ExtensionOAuthConfig): OAuthAuth { + return { + name: config.name, + login: async (callbacks) => { + const legacyCallbacks: OAuthLoginCallbacks & { + onInfo(message: string, links: readonly { label?: string; url: string }[]): void; + } = { + onAuth: (info) => callbacks.notify({ type: "auth_url", ...info }), + onDeviceCode: (info) => callbacks.notify({ type: "device_code", ...info }), + onPrompt: (prompt) => callbacks.prompt({ type: "text", ...prompt }), + onProgress: (message) => callbacks.notify({ type: "progress", message }), + onInfo: (message, links) => callbacks.notify({ type: "info", message, links }), + onManualCodeInput: () => callbacks.prompt({ type: "manual_code", message: "Paste the authorization code" }), + onSelect: (prompt) => callbacks.prompt({ type: "select", ...prompt }), + signal: callbacks.signal, + }; + const credential = await config.login(legacyCallbacks); + return { ...credential, type: "oauth" }; + }, + refresh: async (credential) => ({ ...(await config.refreshToken(credential)), type: "oauth" }), + toAuth: async (credential) => ({ apiKey: config.getApiKey(credential) }), + }; +} + +function withConfiguredAuth( + auth: ModelAuth, + headers: Record | undefined, + authHeader: boolean, +): ModelAuth { + let mergedHeaders: ProviderHeaders | undefined = + auth.headers || headers ? { ...auth.headers, ...headers } : undefined; + if (authHeader) { + if (!auth.apiKey) throw new Error("authHeader requires a resolved API key"); + mergedHeaders = { ...mergedHeaders, Authorization: `Bearer ${auth.apiKey}` }; + } + return { ...auth, headers: mergedHeaders }; +} + +export function configuredApiKey( + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, +): string | undefined { + return extension?.apiKey ?? config?.apiKey; +} + +export function configuredHeaders( + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, +): Record | undefined { + if (!config?.headers && !extension?.headers) return undefined; + return { ...config?.headers, ...extension?.headers }; +} + +async function configContextEnv( + values: readonly string[], + ctx: AuthContext, + explicit?: Record, +): Promise | undefined> { + const env = { ...explicit }; + for (const name of new Set(values.flatMap(getConfigValueEnvVarNames))) { + if (env[name] !== undefined) continue; + const value = await ctx.env(name); + if (value !== undefined) env[name] = value; + } + return Object.keys(env).length > 0 ? env : undefined; +} + +export function composeApiKeyAuth( + providerId: string, + base: Provider | undefined, + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, +): ApiKeyAuth | undefined { + const inherited = base?.auth.apiKey; + const rawKey = configuredApiKey(config, extension); + const oauth = extension?.oauth ?? base?.auth.oauth; + // OAuth-only providers get no fabricated API-key login method. + if (!inherited && rawKey === undefined && oauth) return undefined; + const rawHeaders = configuredHeaders(config, extension); + const authHeader = extension?.authHeader ?? config?.authHeader ?? false; + return { + name: inherited?.name ?? "API key", + login: + inherited?.login ?? + (async (interaction: AuthInteraction) => ({ + type: "api_key", + key: await interaction.prompt({ type: "secret", message: "Enter API key" }), + })), + check: async (input) => { + if (input.credential) { + if (inherited?.check) return inherited.check(input); + if (input.credential.key) return { type: "api_key", source: "stored credential" }; + const resolved = await inherited?.resolve(input); + return resolved ? { type: "api_key", source: resolved.source } : undefined; + } + if (rawKey !== undefined) { + if (isCommandConfigValue(rawKey)) return { type: "api_key", source: "configured API key" }; + const envNames = getConfigValueEnvVarNames(rawKey); + for (const name of envNames) { + if ((await input.ctx.env(name)) === undefined) return undefined; + } + return { type: "api_key", source: "configured API key" }; + } + if (inherited?.check) return inherited.check(input); + const resolved = await inherited?.resolve(input); + return resolved ? { type: "api_key", source: resolved.source } : undefined; + }, + resolve: async (input) => { + let result: AuthResult | undefined; + if (input.credential) { + result = inherited + ? await inherited.resolve(input) + : input.credential.key + ? { auth: { apiKey: input.credential.key }, env: input.credential.env, source: "stored credential" } + : undefined; + } else if (rawKey !== undefined) { + const env = await configContextEnv([rawKey], input.ctx); + const key = resolveConfigValueOrThrow(rawKey, `API key for provider "${providerId}"`, env); + result = inherited + ? await inherited.resolve({ ...input, credential: { type: "api_key", key } }) + : { auth: { apiKey: key }, source: "configured API key" }; + } else { + result = await inherited?.resolve(input); + } + if (!result) return undefined; + const explicitEnv = { ...(input.credential?.env ?? {}), ...(result.env ?? {}) }; + const headerEnv = await configContextEnv(Object.values(rawHeaders ?? {}), input.ctx, explicitEnv); + const headers = resolveHeadersOrThrow(rawHeaders, `provider "${providerId}"`, headerEnv); + return { ...result, auth: withConfiguredAuth(result.auth, headers, authHeader) }; + }, + }; +} + +export function composeOAuthAuth( + providerId: string, + base: Provider | undefined, + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, +): OAuthAuth | undefined { + const oauth = extension?.oauth ? adaptOAuth(extension.oauth) : base?.auth.oauth; + if (!oauth) return undefined; + const rawHeaders = configuredHeaders(config, extension); + const authHeader = extension?.authHeader ?? config?.authHeader ?? false; + return { + ...oauth, + toAuth: async (credential) => { + const auth = await oauth.toAuth(credential); + const env = credential.env; + const headers = resolveHeadersOrThrow( + rawHeaders, + `provider "${providerId}"`, + typeof env === "object" && env !== null ? (env as Record) : undefined, + ); + return withConfiguredAuth(auth, headers, authHeader); + }, + }; +} + +export function rawModelHeaders( + model: Model, + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, +): Record | undefined { + const definition = config?.models?.find((entry) => entry.id === model.id); + const extensionModel = extension?.models?.find((entry) => entry.id === model.id); + const headers = { + ...config?.modelOverrides?.[model.id]?.headers, + ...definition?.headers, + ...extensionModel?.headers, + }; + return Object.keys(headers).length > 0 ? headers : undefined; +} + diff --git a/packages/coding-agent/src/core/provider-composer.ts b/packages/coding-agent/src/core/provider-composer.ts new file mode 100644 index 000000000..5a46c2b7a --- /dev/null +++ b/packages/coding-agent/src/core/provider-composer.ts @@ -0,0 +1,186 @@ +import { + type Api, + type AssistantMessageEventStream, + type Context, + type Credential, + lazyStream, + type Model, + type OAuthCredentials, + type Provider, + type ProviderHeaders, + type SimpleStreamOptions, + type StreamOptions, +} from "@earendil-works/pi-ai"; +import { getApiProvider } from "@earendil-works/pi-ai/compat"; +import type { ModelConfig, ModelsJsonProvider } from "./model-config.ts"; +import { + applyExtension, + applyModelOverride, + applyModelsJson, + composeApiKeyAuth, + composeOAuthAuth, + configuredApiKey, + configuredHeaders, + rawModelHeaders, +} from "./provider-composer-internal.ts"; +import type { AuthStatus, ProviderConfigInput } from "./provider-composer-internal.ts"; +import { + getConfigValueEnvVarNames, + isCommandConfigValue, + isConfigValueConfigured, + resolveHeadersOrThrow, +} from "./resolve-config-value.ts"; + +export { clearApiKeyCache } from "./provider-composer-internal.ts"; +export type { AuthStatus, ExtensionOAuthConfig, ProviderConfigInput } from "./provider-composer-internal.ts"; + +export function validateExtensionProvider( + providerId: string, + base: Provider | undefined, + modelsConfig: ModelsJsonProvider | undefined, + extension: ProviderConfigInput, +): void { + if (extension.streamSimple && !extension.api) { + throw new Error(`Provider ${providerId}: "api" is required when registering streamSimple.`); + } + applyExtension(providerId, applyModelsJson(providerId, base?.getModels() ?? [], modelsConfig), extension); +} + +/** Compose built-in, models.json, and extension layers without reading credentials. */ +export function composeModelProvider( + providerId: string, + base: Provider | undefined, + modelConfig: ModelConfig, + extension: ProviderConfigInput | undefined, +): Provider { + const config = modelConfig.getProvider(providerId); + let extensionOAuthCredential: OAuthCredentials | undefined; + let refreshedExtensionModels: ProviderConfigInput["models"]; + const currentExtension = (): ProviderConfigInput | undefined => + extension && refreshedExtensionModels ? { ...extension, models: refreshedExtensionModels } : extension; + // models.json modelOverrides are the topmost user-config layer: they apply once, + // after custom-model upserts, extension model replacement, and legacy OAuth projection. + const getModels = () => { + let models = applyExtension( + providerId, + applyModelsJson(providerId, base?.getModels() ?? [], config), + currentExtension(), + ); + if (extensionOAuthCredential && extension?.oauth?.modifyModels) { + models = extension.oauth.modifyModels(models, extensionOAuthCredential); + } + return models.map((model) => { + const override = config?.modelOverrides?.[model.id]; + return override ? applyModelOverride(model, override) : model; + }); + }; + // Validate eagerly so registration/reload reports structural errors immediately. + getModels(); + const apiKey = composeApiKeyAuth(providerId, base, config, extension); + const oauth = composeOAuthAuth(providerId, base, config, extension); + if (!apiKey && !oauth) throw new Error(`Provider ${providerId}: no authentication method configured.`); + + const supportsBaseApi = (model: Model) => base?.getModels().some((entry) => entry.api === model.api) ?? false; + const streamWith = ( + model: Model, + context: Context, + options: StreamOptions | undefined, + simple: boolean, + ): AssistantMessageEventStream => + lazyStream(model, async () => { + if (extension?.streamSimple && model.api === extension.api) { + return extension.streamSimple(model, context, options as SimpleStreamOptions); + } + if (base && supportsBaseApi(model)) { + return simple + ? base.streamSimple(model, context, options as SimpleStreamOptions) + : base.stream(model, context, options); + } + const api = getApiProvider(model.api); + if (!api) throw new Error(`No API provider registered for api: ${model.api}`); + return simple + ? api.streamSimple(model, context, options as SimpleStreamOptions) + : api.stream(model, context, options); + }); + + return { + id: providerId, + name: extension?.name ?? config?.name ?? base?.name ?? extension?.oauth?.name ?? providerId, + baseUrl: extension?.baseUrl ?? config?.baseUrl ?? base?.baseUrl, + headers: base?.headers, + auth: { ...(apiKey ? { apiKey } : {}), ...(oauth ? { oauth } : {}) }, + getModels, + refreshModels: + base?.refreshModels || extension?.refreshModels || extension?.oauth?.modifyModels + ? async (context) => { + await base?.refreshModels?.(context); + if (extension?.refreshModels) { + const refreshed = await extension.refreshModels(context); + if (!context.signal?.aborted) { + // Validate before publishing the new synchronous list. + applyExtension(providerId, applyModelsJson(providerId, base?.getModels() ?? [], config), { + ...extension, + models: refreshed, + }); + refreshedExtensionModels = refreshed; + } + } + extensionOAuthCredential = context.credential?.type === "oauth" ? context.credential : undefined; + } + : undefined, + filterModels: base?.filterModels + ? (models, credential: Credential | undefined) => base.filterModels!(models, credential) + : undefined, + stream: (model, context, options) => streamWith(model, context, options, false), + streamSimple: (model, context, options) => streamWith(model, context, options, true), + }; +} + +export function resolveConfiguredModelHeaders( + model: Model, + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, + env?: Record, +): Record | undefined { + return resolveHeadersOrThrow( + rawModelHeaders(model, config, extension), + `model "${model.provider}/${model.id}"`, + env, + ); +} + +export interface CompatibilityRequestConfig { + headers?: ProviderHeaders; + authHeader: boolean; +} + +export function resolveCompatibilityRequestConfig( + model: Model, + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, +): CompatibilityRequestConfig { + const configured = resolveHeadersOrThrow( + { ...configuredHeaders(config, extension), ...rawModelHeaders(model, config, extension) }, + `model "${model.provider}/${model.id}"`, + ); + return { + headers: model.headers || configured ? { ...model.headers, ...configured } : undefined, + authHeader: extension?.authHeader ?? config?.authHeader ?? false, + }; +} + +export function configuredRequestAuthStatus( + config: ModelsJsonProvider | undefined, + extension: ProviderConfigInput | undefined, +): AuthStatus | undefined { + const value = configuredApiKey(config, extension); + if (value === undefined) return undefined; + if (isCommandConfigValue(value)) return { configured: true, source: "models_json_command" }; + const names = getConfigValueEnvVarNames(value); + if (names.length > 0) { + return isConfigValueConfigured(value) + ? { configured: true, source: "environment", label: names.join(", ") } + : { configured: false }; + } + return { configured: true, source: extension?.apiKey !== undefined ? "fallback" : "models_json_key" }; +} diff --git a/packages/coding-agent/src/core/remote-catalog-provider.ts b/packages/coding-agent/src/core/remote-catalog-provider.ts index 5d1e3320c..56147035d 100644 --- a/packages/coding-agent/src/core/remote-catalog-provider.ts +++ b/packages/coding-agent/src/core/remote-catalog-provider.ts @@ -15,18 +15,24 @@ function mergeModels(baseline: readonly Model[], dynamic: readonly Model[] { +function parseCatalog(providerId: string, value: unknown): Model[] { const entries = Array.isArray(value) ? value - : "models" in value && Array.isArray(value.models) + : typeof value === "object" && value !== null && "models" in value && Array.isArray(value.models) ? value.models - : Object.values(value); + : typeof value === "object" && value !== null + ? Object.values(value) + : undefined; + if (!entries) throw new Error(`Invalid model catalog for provider "${providerId}"`); return entries .filter((entry): entry is Model => typeof entry === "object" && entry !== null && "id" in entry) .map((model) => ({ ...model, provider: providerId })); } -function remoteModels(entry: ModelsStoreEntry | undefined, localGeneratedAt: number | undefined): readonly Model[] { +function remoteModels( + entry: ModelsStoreEntry | undefined, + localGeneratedAt: number | undefined, +): readonly Model[] { if (!entry) return []; if (localGeneratedAt !== undefined && (entry.lastModified === undefined || entry.lastModified <= localGeneratedAt)) { return []; @@ -34,30 +40,6 @@ function remoteModels(entry: ModelsStoreEntry | undefined, localGeneratedAt: num return entry.models; } -function settleOnAbort( - operation: Promise, - signal: AbortSignal | undefined, - canSettle: () => boolean, -): Promise { - if (!signal) return operation; - if (signal.aborted && canSettle()) return Promise.resolve(); - return new Promise((resolve, reject) => { - const abort = () => { - if (canSettle()) resolve(); - }; - signal.addEventListener("abort", abort, { once: true }); - operation.then(resolve, reject).finally(() => signal.removeEventListener("abort", abort)); - }); -} - -type RefreshContext = Parameters>[0]; - -interface InflightRefresh { - promise: Promise; - allowNetwork: boolean; - force: boolean; - signal?: AbortSignal; -} /** Add a persisted pi.dev catalog overlay to a static built-in provider. */ export function withRemoteCatalog( provider: Provider, @@ -65,109 +47,77 @@ export function withRemoteCatalog( localGeneratedAt?: number, ): Provider { let dynamicModels: readonly Model[] = []; - let inflightRefresh: InflightRefresh | undefined; - let refreshEpoch = 0; - let persistenceInProgress = false; - - const refreshModels = (context: RefreshContext): Promise => { - if (inflightRefresh) { - const requiresEscalation = (context.allowNetwork && !inflightRefresh.allowNetwork) - || (context.force === true && !inflightRefresh.force) - || (inflightRefresh.signal?.aborted === true && context.signal?.aborted !== true); - if (requiresEscalation) { - const startEscalation = () => { - if (context.signal?.aborted) return; - return refreshModels(context); - }; - const escalation = inflightRefresh.promise.then(startEscalation, startEscalation); - return settleOnAbort(escalation, context.signal, () => !persistenceInProgress); - } - const owner = inflightRefresh; - const retryAfterAbortedOwner = () => { - if (owner.signal?.aborted && !context.signal?.aborted) return refreshModels(context); - }; - const joined = owner.promise.then(retryAfterAbortedOwner, (error) => { - const retry = retryAfterAbortedOwner(); - if (retry) return retry; - throw error; - }); - return settleOnAbort(joined, context.signal, () => !persistenceInProgress); - } - const epoch = ++refreshEpoch; - const isCurrent = () => epoch === refreshEpoch && !context.signal?.aborted; - const persist = async (entry: Parameters[0]) => { - persistenceInProgress = true; - try { - await context.store.write(entry); - } finally { - persistenceInProgress = false; - } - }; - const operation = (async () => { - const stored = await context.store.read(); - if (!isCurrent()) return; - dynamicModels = remoteModels(stored, localGeneratedAt).filter((model) => model.provider === provider.id); - if (!context.allowNetwork) return; - if ( - !context.force && - stored?.checkedAt !== undefined && - stored.lastModified !== undefined && - Date.now() - stored.checkedAt < REMOTE_CATALOG_REFRESH_INTERVAL_MS - ) return; - - const validator = stored?.models.length ? stored.etag : undefined; - const url = new URL(`/api/models/providers/${encodeURIComponent(provider.id)}`, catalogBaseUrl); - const response = await fetch(url, { - headers: { - accept: "application/json", - "User-Agent": getPiUserAgent(VERSION), - ...(validator ? { "if-none-match": validator } : {}), - }, - signal: context.signal, - }); - if (!isCurrent()) return; - const checkedAt = Date.now(); - if (response.status === 304 && stored) { - await persist({ ...stored, checkedAt }); - return; - } - if (response.status === 404 || response.status === 501) { - await persist({ ...(stored ?? { models: [] }), checkedAt, lastModified: 0, etag: undefined }); - return; - } - if (!response.ok) { - await persist({ ...(stored ?? { models: [] }), checkedAt }); - throw new Error(`Model catalog request failed for ${provider.id}: ${response.status}`); - } - const body: object = await response.json(); - if (!isCurrent()) return; - const refreshed = parseCatalog(provider.id, body); - const lastModified = Date.parse(response.headers.get("last-modified") ?? ""); - const entry = { - models: refreshed, - checkedAt, - lastModified: Number.isNaN(lastModified) ? 0 : lastModified, - etag: response.headers.get("etag") ?? undefined, - }; - await persist(entry); - if (isCurrent()) dynamicModels = remoteModels(entry, localGeneratedAt); - })(); - const refresh = settleOnAbort(operation, context.signal, () => !persistenceInProgress); - const published = refresh.finally(() => { - if (inflightRefresh?.promise === published) inflightRefresh = undefined; - }); - inflightRefresh = { - promise: published, - allowNetwork: context.allowNetwork, - force: context.force === true, - signal: context.signal, - }; - return published; - }; + let inflightRefresh: Promise | undefined; return { ...provider, getModels: () => mergeModels(provider.getModels(), dynamicModels), - refreshModels, + refreshModels: (context) => { + inflightRefresh ??= (async () => { + try { + const stored = await context.store.read(); + dynamicModels = remoteModels(stored, localGeneratedAt).filter((model) => model.provider === provider.id); + if (!context.allowNetwork || context.signal?.aborted) return; + if ( + !context.force && + stored?.checkedAt !== undefined && + stored.lastModified !== undefined && + Date.now() - stored.checkedAt < REMOTE_CATALOG_REFRESH_INTERVAL_MS + ) { + return; + } + + // Only revalidate when a cached body backs the validator, so a 304 can never + // leave the overlay empty. + const validator = stored?.models.length ? stored.etag : undefined; + const url = new URL(`/api/models/providers/${encodeURIComponent(provider.id)}`, catalogBaseUrl); + const response = await fetch(url, { + headers: { + accept: "application/json", + "User-Agent": getPiUserAgent(VERSION), + ...(validator ? { "if-none-match": validator } : {}), + }, + signal: context.signal, + }); + if (context.signal?.aborted) return; + const checkedAt = Date.now(); + // Unchanged: dynamicModels already holds the stored overlay, so only the + // freshness window moves. + if (response.status === 304 && stored) { + await context.store.write({ ...stored, checkedAt }); + return; + } + if (response.status === 404 || response.status === 501) { + await context.store.write({ + ...(stored ?? { models: [] }), + checkedAt, + lastModified: 0, + etag: undefined, + }); + return; + } + if (!response.ok) { + // Transient failure: the cached body and its validator stay valid, so keep the + // etag and let the next refresh revalidate instead of downloading the catalog. + await context.store.write({ ...(stored ?? { models: [] }), checkedAt }); + throw new Error(`Model catalog request failed for ${provider.id}: ${response.status}`); + } + const refreshed = parseCatalog(provider.id, await response.json()); + const lastModified = Date.parse(response.headers.get("last-modified") ?? ""); + if (context.signal?.aborted) return; + const entry = { + models: refreshed, + checkedAt, + lastModified: Number.isNaN(lastModified) ? 0 : lastModified, + etag: response.headers.get("etag") ?? undefined, + }; + dynamicModels = remoteModels(entry, localGeneratedAt); + await context.store.write(entry); + } finally { + inflightRefresh = undefined; + } + })(); + return inflightRefresh; + }, }; } diff --git a/packages/coding-agent/src/core/resolve-config-value.ts b/packages/coding-agent/src/core/resolve-config-value.ts index 119e379d6..db6ce1bd7 100644 --- a/packages/coding-agent/src/core/resolve-config-value.ts +++ b/packages/coding-agent/src/core/resolve-config-value.ts @@ -86,8 +86,8 @@ function parseConfigValueReference(config: string): ConfigValueReference { return { type: "template", parts: parseConfigValueTemplate(config) }; } -function resolveEnvConfigValue(name: string): string | undefined { - return process.env[name] || undefined; +function resolveEnvConfigValue(name: string, env?: Record): string | undefined { + return env?.[name] || process.env[name] || undefined; } function getTemplateEnvVarNames(parts: TemplatePart[]): string[] { @@ -99,14 +99,14 @@ function getTemplateEnvVarNames(parts: TemplatePart[]): string[] { return names; } -function resolveTemplate(parts: TemplatePart[]): string | undefined { +function resolveTemplate(parts: TemplatePart[], env?: Record): string | undefined { let resolved = ""; for (const part of parts) { if (part.type === "literal") { resolved += part.value; continue; } - const envValue = resolveEnvConfigValue(part.name); + const envValue = resolveEnvConfigValue(part.name, env); if (envValue === undefined) return undefined; resolved += envValue; } @@ -124,21 +124,21 @@ export function getConfigValueEnvVarNames(config: string): string[] { return reference.type === "template" ? getTemplateEnvVarNames(reference.parts) : []; } -export function getMissingConfigValueEnvVarNames(config: string): string[] { - return getConfigValueEnvVarNames(config).filter((name) => resolveEnvConfigValue(name) === undefined); +export function getMissingConfigValueEnvVarNames(config: string, env?: Record): string[] { + return getConfigValueEnvVarNames(config).filter((name) => resolveEnvConfigValue(name, env) === undefined); } export function isCommandConfigValue(config: string): boolean { return parseConfigValueReference(config).type === "command"; } -export function isConfigValueConfigured(config: string): boolean { - return getMissingConfigValueEnvVarNames(config).length === 0; -} - +/** Atomic migration compatibility for pre-template bare environment variable names. */ export function isLegacyEnvVarNameConfigValue(config: string): boolean { return LEGACY_ENV_VAR_NAME_RE.test(config); } +export function isConfigValueConfigured(config: string, env?: Record): boolean { + return getMissingConfigValueEnvVarNames(config, env).length === 0; +} /** * Resolve a config value (API key, header value, etc.) to an actual value. @@ -147,21 +147,23 @@ export function isLegacyEnvVarNameConfigValue(config: string): boolean { * - In non-command values, "$$" escapes a literal "$" and "$!" escapes a literal "!" * - Otherwise treats the value as a literal */ -export function resolveConfigValue(config: string): string | undefined { +export function resolveConfigValue(config: string, env?: Record): string | undefined { const reference = parseConfigValueReference(config); if (reference.type === "command") { return executeCommand(reference.config); } - return resolveTemplate(reference.parts); + return resolveTemplate(reference.parts, env); } function executeWithConfiguredShell(command: string): { executed: boolean; value: string | undefined } { try { - const { shell, args } = getShellConfig(); - const result = spawnSync(shell, [...args, command], { + const { shell, args, commandTransport } = getShellConfig(); + const commandFromStdin = commandTransport === "stdin"; + const result = spawnSync(shell, commandFromStdin ? args : [...args, command], { encoding: "utf-8", + input: commandFromStdin ? command : undefined, timeout: 10000, - stdio: ["ignore", "pipe", "ignore"], + stdio: [commandFromStdin ? "pipe" : "ignore", "pipe", "ignore"], shell: false, windowsHide: true, }); @@ -221,16 +223,16 @@ function executeCommand(commandConfig: string): string | undefined { /** * Resolve all header values using the same resolution logic as API keys. */ -export function resolveConfigValueUncached(config: string): string | undefined { +export function resolveConfigValueUncached(config: string, env?: Record): string | undefined { const reference = parseConfigValueReference(config); if (reference.type === "command") { return executeCommandUncached(reference.config); } - return resolveTemplate(reference.parts); + return resolveTemplate(reference.parts, env); } -export function resolveConfigValueOrThrow(config: string, description: string): string { - const resolvedValue = resolveConfigValueUncached(config); +export function resolveConfigValueOrThrow(config: string, description: string, env?: Record): string { + const resolvedValue = resolveConfigValueUncached(config, env); if (resolvedValue !== undefined) { return resolvedValue; } @@ -241,7 +243,7 @@ export function resolveConfigValueOrThrow(config: string, description: string): } if (reference.type === "template") { - const missingEnvVars = getMissingConfigValueEnvVarNames(config); + const missingEnvVars = getMissingConfigValueEnvVarNames(config, env); if (missingEnvVars.length === 1) { throw new Error(`Failed to resolve ${description} from environment variable: ${missingEnvVars[0]}`); } @@ -256,11 +258,14 @@ export function resolveConfigValueOrThrow(config: string, description: string): /** * Resolve all header values using the same resolution logic as API keys. */ -export function resolveHeaders(headers: Record | undefined): Record | undefined { +export function resolveHeaders( + headers: Record | undefined, + env?: Record, +): Record | undefined { if (!headers) return undefined; const resolved: Record = {}; for (const [key, value] of Object.entries(headers)) { - const resolvedValue = resolveConfigValue(value); + const resolvedValue = resolveConfigValue(value, env); if (resolvedValue) { resolved[key] = resolvedValue; } @@ -271,11 +276,12 @@ export function resolveHeaders(headers: Record | undefined): Rec export function resolveHeadersOrThrow( headers: Record | undefined, description: string, + env?: Record, ): Record | undefined { if (!headers) return undefined; const resolved: Record = {}; for (const [key, value] of Object.entries(headers)) { - resolved[key] = resolveConfigValueOrThrow(value, `${description} header "${key}"`); + resolved[key] = resolveConfigValueOrThrow(value, `${description} header "${key}"`, env); } return Object.keys(resolved).length > 0 ? resolved : undefined; } diff --git a/packages/coding-agent/src/core/runtime-credentials.ts b/packages/coding-agent/src/core/runtime-credentials.ts new file mode 100644 index 000000000..52c38ebda --- /dev/null +++ b/packages/coding-agent/src/core/runtime-credentials.ts @@ -0,0 +1,61 @@ +import type { Credential, CredentialInfo, CredentialStore } from "@earendil-works/pi-ai"; + +interface ReloadableCredentialStore { + reload(): void | Promise; +} + +function getReloadableStore(store: CredentialStore): ReloadableCredentialStore | undefined { + const reloadable = store as CredentialStore & Partial; + return typeof reloadable.reload === "function" ? (reloadable as ReloadableCredentialStore) : undefined; +} + +/** Async credential store overlay for non-persistent runtime API keys. */ +export class RuntimeCredentials implements CredentialStore { + private readonly store: CredentialStore; + private readonly overrides = new Map(); + + constructor(store: CredentialStore) { + this.store = store; + } + + setRuntimeApiKey(providerId: string, apiKey: string): void { + this.overrides.set(providerId, apiKey); + } + + removeRuntimeApiKey(providerId: string): void { + this.overrides.delete(providerId); + } + + hasRuntimeApiKey(providerId: string): boolean { + return this.overrides.has(providerId); + } + + async reload(): Promise { + await getReloadableStore(this.store)?.reload(); + } + + async read(providerId: string): Promise { + const override = this.overrides.get(providerId); + return override ? { type: "api_key", key: override } : this.store.read(providerId); + } + + async list(): Promise { + const entries = new Map((await this.store.list()).map((entry) => [entry.providerId, entry])); + for (const providerId of this.overrides.keys()) { + entries.set(providerId, { providerId, type: "api_key" }); + } + return [...entries.values()]; + } + + modify( + providerId: string, + fn: (current: Credential | undefined) => Promise, + ): Promise { + return this.store.modify(providerId, fn); + } + + async delete(providerId: string): Promise { + this.overrides.delete(providerId); + await this.store.delete(providerId); + } +} diff --git a/packages/coding-agent/src/core/sdk-types.ts b/packages/coding-agent/src/core/sdk-types.ts index 55919e533..14f5cb68a 100644 --- a/packages/coding-agent/src/core/sdk-types.ts +++ b/packages/coding-agent/src/core/sdk-types.ts @@ -1,7 +1,5 @@ import type { ThinkingLevel } from "@earendil-works/pi-agent-core"; import type { Api, Model } from "@earendil-works/pi-ai/compat"; -import type { AuthStorage } from "./auth-storage.ts"; -import type { ModelRegistry } from "./model-registry.ts"; import type { ModelRuntime } from "./model-runtime.ts"; import type { ResourceLoader } from "./resource-loader.ts"; import type { SessionManager } from "./session-manager.ts"; @@ -16,12 +14,8 @@ export interface CreateAgentSessionOptions { /** Global config directory. Default: ~/.atomic/agent */ agentDir?: string; - /** Auth storage for credentials. Default: AuthStorage.create(agentDir/auth.json) */ - authStorage?: AuthStorage; - /** Model registry. Default: ModelRegistry.create(authStorage, agentDir/models.json) */ - modelRegistry?: ModelRegistry; - /** Canonical model/auth facade. Takes precedence over modelRegistry and authStorage. */ - modelRuntime?: ModelRuntime; + /** Canonical model/auth runtime. Defaults to agentDir/auth.json and models.json. */ + modelRuntime?: ModelRuntime; /** Model to use. Default: from settings, else first available */ model?: Model; diff --git a/packages/coding-agent/src/core/sdk.ts b/packages/coding-agent/src/core/sdk.ts index 1b5048c4a..974660bc7 100644 --- a/packages/coding-agent/src/core/sdk.ts +++ b/packages/coding-agent/src/core/sdk.ts @@ -5,7 +5,6 @@ import { getAgentDir } from "../config.ts"; import { resolvePath } from "../utils/paths.ts"; import { AgentSession } from "./agent-session.ts"; import { formatNoModelsAvailableMessage } from "./auth-guidance.ts"; -import { AuthStorage } from "./auth-storage.ts"; import { shouldApplyCodexFastMode, streamWithCodexFastMode, @@ -16,7 +15,7 @@ import { restoreAnthropicReplayThinkingBlocks } from "./anthropic-thinking-guard import { DEFAULT_THINKING_LEVEL } from "./defaults.ts"; import type { ExtensionRunner } from "./extensions/index.ts"; import { convertToLlm } from "./messages.ts"; -import { ModelRegistry } from "./model-registry.ts"; +import { ModelRuntime } from "./model-runtime.ts"; import { findInitialModel, resolveRestoredModelReference } from "./model-resolver.ts"; import { DefaultResourceLoader } from "./resource-loader.ts"; import { getDefaultSessionDir, SessionManager } from "./session-manager.ts"; @@ -85,22 +84,10 @@ export async function createAgentSession( const agentDir = options.agentDir ? resolvePath(options.agentDir) : getDefaultAgentDir(); let resourceLoader = options.resourceLoader; - // Use provided or create AuthStorage and ModelRegistry. When a modelRegistry - // is supplied (e.g. a workflow stage reusing one registry across model - // fallback candidates), do NOT also build a fresh AuthStorage: its - // constructor eagerly calls reload(), which acquires the auth.json file lock - // and, under contention, can fail and leave an empty credential set. Reusing - // the supplied registry's already-loaded auth avoids that race (issue #1431). const authPath = options.agentDir ? join(agentDir, "auth.json") : undefined; - const modelsPath = options.agentDir - ? join(agentDir, "models.json") - : undefined; - const modelRegistry = - options.modelRuntime?.modelRegistry ?? - options.modelRegistry ?? - ModelRegistry.create(options.authStorage ?? AuthStorage.create(authPath), modelsPath); - // Restore persisted provider-owned catalogs before any synchronous model reads. - await modelRegistry.refresh({ allowNetwork: false }); + const modelsPath = options.agentDir ? join(agentDir, "models.json") : undefined; + const modelRuntime = options.modelRuntime ?? (await ModelRuntime.create({ authPath, modelsPath })); + await modelRuntime.refresh({ allowNetwork: false }); const settingsManager = options.settingsManager ?? SettingsManager.create(cwd, agentDir); @@ -144,14 +131,11 @@ export async function createAgentSession( // If session has data, try to restore model from it if (!model && hasExistingSession && existingSession.model) { - const restoredModel = await resolveRestoredModelReference( + model = await resolveRestoredModelReference( existingSession.model.provider, existingSession.model.modelId, - modelRegistry, + modelRuntime, ); - if (restoredModel && modelRegistry.hasConfiguredAuth(restoredModel)) { - model = restoredModel; - } if (!model) { modelFallbackMessage = `Could not restore model ${existingSession.model.provider}/${existingSession.model.modelId}`; modelFallbackReason = "session-restore"; @@ -164,7 +148,7 @@ export async function createAgentSession( defaultProvider: settingsManager.getDefaultProvider(), defaultModelId: settingsManager.getDefaultModel(), defaultThinkingLevel: settingsManager.getDefaultThinkingLevel(), - modelRegistry, + modelRuntime, }); model = result.model; if (!model) { @@ -278,10 +262,16 @@ export async function createAgentSession( }, convertToLlm: convertToLlmWithBlockImages, streamFn: async (model, context, streamOptions) => { - const auth = await modelRegistry.getApiKeyAndHeaders(model); - if (!auth.ok) { - throw new Error(auth.error); + const authResult = await modelRuntime.getAuth(model); + const compatibility = authResult ? undefined : modelRuntime.getCompatibilityRequestConfig(model); + if (!authResult && compatibility?.authHeader) { + throw new Error(`No API key found for "${model.provider}"`); } + const auth = { + apiKey: authResult?.auth.apiKey, + headers: authResult?.auth.headers ?? compatibility?.headers, + baseUrl: authResult?.auth.baseUrl, + }; const requestModel = auth.baseUrl !== undefined && auth.baseUrl !== model.baseUrl ? { ...model, baseUrl: auth.baseUrl } : model; @@ -300,6 +290,12 @@ export async function createAgentSession( const attributionHeaders = headerRunner?.hasHandlers("before_provider_headers") ? await headerRunner.emitBeforeProviderHeaders(mergedHeaders ?? {}) : mergedHeaders; const fastModeEnabled = isCodexFastModeEnabled(model); + const extensionProvider = modelRuntime.getRegisteredProviderConfig(requestModel.provider); + const usesExtensionStream = extensionProvider?.streamSimple !== undefined + && requestModel.api === extensionProvider.api; + if (fastModeEnabled && !usesExtensionStream && !authResult) { + throw new Error(`No API key found for "${model.provider}"`); + } const codexFastModeStreamOptions = withCodexFastModeStreamOptions( { ...streamOptions, @@ -313,10 +309,12 @@ export async function createAgentSession( }, fastModeEnabled, ); - if (modelRegistry.hasRegisteredStreamSimpleForApi(requestModel.api)) { - return streamSimple(requestModel, context, codexFastModeStreamOptions); + if (usesExtensionStream) { + return modelRuntime.streamSimple(requestModel, context, codexFastModeStreamOptions); } - return streamWithCodexFastMode(requestModel, context, codexFastModeStreamOptions); + return fastModeEnabled + ? streamWithCodexFastMode(requestModel, context, codexFastModeStreamOptions) + : modelRuntime.streamSimple(requestModel, context, codexFastModeStreamOptions); }, onPayload: async (payload, model) => { const fastModeEnabled = isCodexFastModeEnabled(model); @@ -386,7 +384,7 @@ export async function createAgentSession( fallbackModels: options.fallbackModels ?? settingsManager.getFallbackModels(), resourceLoader, customTools: options.customTools, - modelRegistry, + modelRuntime, initialActiveToolNames, allowedToolNames, excludedToolNames: options.excludedTools, diff --git a/packages/coding-agent/src/extensions/llama/index.ts b/packages/coding-agent/src/extensions/llama/index.ts index f9d09d026..a71037208 100644 --- a/packages/coding-agent/src/extensions/llama/index.ts +++ b/packages/coding-agent/src/extensions/llama/index.ts @@ -41,7 +41,7 @@ async function configuredClient(ctx: ExtensionCommandContext): Promise { const configured = credentialServerUrl(credential) ?? (await ctx.env("LLAMA_BASE_URL"))?.trim(); return configured ? normalizeLlamaServerUrl(configured) : undefined; } -export function toLlamaModel(model: LlamaModelInfo, serverUrl: string): Model<"openai-completions"> { +function toPiModel(model: LlamaModelInfo, serverUrl: string): Model<"openai-completions"> { const reportedContextWindow = model.meta?.n_ctx ?? model.meta?.n_ctx_train; const contextWindow = reportedContextWindow && reportedContextWindow > 0 ? reportedContextWindow : 128000; return { id: model.id, - api: "openai-completions", name: model.id, + api: "openai-completions", provider: LLAMA_PROVIDER_ID, baseUrl: llamaInferenceUrl(serverUrl), reasoning: false, @@ -45,23 +51,25 @@ export function toLlamaModel(model: LlamaModelInfo, serverUrl: string): Model<"o } export interface LlamaProviderController { - config: ProviderConfig; + provider: Provider<"openai-completions">; setCatalog(models: readonly LlamaModelInfo[], serverUrl: string): void; } export function createLlamaProvider(): LlamaProviderController { - let models: Model<"openai-completions">[] = []; + let models: readonly Model<"openai-completions">[] = []; + const setCatalog = (catalog: readonly LlamaModelInfo[], serverUrl: string): void => { - models = catalog.filter((model) => model.status.value === "loaded").map((model) => toLlamaModel(model, serverUrl)); + models = catalog.filter((model) => model.status.value === "loaded").map((model) => toPiModel(model, serverUrl)); }; - const config: ProviderConfig = { + + const provider: Provider<"openai-completions"> = { + id: LLAMA_PROVIDER_ID, name: "llama.cpp", baseUrl: llamaInferenceUrl(DEFAULT_LLAMA_SERVER_URL), - api: "openai-completions", auth: { apiKey: { name: "llama.cpp server", - login: async (interaction) => { + login: async (interaction): Promise => { const enteredUrl = await interaction.prompt({ type: "text", message: "llama.cpp server URL", @@ -70,15 +78,26 @@ export function createLlamaProvider(): LlamaProviderController { const serverUrl = normalizeLlamaServerUrl( enteredUrl.trim() || process.env.LLAMA_BASE_URL || DEFAULT_LLAMA_SERVER_URL, ); - const apiKey = (await interaction.prompt({ type: "secret", message: "API key (optional)" })).trim(); + const apiKey = ( + await interaction.prompt({ + type: "secret", + message: "API key (optional)", + }) + ).trim(); await new LlamaClient(serverUrl, apiKey || undefined).list({ signal: interaction.signal }); - return { type: "api_key", key: apiKey || undefined, env: { LLAMA_BASE_URL: serverUrl } }; + return { + type: "api_key", + key: apiKey || undefined, + env: { LLAMA_BASE_URL: serverUrl }, + }; }, check: async ({ ctx, credential }) => { const serverUrl = await resolveServerUrl(ctx, credential); - return serverUrl ? { type: "api_key", source: credential ? "stored credential" : "LLAMA_BASE_URL" } : undefined; + return serverUrl + ? { type: "api_key", source: credential ? "stored credential" : "LLAMA_BASE_URL" } + : undefined; }, - resolve: async ({ ctx, credential }) => { + resolve: async ({ ctx, credential }): Promise => { const serverUrl = await resolveServerUrl(ctx, credential); if (!serverUrl) return undefined; const apiKey = credential?.key ?? (await ctx.env("LLAMA_API_KEY")) ?? "local"; @@ -90,8 +109,8 @@ export function createLlamaProvider(): LlamaProviderController { }, }, }, - models, - refreshModels: async (context) => { + getModels: () => models, + refreshModels: async (context: RefreshModelsContext): Promise => { const stored = await context.store.read(); if (stored) { models = stored.models.filter( @@ -99,16 +118,17 @@ export function createLlamaProvider(): LlamaProviderController { model.provider === LLAMA_PROVIDER_ID && model.api === "openai-completions", ); } - if (!context.allowNetwork || context.signal?.aborted || context.credential?.type !== "api_key") return models; - const credential = context.credential as ApiKeyCredential; - const serverUrl = credentialServerUrl(credential); - if (!serverUrl) return models; - const catalog = await new LlamaClient(serverUrl, credential.key).list({ signal: context.signal }); + + if (!context.allowNetwork || context.signal?.aborted || context.credential?.type !== "api_key") return; + const serverUrl = credentialServerUrl(context.credential); + if (!serverUrl) return; + const catalog = await new LlamaClient(serverUrl, context.credential.key).list({ signal: context.signal }); setCatalog(catalog, serverUrl); if (!context.signal?.aborted) await context.store.write({ models, checkedAt: Date.now() }); - return models; }, + stream: (model, context, options) => stream(model, context, options as ProviderStreamOptions | undefined), streamSimple: (model, context, options) => streamSimple(model, context, options), }; - return { config, setCatalog }; + + return { provider, setCatalog }; } diff --git a/packages/coding-agent/src/index.ts b/packages/coding-agent/src/index.ts index 7052a86d2..5e6c813eb 100644 --- a/packages/coding-agent/src/index.ts +++ b/packages/coding-agent/src/index.ts @@ -55,28 +55,9 @@ export { parseSkillBlock, type SessionStats, } from "./core/agent-session.ts"; -// Auth and model registry -export { - type ApiKeyCredential, - type AuthCredential, - type AuthStatus, - AuthStorage, - type AuthStorageBackend, - FileAuthStorageBackend, - InMemoryAuthStorageBackend, - type OAuthCredential, - readStoredCredential, -} from "./core/auth-storage.ts"; -import "./core/oauth-compat.js"; -export { - getOAuthApiKey, - getOAuthProvider, - getOAuthProviders, - type LegacyOAuthProvider, - type OAuthProviderDescriptor, - registerOAuthProvider, - resetOAuthProviders, -} from "./core/oauth-provider-bridge.ts"; +// Auth and model runtime +export { AuthStorage, readStoredCredential } from "./core/auth-storage.ts"; +export { type AuthStorageBackend, FileAuthStorageBackend, InMemoryAuthStorageBackend } from "./core/auth-storage-backends.ts"; // Compaction export { type BranchPreparation, diff --git a/packages/coding-agent/src/main-runtime-api-key.ts b/packages/coding-agent/src/main-runtime-api-key.ts new file mode 100644 index 000000000..6b04513b7 --- /dev/null +++ b/packages/coding-agent/src/main-runtime-api-key.ts @@ -0,0 +1,13 @@ +import type { ModelRuntime } from "./core/model-runtime.ts"; + +type CliApiKeyRuntime = Pick; + +/** Apply a CLI-only runtime key without blocking startup on remote catalog I/O. */ +export async function applyCliRuntimeApiKey( + modelRuntime: CliApiKeyRuntime, + providerId: string, + apiKey: string, +): Promise { + await modelRuntime.setRuntimeApiKey(providerId, apiKey, { allowNetwork: false }); + await modelRuntime.getAvailable(); +} diff --git a/packages/coding-agent/src/main-session-options.ts b/packages/coding-agent/src/main-session-options.ts index 2c2020c61..4fd801454 100644 --- a/packages/coding-agent/src/main-session-options.ts +++ b/packages/coding-agent/src/main-session-options.ts @@ -1,7 +1,7 @@ import { modelsAreEqual } from "@earendil-works/pi-ai/compat"; import type { Args } from "./cli/args.ts"; import type { AgentSessionRuntimeDiagnostic } from "./core/agent-session-services.ts"; -import type { ModelRegistry } from "./core/model-registry.ts"; +import type { ModelRuntime } from "./core/model-runtime.ts"; import { resolveCliModel, type ScopedModel } from "./core/model-resolver.ts"; import type { CreateAgentSessionOptions } from "./core/sdk.ts"; import type { SettingsManager } from "./core/settings-manager.ts"; @@ -10,7 +10,7 @@ export function buildSessionOptions( parsed: Args, scopedModels: ScopedModel[], hasExistingSession: boolean, - modelRegistry: ModelRegistry, + modelRuntime: ModelRuntime, settingsManager: SettingsManager, ): { options: CreateAgentSessionOptions; @@ -28,7 +28,7 @@ export function buildSessionOptions( const resolved = resolveCliModel({ cliProvider: parsed.provider, cliModel: parsed.model, - modelRegistry, + modelRuntime, }); if (resolved.warning) { diagnostics.push({ type: "warning", message: resolved.warning }); @@ -51,7 +51,7 @@ export function buildSessionOptions( // Check if saved default is in scoped models - use it if so, otherwise first scoped model const savedProvider = settingsManager.getDefaultProvider(); const savedModelId = settingsManager.getDefaultModel(); - const savedModel = savedProvider && savedModelId ? modelRegistry.find(savedProvider, savedModelId) : undefined; + const savedModel = savedProvider && savedModelId ? modelRuntime.getModel(savedProvider, savedModelId) : undefined; const savedInScope = savedModel ? scopedModels.find((sm) => modelsAreEqual(sm.model, savedModel)) : undefined; if (savedInScope) { diff --git a/packages/coding-agent/src/main.ts b/packages/coding-agent/src/main.ts index 8efcce68c..a612d31f8 100644 --- a/packages/coding-agent/src/main.ts +++ b/packages/coding-agent/src/main.ts @@ -7,7 +7,6 @@ import { ENV_OFFLINE, ENV_SESSION_DIR, ENV_SKIP_VERSION_CHECK, ENV_STARTUP_BENCH import type { CreateAgentSessionRuntimeFactory } from "./core/agent-session-runtime.ts"; import { type AgentSessionRuntimeDiagnostic, createAgentSessionFromServices, createAgentSessionServices } from "./core/agent-session-services.ts"; import { formatNoModelsAvailableMessage } from "./core/auth-guidance.ts"; -import { AuthStorage } from "./core/auth-storage.ts"; import { getBuiltinPackagePaths } from "./core/builtin-packages.ts"; import { applyHttpProxySettings, configureHttpDispatcher } from "./core/http-dispatcher.ts"; import { resolveModelScope, resolveModelScopeWithDiagnostics } from "./core/model-resolver.ts"; @@ -22,6 +21,7 @@ import { runMigrations, showDeprecationWarnings } from "./migrations.ts"; import { builtInExtensions } from "./extensions/index.ts"; import { type AppMode, isPlainRuntimeMetadataCommand, prepareInitialMessage, resolveAppMode, resolveCliPaths, resolveExcludedToolsForAppMode, toPrintOutputMode } from "./main-app-mode.ts"; import { type EarlyInputCapture, startEarlyInputCapture } from "./main-early-input.ts"; +import { applyCliRuntimeApiKey } from "./main-runtime-api-key.ts"; import { computeDeferExtensions, computeStartupInputCaptureEnabled, formatScopedModelList } from "./main-deferred-startup.ts"; import { applyInheritedWorkflowSessionClassification, createSessionManager, promptForMissingSessionCwd, validateForkFlags, validateSessionIdFlags } from "./main-session.ts"; import { buildSessionOptions } from "./main-session-options.ts"; @@ -169,7 +169,6 @@ export async function main(args: string[], options?: MainOptions) { parsed.projectTrustOverride === undefined && !hasProjectTrustInputs(sessionCwd) ? sessionCwd : undefined; const builtinPackagePaths = options?.builtinPackagePaths ?? getBuiltinPackagePaths(); - const authStorage = AuthStorage.create(); const trustPromptMode: AppMode = parsed.help || parsed.listModels !== undefined ? "print" : appMode; const projectTrustByCwd = new Map(); const borrowedExtensionSourceTrustByPath = new Map(); @@ -218,7 +217,6 @@ export async function main(args: string[], options?: MainOptions) { const services = await createAgentSessionServices({ cwd, agentDir, - authStorage, settingsManager: runtimeSettingsManager, extensionFlagValues: parsed.unknownFlags, resourceLoaderReloadOptions: @@ -279,7 +277,7 @@ export async function main(args: string[], options?: MainOptions) { extensionFactories: isolateInteractiveHost ? undefined : extensionFactories, }, }); - const { settingsManager, modelRegistry, resourceLoader } = services; + const { settingsManager, modelRuntime, resourceLoader } = services; const diagnostics: AgentSessionRuntimeDiagnostic[] = [ ...services.diagnostics, ...collectSettingsDiagnostics(settingsManager, "runtime creation"), @@ -292,8 +290,8 @@ export async function main(args: string[], options?: MainOptions) { const scopedModels = modelPatterns && modelPatterns.length > 0 ? deferredExtensionLoad - ? (await resolveModelScopeWithDiagnostics(modelPatterns, modelRegistry)).scopedModels - : await resolveModelScope(modelPatterns, modelRegistry) + ? (await resolveModelScopeWithDiagnostics(modelPatterns, modelRuntime)).scopedModels + : await resolveModelScope(modelPatterns, modelRuntime) : []; const sessionArgs = isolateInteractiveHost ? { ...parsed, provider: undefined, model: undefined, apiKey: undefined, models: undefined } : parsed; @@ -305,7 +303,7 @@ export async function main(args: string[], options?: MainOptions) { sessionArgs, scopedModels, sessionManager.buildSessionContext().messages.length > 0, - modelRegistry, + modelRuntime, settingsManager, ); diagnostics.push(...sessionOptionDiagnostics); @@ -317,7 +315,7 @@ export async function main(args: string[], options?: MainOptions) { message: "--api-key requires a model to be specified via --model, --provider/--model, or --models", }); } else { - authStorage.setRuntimeApiKey(sessionOptions.model.provider, parsed.apiKey); + await applyCliRuntimeApiKey(modelRuntime, sessionOptions.model.provider, parsed.apiKey); } } @@ -350,7 +348,7 @@ export async function main(args: string[], options?: MainOptions) { }); endTimingSpan(runtimeCreationSpan); const { services, session, modelFallbackMessage } = runtime; - const { settingsManager, modelRegistry, resourceLoader } = services; + const { settingsManager, modelRuntime, resourceLoader } = services; applyHttpProxySettings(settingsManager.getGlobalSettings().httpProxy); configureHttpDispatcher(settingsManager.getHttpIdleTimeoutMs()); if (parsed.help) { @@ -369,7 +367,7 @@ export async function main(args: string[], options?: MainOptions) { if (shouldRestoreStdoutForMetadata) { restoreStdout(); } - await listModels(modelRegistry, searchPattern); + await listModels(modelRuntime, searchPattern); process.exit(0); } @@ -421,7 +419,7 @@ export async function main(args: string[], options?: MainOptions) { } if (appMode === "rpc") { - if (!offlineMode) void modelRegistry.refresh().catch(() => {}); + if (!offlineMode) void modelRuntime.refresh().catch(() => {}); printTimings(); await runRpcMode(runtime); } else if (appMode === "interactive") { diff --git a/packages/coding-agent/src/modes/interactive-engine/isolated-auth.ts b/packages/coding-agent/src/modes/interactive-engine/isolated-auth.ts index bcb7a905f..6cada98d5 100644 --- a/packages/coding-agent/src/modes/interactive-engine/isolated-auth.ts +++ b/packages/coding-agent/src/modes/interactive-engine/isolated-auth.ts @@ -2,11 +2,11 @@ import type { AgentSession } from "../../core/agent-session.ts"; import { normalizeOAuthLoginError, type AtomicOAuthLoginCallbacks, -} from "../../core/oauth-provider-bridge.ts"; +} from "../../core/oauth-login.ts"; import type { RpcClient } from "../rpc/rpc-client.ts"; import { loginRpcOAuthProvider } from "../rpc/rpc-oauth-client.ts"; import type { RemoteModelCatalog } from "./remote-model-catalog.ts"; -export type { AtomicOAuthLoginCallbacks } from "../../core/oauth-provider-bridge.ts"; +export type { AtomicOAuthLoginCallbacks } from "../../core/oauth-login.ts"; /** Acquire and persist OAuth entirely in the engine, transporting only UI callbacks and catalog metadata. */ export async function loginIsolatedOAuthProvider( @@ -21,7 +21,8 @@ export async function loginIsolatedOAuthProvider( throw normalizeOAuthLoginError(callbacks.signal?.reason ?? new Error("Login cancelled"), callbacks.signal); } catalog.apply(remoteCatalog); - // Reload only after the engine transaction and catalog refresh both succeed. - session.modelRegistry.authStorage.reload(); + // The engine owns persistence. Refresh the frontend snapshot only after its + // credential transaction and catalog refresh both complete successfully. + await session.modelRuntime.reloadCredentials(); return { modelsRefreshed: true }; } diff --git a/packages/coding-agent/src/modes/interactive-engine/isolated-runtime.ts b/packages/coding-agent/src/modes/interactive-engine/isolated-runtime.ts index e988569ea..d4f2104dc 100644 --- a/packages/coding-agent/src/modes/interactive-engine/isolated-runtime.ts +++ b/packages/coding-agent/src/modes/interactive-engine/isolated-runtime.ts @@ -95,7 +95,8 @@ export class IsolatedInteractiveRuntime extends AgentSessionRuntime { override async logoutProvider(provider: string) { const result = await this.client.logoutProvider(provider); this.remoteModelCatalog.applyModels({ models: result.models, scopedModels: result.scopedModels ?? [] }); - super.session.modelRegistry.authStorage.reload(); + await super.session.modelRuntime.reloadCredentials(); + super.session.refreshCurrentModelFromRegistry(); return result; } @@ -310,7 +311,7 @@ export class IsolatedInteractiveRuntime extends AgentSessionRuntime { configurable: true, value: async (model: Model) => { const selected = await this.client.setModel(model.provider, model.id); - session.agent.state.model = session.modelRegistry.find(selected.provider, selected.id) ?? model; + session.agent.state.model = session.modelRuntime.getModel(selected.provider, selected.id) ?? model; this.resolveModelFallback(); }, }, @@ -327,7 +328,7 @@ export class IsolatedInteractiveRuntime extends AgentSessionRuntime { const previousModel = session.model; const result = await this.client.cycleModel(direction); if (!result) return undefined; - const model = session.modelRegistry.find(result.model.provider, result.model.id) ?? result.model; + const model = session.modelRuntime.getModel(result.model.provider, result.model.id) ?? result.model; session.agent.state.model = model; session.agent.state.thinkingLevel = result.thinkingLevel; this.resolveModelFallbackAfterExplicitModelSelection(previousModel, model); diff --git a/packages/coding-agent/src/modes/interactive-engine/remote-model-catalog.ts b/packages/coding-agent/src/modes/interactive-engine/remote-model-catalog.ts index 3d05c67f6..885b01bc3 100644 --- a/packages/coding-agent/src/modes/interactive-engine/remote-model-catalog.ts +++ b/packages/coding-agent/src/modes/interactive-engine/remote-model-catalog.ts @@ -1,8 +1,6 @@ import type { ModelsRefreshResult } from "@earendil-works/pi-ai"; import type { Api, Model } from "@earendil-works/pi-ai/compat"; import type { AgentSession } from "../../core/agent-session.ts"; -import type { ProviderApiKeyAuth } from "../../core/extensions/provider-types.ts"; -import type { OAuthProviderMetadata } from "../../core/oauth-provider-bridge.ts"; import type { RpcClient } from "../rpc/rpc-client.ts"; import type { RpcModelCatalog } from "../rpc/rpc-types.ts"; @@ -17,8 +15,7 @@ export class RemoteModelCatalog { private readonly client: RpcClient; private models: Model[] = []; private scopedModels: Array<{ model: Model; thinkingLevel?: AgentSession["thinkingLevel"] }> = []; - private customAuthProviders = new Map(); - private oauthProviders: OAuthProviderMetadata[] = []; + private oauthProviders: NonNullable = []; private refreshGeneration = 0; constructor(client: RpcClient) { @@ -27,8 +24,7 @@ export class RemoteModelCatalog { apply(catalog: RpcModelCatalog): void { this.applyModels(catalog); - this.customAuthProviders = new Map(catalog.customAuthProviders.map(({ id, name }) => [id, name])); - this.oauthProviders = catalog.oauthProviders ?? []; + if (catalog.oauthProviders) this.oauthProviders = catalog.oauthProviders; } applyModels(catalog: Pick): void { @@ -37,59 +33,20 @@ export class RemoteModelCatalog { } patch(session: AgentSession): void { - const registry = session.modelRegistry; - const localGetCustomAuth = registry.getCustomApiKeyAuth?.bind(registry) ?? (() => undefined); - const localGetDisplayName = registry.getProviderDisplayName?.bind(registry) ?? ((provider: string) => provider); - const localOAuthProviders = registry.authStorage.getOAuthProviders.bind(registry.authStorage); - Object.defineProperty(registry.authStorage, "getOAuthProviders", { - configurable: true, - value: () => this.oauthProviders.length > 0 ? [...this.oauthProviders] : localOAuthProviders(), - }); - Object.defineProperties(registry, { + const runtime = session.modelRuntime; + Object.defineProperties(runtime, { refresh: { configurable: true, value: (options = {}) => this.refresh(options) }, - getAvailable: { configurable: true, value: () => [...this.models] }, - find: { - configurable: true, - value: (provider: string, modelId: string) => - this.models.find((model) => model.provider === provider && model.id === modelId), - }, + getAvailableSnapshot: { configurable: true, value: () => [...this.models] }, + getModels: { configurable: true, value: (provider?: string) => provider ? this.models.filter((model) => model.provider === provider) : [...this.models] }, + getModel: { configurable: true, value: (provider: string, modelId: string) => this.models.find((model) => model.provider === provider && model.id === modelId) }, hasConfiguredAuth: { configurable: true, - value: (model: Model) => this.models.some( - (candidate) => candidate.provider === model.provider && candidate.id === model.id, - ), - }, - getCustomApiKeyAuthProviders: { - configurable: true, - value: () => [...this.customAuthProviders].map(([id, name]) => ({ id, name })), + value: (provider: string) => runtime.getProviderAuthStatus(provider).configured + || this.models.some((model) => model.provider === provider), }, - getProviderDisplayName: { - configurable: true, - value: (provider: string) => this.oauthProviders.find(({ id }) => id === provider)?.name - ?? this.customAuthProviders.get(provider) - ?? localGetDisplayName(provider), - }, - getCustomApiKeyAuth: { - configurable: true, - value: (provider: string): ProviderApiKeyAuth | undefined => { - const name = this.customAuthProviders.get(provider); - if (!name) return localGetCustomAuth(provider); - return { - name, - login: async ({ signal }) => { - const result = await this.client.loginProvider(provider, signal); - if (result.cancelled) throw new Error("Login cancelled"); - this.apply(result); - return result.credential; - }, - }; - }, - }, - }); - Object.defineProperty(session, "scopedModels", { - configurable: true, - get: () => this.scopedModels, + getOAuthProviderMetadata: { configurable: true, value: () => [...this.oauthProviders] }, }); + Object.defineProperty(session, "scopedModels", { configurable: true, get: () => this.scopedModels }); } private async refresh(options: RemoteModelRefreshOptions = {}): Promise { diff --git a/packages/coding-agent/src/modes/interactive/components/footer.ts b/packages/coding-agent/src/modes/interactive/components/footer.ts index 0f3b5371e..2e2ce210f 100644 --- a/packages/coding-agent/src/modes/interactive/components/footer.ts +++ b/packages/coding-agent/src/modes/interactive/components/footer.ts @@ -114,7 +114,7 @@ function getUsageLine( // Kimi Coding is subscription-backed despite using API-key authentication. const usingSubscription = state.model - ? state.model.provider === "kimi-coding" || session.modelRegistry.isUsingOAuth(state.model) + ? state.model.provider === "kimi-coding" || session.modelRuntime.isUsingOAuth(state.model.provider) : false; if (totals.cost || usingSubscription) { usageParts.push( diff --git a/packages/coding-agent/src/modes/interactive/components/login-dialog.ts b/packages/coding-agent/src/modes/interactive/components/login-dialog.ts index 09f3d4aef..a9feaeeed 100644 --- a/packages/coding-agent/src/modes/interactive/components/login-dialog.ts +++ b/packages/coding-agent/src/modes/interactive/components/login-dialog.ts @@ -1,5 +1,4 @@ import type { AuthInfoLink, OAuthDeviceCodeInfo } from "@earendil-works/pi-ai"; -import { getOAuthProviderDescriptors } from "../../../core/oauth-provider-bridge.ts"; import { Container, type Focusable, getKeybindings, Input, Spacer, Text, type TUI } from "@earendil-works/pi-tui"; import { openBrowser } from "../../../utils/open-browser.ts"; import { theme } from "../theme/theme.ts"; @@ -41,8 +40,7 @@ export class LoginDialogComponent extends Container implements Focusable { this.onComplete = onComplete; this.tui = tui; - const providerInfo = getOAuthProviderDescriptors().find((p) => p.id === providerId); - const providerName = providerNameOverride || providerInfo?.name || providerId; + const providerName = providerNameOverride || providerId; const title = titleOverride ?? `Login to ${providerName}`; // Top border diff --git a/packages/coding-agent/src/modes/interactive/components/model-selector.ts b/packages/coding-agent/src/modes/interactive/components/model-selector.ts index 8d5fdbd58..c7a0e5bd8 100644 --- a/packages/coding-agent/src/modes/interactive/components/model-selector.ts +++ b/packages/coding-agent/src/modes/interactive/components/model-selector.ts @@ -9,8 +9,7 @@ import { Text, type TUI, } from "@earendil-works/pi-tui"; -import type { ModelRegistry } from "../../../core/model-registry.ts"; -import { isOfflineModeEnabled } from "../../../core/package-manager-env.ts"; +import type { ModelRuntime } from "../../../core/model-runtime.ts"; import type { SettingsManager } from "../../../core/settings-manager.ts"; import { getModelSelectorSearchText } from "../model-search.ts"; import { theme } from "../theme/theme.ts"; @@ -53,7 +52,7 @@ export class ModelSelectorComponent extends Container implements Focusable { private selectedIndex: number = 0; private currentModel?: Model; private settingsManager: SettingsManager; - private modelRegistry: ModelRegistry; + private modelRuntime: ModelRuntime; private onSelectCallback: (model: Model) => void; private onCancelCallback: () => void; private errorMessage?: string; @@ -65,13 +64,14 @@ export class ModelSelectorComponent extends Container implements Focusable { private scopeText?: Text; private scopeHintText?: Text; private readonly refreshAbortController = new AbortController(); + private refreshTimeout?: ReturnType; private closed = false; constructor( tui: TUI, currentModel: Model | undefined, settingsManager: SettingsManager, - modelRegistry: ModelRegistry, + modelRuntime: ModelRuntime, scopedModels: ReadonlyArray, onSelect: (model: Model) => void, onCancel: () => void, @@ -82,7 +82,7 @@ export class ModelSelectorComponent extends Container implements Focusable { this.tui = tui; this.currentModel = currentModel; this.settingsManager = settingsManager; - this.modelRegistry = modelRegistry; + this.modelRuntime = modelRuntime; this.scopedModels = scopedModels; this.scope = scopedModels.length > 0 ? "scoped" : "all"; this.onSelectCallback = onSelect; @@ -128,52 +128,23 @@ export class ModelSelectorComponent extends Container implements Focusable { // Add bottom border this.addChild(new DynamicBorder()); - // Show the current snapshot first, then refresh configured provider catalogs in the background. - void this.loadModelsFromSnapshot() - .then(() => { - if (initialSearchInput) this.filterModels(initialSearchInput); - else this.updateList(); - this.tui.requestRender(); - return this.refreshModels(); - }) - .catch((error) => { - if (this.closed) return; - this.refreshStatusSuccess = false; - this.refreshStatusMessage = `Could not refresh model catalogs: ${error instanceof Error ? error.message : String(error)}`; - this.tui.requestRender(); - }); + // Render the current snapshot immediately, then refresh in the background. + this.loadModelsFromSnapshot(); + if (initialSearchInput) this.filterModels(initialSearchInput); + else this.updateList(); + this.tui.requestRender(); + void this.refreshModels(); } - private async loadModelsFromSnapshot(): Promise { - this.errorMessage = undefined; - let models: ModelItem[]; - - // Check for models.json errors - const loadError = this.modelRegistry.getError(); - if (loadError) { - this.errorMessage = loadError; - } - - // Load available models (built-in models still work even if models.json failed) - try { - const availableModels = await this.modelRegistry.getAvailable(); - models = availableModels.map((model: Model) => ({ - provider: model.provider, - id: model.id, - model, - })); - } catch (error) { - this.allModels = []; - this.scopedModelItems = []; - this.activeModels = []; - this.filteredModels = []; - this.errorMessage = error instanceof Error ? error.message : String(error); - return; - } - + private loadModelsFromSnapshot(): void { + const models = this.modelRuntime.getAvailableSnapshot().map((model: Model) => ({ + provider: model.provider, + id: model.id, + model, + })); this.allModels = this.sortModels(models); this.scopedModels = this.scopedModels.map((scoped) => { - const refreshed = this.modelRegistry.find(scoped.model.provider, scoped.model.id); + const refreshed = this.modelRuntime.getModel(scoped.model.provider, scoped.model.id); return refreshed ? { ...scoped, model: refreshed } : scoped; }); this.scopedModelItems = this.scopedModels.map((scoped) => ({ @@ -189,30 +160,40 @@ export class ModelSelectorComponent extends Container implements Focusable { } private async refreshModels(): Promise { - const result = await this.modelRegistry.refresh({ - allowNetwork: !isOfflineModeEnabled(), - signal: this.refreshAbortController.signal, - timeoutMs: 15_000, - }); - if (this.closed) return; - this.refreshStatusSuccess = false; - if (result.aborted) { - this.refreshStatusMessage = "Model refresh timed out; showing cached models."; - } else if (result.errors.size === 1) { - this.refreshStatusMessage = `Could not refresh ${result.errors.keys().next().value}; showing available models.`; - } else if (result.errors.size > 1) { - this.refreshStatusMessage = `Could not refresh ${result.errors.size} model catalogs; showing available models.`; - } else { - this.refreshStatusMessage = "Model catalogs refreshed."; - this.refreshStatusSuccess = true; + const timeoutMs = 15_000; + let timedOut = false; + this.refreshTimeout = setTimeout(() => { + timedOut = true; + this.refreshAbortController.abort(); + }, timeoutMs); + try { + const result = await this.modelRuntime.refresh({ signal: this.refreshAbortController.signal }); + if (this.closed) return; + this.refreshStatusMessage = ""; + if (result.aborted && timedOut) { + this.errorMessage = "Model refresh timed out; showing cached models."; + } else if (result.errors.size === 1) { + this.errorMessage = `Could not refresh ${result.errors.keys().next().value}; showing cached models.`; + } else if (result.errors.size > 1) { + this.errorMessage = `Could not refresh ${result.errors.size} model catalogs; showing cached models.`; + } else { + this.errorMessage = this.modelRuntime.getError(); + if (!this.errorMessage) { + this.refreshStatusMessage = "Model catalogs refreshed."; + this.refreshStatusSuccess = true; + } + } + this.loadModelsFromSnapshot(); + this.filterModels(this.searchInput.getValue()); + this.tui.requestRender(); + } finally { + if (this.refreshTimeout) clearTimeout(this.refreshTimeout); } - await this.loadModelsFromSnapshot(); - this.filterModels(this.searchInput.getValue()); - this.tui.requestRender(); } private close(): void { this.closed = true; + if (this.refreshTimeout) clearTimeout(this.refreshTimeout); this.refreshAbortController.abort(); } @@ -302,12 +283,6 @@ export class ModelSelectorComponent extends Container implements Focusable { this.listContainer.addChild(new Text(scrollInfo, 0, 0)); } - if (this.refreshStatusMessage) { - this.listContainer.addChild(new Spacer(1)); - const color = this.refreshStatusSuccess ? "success" : "warning"; - this.listContainer.addChild(new Text(theme.fg(color, ` ${this.refreshStatusMessage}`), 0, 0)); - } - // Show error message or "no results" if empty if (this.errorMessage) { // Show error in red @@ -322,6 +297,12 @@ export class ModelSelectorComponent extends Container implements Focusable { this.listContainer.addChild(new Spacer(1)); this.listContainer.addChild(new Text(theme.fg("muted", ` Model Name: ${selected.model.name}`), 0, 0)); } + if (this.refreshStatusMessage) { + this.listContainer.addChild(new Spacer(1)); + this.listContainer.addChild( + new Text(theme.fg(this.refreshStatusSuccess ? "success" : "muted", ` ${this.refreshStatusMessage}`), 0, 0), + ); + } } handleInput(keyData: string): void { diff --git a/packages/coding-agent/src/modes/interactive/components/oauth-selector.ts b/packages/coding-agent/src/modes/interactive/components/oauth-selector.ts index 72669dd2a..b54aab7f8 100644 --- a/packages/coding-agent/src/modes/interactive/components/oauth-selector.ts +++ b/packages/coding-agent/src/modes/interactive/components/oauth-selector.ts @@ -7,7 +7,8 @@ import { Spacer, TruncatedText, } from "@earendil-works/pi-tui"; -import type { AuthStatus, AuthStorage } from "../../../core/auth-storage.ts"; +import type { AuthStatus } from "../../../core/provider-composer.ts"; +import type { ModelRuntime } from "../../../core/model-runtime.ts"; import { theme } from "../theme/theme.ts"; import { DynamicBorder } from "./dynamic-border.ts"; @@ -38,14 +39,14 @@ export class OAuthSelectorComponent extends Container implements Focusable { private filteredProviders: AuthSelectorProvider[]; private selectedIndex: number = 0; private mode: "login" | "logout"; - private authStorage: AuthStorage; + private modelRuntime: ModelRuntime; private getAuthStatus: (providerId: string) => AuthStatus; private onSelectCallback: (providerId: string, authType: AuthSelectorProvider["authType"]) => void; private onCancelCallback: () => void; constructor( mode: "login" | "logout", - authStorage: AuthStorage, + modelRuntime: ModelRuntime, providers: AuthSelectorProvider[], onSelect: (providerId: string, authType: AuthSelectorProvider["authType"]) => void, onCancel: () => void, @@ -55,8 +56,8 @@ export class OAuthSelectorComponent extends Container implements Focusable { super(); this.mode = mode; - this.authStorage = authStorage; - this.getAuthStatus = getAuthStatus ?? ((providerId) => this.authStorage.getAuthStatus(providerId)); + this.modelRuntime = modelRuntime; + this.getAuthStatus = getAuthStatus ?? ((providerId) => this.modelRuntime.getProviderAuthStatus(providerId)); this.allProviders = providers; this.filteredProviders = providers; this.onSelectCallback = onSelect; @@ -156,28 +157,19 @@ export class OAuthSelectorComponent extends Container implements Focusable { } private formatStatusIndicator(provider: AuthSelectorProvider): string { - const credential = this.authStorage.get(provider.id); - if (credential?.type === provider.authType) return theme.fg("success", " ✓ configured"); - if (credential) { - const label = credential.type === "oauth" ? "subscription configured" : "API key configured"; - return theme.fg("muted", " • ") + theme.fg("warning", label); - } - if (provider.authType !== "api_key") return theme.fg("muted", " • unconfigured"); - const status = this.getAuthStatus(provider.id); + if (!status.configured) return theme.fg("muted", " • unconfigured"); switch (status.source) { - case "environment": - return theme.fg("success", ` ✓ env: ${status.label ?? "API key"}`); - case "runtime": - return theme.fg("success", " ✓ runtime API key"); - case "fallback": - return theme.fg("success", " ✓ custom API key"); - case "models_json_key": - return theme.fg("success", " ✓ key in models.json"); - case "models_json_command": - return theme.fg("success", " ✓ command in models.json"); - default: - return theme.fg("muted", " • unconfigured"); + case "environment": return theme.fg("success", ` ✓ env: ${status.label ?? "API key"}`); + case "runtime": return theme.fg("success", " ✓ runtime API key"); + case "stored": { + const storedType = this.modelRuntime.getStoredCredentialType(provider.id); + return theme.fg("success", ` ✓ ${storedType === "oauth" ? "subscription" : "API key"} configured`); + } + case "fallback": return theme.fg("success", " ✓ configured"); + case "models_json_key": return theme.fg("success", " ✓ key in models.json"); + case "models_json_command": return theme.fg("success", " ✓ command in models.json"); + default: return theme.fg("muted", " • unconfigured"); } } diff --git a/packages/coding-agent/src/modes/interactive/interactive-agent-events.ts b/packages/coding-agent/src/modes/interactive/interactive-agent-events.ts index 92681b870..3d5361c55 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-agent-events.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-agent-events.ts @@ -222,7 +222,7 @@ InteractiveModeBase.prototype.handleEvent = async function(this: InteractiveMode this.footer.invalidate(); } if (event.message.role === "assistant" && this.settingsManager.getShowCacheMissNotices()) { - const miss = detectCacheMiss(this.sessionManager.getEntries(), event.message, { getModel: (provider, model) => this.session.modelRegistry.find(provider, model) }); + const miss = detectCacheMiss(this.sessionManager.getEntries(), event.message, { getModel: (provider, model) => this.session.modelRuntime.getModel(provider, model) }); if (miss) { const cause = miss.modelChanged ? " after model switch" : miss.idleMs >= CACHE_TTL_MS ? " after cache TTL expiry" : ""; this.chatContainer.addChild(new Text(theme.fg("warning", `Prompt cache miss${cause}: ${miss.missedTokens.toLocaleString()} tokens re-billed ($${miss.missedCost.toFixed(3)})`), 1, 0)); diff --git a/packages/coding-agent/src/modes/interactive/interactive-auth-login.ts b/packages/coding-agent/src/modes/interactive/interactive-auth-login.ts index 10317d94f..f4f8bfe51 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-auth-login.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-auth-login.ts @@ -1,13 +1,13 @@ import { InteractiveModeBase } from "./interactive-mode-base.ts"; import { type Api, type Model, type OAuthSelectPrompt, path, getAuthPath, getDocsPath, defaultModelPerProvider, ExtensionSelectorComponent, LoginDialogComponent, theme } from "./interactive-mode-deps.ts"; import { hasDefaultModelProvider, isUnknownModel } from "./interactive-mode-helpers.ts"; -import { isOAuthLoginCancelled } from "../../core/oauth-provider-bridge.ts"; +import { isOAuthLoginCancelled } from "../../core/oauth-login.ts"; InteractiveModeBase.prototype.completeProviderAuthentication = async function(this: InteractiveModeBase, providerId: string, providerName: string, authType: "oauth" | "api_key", previousModel: Model | undefined, options: { modelsRefreshed?: boolean } = {}): Promise { if (!options.modelsRefreshed) { - // Upstream pi parity: a failed or timeout-aborted catalog refresh after a - // completed login is non-fatal; models fall back to the cached snapshot. - await this.session.modelRegistry.refresh(); + // Match pi: after authentication persists the credential, a thrown catalog + // refresh failure remains visible to the caller without rolling it back. + await this.session.modelRuntime.refresh(); } const actionLabel = @@ -18,7 +18,7 @@ InteractiveModeBase.prototype.completeProviderAuthentication = async function(th let selectedModel: Model | undefined; let selectionError: string | undefined; if (isUnknownModel(previousModel)) { - const availableModels = this.session.modelRegistry.getAvailable(); + const availableModels = this.session.modelRuntime.getAvailableSnapshot(); const providerModels = availableModels.filter( (model) => model.provider === providerId, ); @@ -126,26 +126,16 @@ InteractiveModeBase.prototype.showApiKeyLoginDialog = async function(this: Inter }; try { - const customAuth = this.session.modelRegistry.getCustomApiKeyAuth(providerId); - if (customAuth) { - const credential = await customAuth.login({ - signal: dialog.signal, - prompt: (prompt) => dialog.showPrompt(prompt.message, prompt.placeholder), - }); - this.session.modelRegistry.authStorage.set(providerId, credential); - } else { - const apiKey = (await dialog.showPrompt("Enter API key:")).trim(); - if (!apiKey) throw new Error("API key cannot be empty."); - this.session.modelRegistry.authStorage.set(providerId, { type: "api_key", key: apiKey }); - } - + await this.session.modelRuntime.login(providerId, "api_key", { + signal: dialog.signal, + prompt: (prompt) => dialog.showPrompt(prompt.message, "placeholder" in prompt ? prompt.placeholder : undefined), + notify: (event) => { + if (event.type === "info") dialog.showInfo(event.message, event.links); + else if (event.type === "progress") dialog.showProgress(event.message); + }, + }); restoreEditor(); - await this.completeProviderAuthentication( - providerId, - providerName, - "api_key", - previousModel, - ); + await this.completeProviderAuthentication(providerId, providerName, "api_key", previousModel, { modelsRefreshed: true }); } catch (error: unknown) { restoreEditor(); const errorMsg = error instanceof Error ? error.message : String(error); @@ -186,15 +176,10 @@ InteractiveModeBase.prototype.showOAuthLoginSelect = function(this: InteractiveM }; InteractiveModeBase.prototype.showLoginDialog = async function(this: InteractiveModeBase, providerId: string, providerName: string): Promise { - const providerInfo = this.session.modelRegistry.authStorage - .getOAuthProviders() - .find((provider) => provider.id === providerId); const previousModel = this.session.model; + const metadata = this.session.modelRuntime?.getOAuthProviderMetadata().find(({ id }) => id === providerId); + const usesCallbackServer = metadata?.usesCallbackServer === true; - // Providers that use callback servers (can paste redirect URL) - const usesCallbackServer = providerInfo?.usesCallbackServer ?? false; - - // Create login dialog component const dialog = new LoginDialogComponent( this.ui, providerId, @@ -202,6 +187,7 @@ InteractiveModeBase.prototype.showLoginDialog = async function(this: Interactive // Completion handled below }, providerName, + metadata?.loginLabel, ); // Show dialog in editor container @@ -267,7 +253,7 @@ InteractiveModeBase.prototype.showLoginDialog = async function(this: Interactive message: string; placeholder?: string; }) => { - return dialog.showPrompt(prompt.message, prompt.placeholder); + return dialog.showPrompt(prompt.message, "placeholder" in prompt ? prompt.placeholder : undefined); }, onProgress: (message: string) => { diff --git a/packages/coding-agent/src/modes/interactive/interactive-auth-routing.ts b/packages/coding-agent/src/modes/interactive/interactive-auth-routing.ts index 0ad3fdc3b..1a776a753 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-auth-routing.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-auth-routing.ts @@ -1,5 +1,5 @@ import { builtinProviders } from "@earendil-works/pi-ai/providers/all"; -import type { AuthStatus } from "../../core/auth-storage.ts"; +import type { AuthStatus } from "../../core/provider-composer.ts"; import { InteractiveModeBase } from "./interactive-mode-base.ts"; import { type AuthSelectorProvider, @@ -38,64 +38,27 @@ InteractiveModeBase.prototype.getLoginProviderOptions = function( this: InteractiveModeBase, authType?: "oauth" | "api_key", ): AuthSelectorProvider[] { - const authStorage = this.session.modelRegistry.authStorage; - const oauthProviders = authStorage.getOAuthProviders(); - const options: AuthSelectorProvider[] = oauthProviders.map((provider) => ({ - id: provider.id, - name: this.session.modelRegistry.getProviderDisplayName(provider.id), - authType: "oauth", - })); - - const builtins = builtinProviders(); - const builtinIds = new Set(builtins.map((provider) => provider.id)); - options.push(...getBuiltinApiKeyLoginOptions( - (providerId) => this.session.modelRegistry.getProviderDisplayName(providerId), - )); - const customApiKeyProviders = this.session.modelRegistry.getCustomApiKeyAuthProviders(); - const customApiKeyProviderIds = new Set(customApiKeyProviders.map((provider) => provider.id)); - options.push(...customApiKeyProviders.map((provider) => ({ - ...provider, - authType: "api_key" as const, - }))); - - // Legacy extension/config providers do not expose pi-ai auth metadata. Keep - // Atomic's existing API-key behavior for model-backed, non-OAuth providers. - const oauthProviderIds = new Set(oauthProviders.map((provider) => provider.id)); - const modelProviderIds = new Set( - this.session.modelRegistry.getAll().map((model) => model.provider), - ); - for (const providerId of modelProviderIds) { - if (builtinIds.has(providerId) || oauthProviderIds.has(providerId) || customApiKeyProviderIds.has(providerId)) continue; - options.push({ - id: providerId, - name: this.session.modelRegistry.getProviderDisplayName(providerId), - authType: "api_key", - }); + const options: AuthSelectorProvider[] = this.session.modelRuntime + .getOAuthProviderMetadata() + .map((provider) => ({ id: provider.id, name: provider.name, authType: "oauth" as const })); + for (const provider of this.session.modelRuntime.getProviders()) { + if (provider.auth.apiKey) options.push({ id: provider.id, name: provider.name ?? provider.id, authType: "api_key" }); } - - const filtered = authType - ? options.filter((option) => option.authType === authType) - : options; - return filtered.sort((a, b) => a.name.localeCompare(b.name)); + return (authType ? options.filter((option) => option.authType === authType) : options) + .sort((a, b) => a.name.localeCompare(b.name)); }; - InteractiveModeBase.prototype.getLogoutProviderOptions = function( this: InteractiveModeBase, ): AuthSelectorProvider[] { - const authStorage = this.session.modelRegistry.authStorage; - const supportedProviderIds = new Set( - this.getLoginProviderOptions().map((provider) => provider.id), - ); + const runtime = this.session.modelRuntime; + const providers = runtime.getProviders(); + const providerNames = new Map(providers.map((provider) => [provider.id, provider.name ?? provider.id])); + for (const provider of runtime.getOAuthProviderMetadata()) providerNames.set(provider.id, provider.name); const options: AuthSelectorProvider[] = []; - for (const providerId of authStorage.list()) { - if (!supportedProviderIds.has(providerId)) continue; - const credential = authStorage.get(providerId); - if (!credential) continue; - options.push({ - id: providerId, - name: this.session.modelRegistry.getProviderDisplayName(providerId), - authType: credential.type, - }); + for (const [providerId, name] of providerNames) { + if (runtime.getProviderAuthStatus(providerId).source !== "stored") continue; + const authType = runtime.getStoredCredentialType(providerId); + if (authType) options.push({ id: providerId, name, authType }); } return options.sort((a, b) => a.name.localeCompare(b.name)); }; @@ -186,7 +149,7 @@ InteractiveModeBase.prototype.showLoginProviderSelector = function( this.showSelector((done) => { const selector = new OAuthSelectorComponent( "login", - this.session.modelRegistry.authStorage, + this.session.modelRuntime, providerOptions, async (providerId, selectedAuthType) => { done(); @@ -201,7 +164,7 @@ InteractiveModeBase.prototype.showLoginProviderSelector = function( if (authType) this.showLoginAuthTypeSelector(); else this.ui.requestRender(); }, - (providerId) => this.session.modelRegistry.getProviderAuthStatus(providerId), + (providerId) => this.session.modelRuntime.getProviderAuthStatus(providerId), initialSearchInput, ); return { component: selector, focus: selector }; @@ -226,7 +189,7 @@ InteractiveModeBase.prototype.showOAuthSelector = async function( this.showSelector((done) => { const selector = new OAuthSelectorComponent( mode, - this.session.modelRegistry.authStorage, + this.session.modelRuntime, providerOptions, async (providerId, selectedAuthType) => { done(); diff --git a/packages/coding-agent/src/modes/interactive/interactive-autocomplete.ts b/packages/coding-agent/src/modes/interactive/interactive-autocomplete.ts index 6a4a27685..6c7cc51af 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-autocomplete.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-autocomplete.ts @@ -164,10 +164,10 @@ InteractiveModeBase.prototype.getCodexFastModeCandidateModels = function(this: I if (this.session.scopedModels.length > 0) { return this.session.scopedModels .map((scoped) => scoped.model) - .filter((model) => this.session.modelRegistry.hasConfiguredAuth(model)); + .filter((model) => this.session.modelRuntime.hasConfiguredAuth(model.provider)); } - return this.session.modelRegistry.getAvailable(); + return [...this.session.modelRuntime.getAvailableSnapshot()]; }; InteractiveModeBase.prototype.hasCodexFastModeSupportedModels = function(this: InteractiveModeBase): boolean { @@ -240,7 +240,7 @@ InteractiveModeBase.prototype.createBaseAutocompleteProvider = function(this: In const models = this.session.scopedModels.length > 0 ? this.session.scopedModels.map((s) => s.model) - : this.session.modelRegistry.getAvailable(); + : this.session.modelRuntime.getAvailableSnapshot(); if (models.length === 0) return null; diff --git a/packages/coding-agent/src/modes/interactive/interactive-deferred-startup.ts b/packages/coding-agent/src/modes/interactive/interactive-deferred-startup.ts index c1d76562e..563699aa9 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-deferred-startup.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-deferred-startup.ts @@ -65,7 +65,7 @@ InteractiveModeBase.prototype.completeDeferredStartup = async function(this: Int // Keep the subscription warning after the RESOURCES disclosure. void this.maybeWarnAboutAnthropicSubscriptionAuth(undefined, this.startupNoticesContainer); this.showStartupNoticesIfNeeded(this.startupNoticesContainer); - const modelsJsonError = this.session.modelRegistry.getError(); + const modelsJsonError = this.session.modelRuntime.getError(); if (modelsJsonError) { this.showError(`models.json error: ${modelsJsonError}`); } @@ -97,7 +97,7 @@ export async function applyDeferredModelScope(mode: InteractiveModeBase): Promis return; } - const { scopedModels, diagnostics } = await resolveModelScopeWithDiagnostics(patterns, mode.session.modelRegistry); + const { scopedModels, diagnostics } = await resolveModelScopeWithDiagnostics(patterns, mode.session.modelRuntime); for (const diagnostic of diagnostics) { mode.showWarning(diagnostic.message); } @@ -114,10 +114,10 @@ export async function applyDeferredModelScope(mode: InteractiveModeBase): Promis const savedProvider = mode.settingsManager.getDefaultProvider(); const savedModelId = mode.settingsManager.getDefaultModel(); - const savedModel = savedProvider && savedModelId ? mode.session.modelRegistry.find(savedProvider, savedModelId) : undefined; + const savedModel = savedProvider && savedModelId ? mode.session.modelRuntime.getModel(savedProvider, savedModelId) : undefined; const preferred = savedModel ? scopedModels.find((scoped) => modelsAreEqual(scoped.model, savedModel)) : undefined; const nextScopedModel = preferred ?? scopedModels[0]; - if (mode.session.modelRegistry.hasConfiguredAuth(nextScopedModel.model)) { + if (mode.session.modelRuntime.hasConfiguredAuth(nextScopedModel.model.provider)) { await mode.session.setModel(nextScopedModel.model); if (nextScopedModel.thinkingLevel && !mode.options.deferredModelScopePreserveThinking) { mode.session.setThinkingLevel(nextScopedModel.thinkingLevel); @@ -135,16 +135,16 @@ InteractiveModeBase.prototype.retryDeferredModelRestore = async function(this: I if (!preliminaryFallbackMessage) return; const savedModel = this.sessionManager.buildSessionContext().model; - if (!savedModel && this.session.model && this.session.modelRegistry.hasConfiguredAuth(this.session.model)) { + if (!savedModel && this.session.model && this.session.modelRuntime.hasConfiguredAuth(this.session.model.provider)) { return; } if (savedModel) { const restoredModel = await resolveRestoredModelReference( savedModel.provider, savedModel.modelId, - this.session.modelRegistry, + this.session.modelRuntime, ); - if (restoredModel && this.session.modelRegistry.hasConfiguredAuth(restoredModel)) { + if (restoredModel && this.session.modelRuntime.hasConfiguredAuth(restoredModel.provider)) { await this.session.setModel(restoredModel); return; } @@ -161,7 +161,7 @@ InteractiveModeBase.prototype.retryDeferredModelRestore = async function(this: I defaultProvider, defaultModelId, defaultThinkingLevel: this.settingsManager.getDefaultThinkingLevel(), - modelRegistry: this.session.modelRegistry, + modelRuntime: this.session.modelRuntime, }); if (result.model) { await this.session.setModel(result.model); diff --git a/packages/coding-agent/src/modes/interactive/interactive-extension-runtime.ts b/packages/coding-agent/src/modes/interactive/interactive-extension-runtime.ts index 4cd208015..75e50baf3 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-extension-runtime.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-extension-runtime.ts @@ -1,3 +1,4 @@ +import { ModelRegistry } from "../../core/model-registry.ts"; import { InteractiveModeBase } from "./interactive-mode-base.ts"; import { type KeyId, type Component, type LoaderIndicatorOptions, type ExtensionContext, type ExtensionRunner, type ExtensionWidgetOptions, Container, matchesKey, Text, TUI, AssistantMessageComponent, keyText, Theme, theme } from "./interactive-mode-deps.ts"; import { mountIdleStatus } from "./components/idle-status.ts"; @@ -16,7 +17,7 @@ InteractiveModeBase.prototype.setupExtensionShortcuts = function(this: Interacti hasUI: true, cwd: this.sessionManager.getCwd(), sessionManager: this.sessionManager, - modelRegistry: this.session.modelRegistry, + modelRegistry: new ModelRegistry(this.session.modelRuntime), model: this.session.model, isIdle: () => !this.session.isStreaming, isProjectTrusted: () => this.session.settingsManager.isProjectTrusted(), diff --git a/packages/coding-agent/src/modes/interactive/interactive-model-catalog-startup.ts b/packages/coding-agent/src/modes/interactive/interactive-model-catalog-startup.ts index e46d19609..4614e340a 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-model-catalog-startup.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-model-catalog-startup.ts @@ -4,7 +4,7 @@ import type { InteractiveModeBase } from "./interactive-mode-base.ts"; export function updateProviderCountFromSnapshot(mode: InteractiveModeBase): void { const models = mode.session.scopedModels.length > 0 ? mode.session.scopedModels.map((scoped) => scoped.model) - : mode.session.modelRegistry.getAvailable(); + : mode.session.modelRuntime.getAvailableSnapshot(); mode.footerDataProvider.setAvailableProviderCount(new Set(models.map((model) => model.provider)).size); } @@ -16,7 +16,7 @@ export function updateProviderCountFromSnapshot(mode: InteractiveModeBase): void * (the caller already gates on offline mode). */ export function refreshCatalogsAfterTuiStartup(mode: InteractiveModeBase): Promise { - return mode.session.modelRegistry.refresh({ allowNetwork: true }) + return mode.session.modelRuntime.refresh({ allowNetwork: true }) .catch(() => {}) .then(() => updateProviderCountFromSnapshot(mode)) .catch(() => {}); diff --git a/packages/coding-agent/src/modes/interactive/interactive-model-routing.ts b/packages/coding-agent/src/modes/interactive/interactive-model-routing.ts index 9f87e414f..7eff1eb8d 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-model-routing.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-model-routing.ts @@ -38,9 +38,9 @@ InteractiveModeBase.prototype.getModelCandidates = async function(this: Interact } const allowNetwork = !isOfflineModeEnabled(); - await this.session.modelRegistry.refresh({ allowNetwork }); + await this.session.modelRuntime.refresh({ allowNetwork }); try { - return await this.session.modelRegistry.getAvailable(); + return [...this.session.modelRuntime.getAvailableSnapshot()]; } catch { return []; } @@ -49,7 +49,7 @@ InteractiveModeBase.prototype.getModelCandidates = async function(this: Interact InteractiveModeBase.prototype.updateAvailableProviderCount = async function(this: InteractiveModeBase): Promise { const models = this.session.scopedModels.length > 0 ? this.session.scopedModels.map((scoped) => scoped.model) - : this.session.modelRegistry.getAvailable(); + : this.session.modelRuntime.getAvailableSnapshot(); this.footerDataProvider.setAvailableProviderCount(new Set(models.map((model) => model.provider)).size); }; @@ -64,18 +64,15 @@ InteractiveModeBase.prototype.maybeWarnAboutAnthropicSubscriptionAuth = async fu return; } - const storedCredential = - this.session.modelRegistry.authStorage.get("anthropic"); - if (storedCredential?.type === "oauth") { + if (this.session.modelRuntime.isUsingOAuth("anthropic")) { this.anthropicSubscriptionWarningShown = true; this.showWarning(ANTHROPIC_SUBSCRIPTION_AUTH_WARNING, targetContainer); return; } try { - const apiKey = await this.session.modelRegistry.getApiKeyForProvider( - model.provider, - ); + const storedCredential = await this.session.modelRuntime.getAuth("anthropic"); + const apiKey = storedCredential?.auth.apiKey; if (!isAnthropicSubscriptionAuthKey(apiKey)) { return; } @@ -92,7 +89,7 @@ InteractiveModeBase.prototype.showModelSelector = function(this: InteractiveMode this.ui, this.session.model, this.settingsManager, - this.session.modelRegistry, + this.session.modelRuntime, this.session.scopedModels, async (model) => { try { @@ -121,8 +118,8 @@ InteractiveModeBase.prototype.showModelSelector = function(this: InteractiveMode }; InteractiveModeBase.prototype.showModelsSelector = async function(this: InteractiveModeBase): Promise { - await this.session.modelRegistry.refresh({ allowNetwork: !isOfflineModeEnabled() }); - const allModels = this.session.modelRegistry.getAvailable(); + await this.session.modelRuntime.refresh({ allowNetwork: !isOfflineModeEnabled() }); + const allModels = [...this.session.modelRuntime.getAvailableSnapshot()]; const allModelIds = new Set(allModels.map((model) => `${model.provider}/${model.id}`)); const configuredPatterns = this.settingsManager.getEnabledModels(); const sessionScopedModels = this.session.scopedModels; @@ -133,7 +130,7 @@ InteractiveModeBase.prototype.showModelsSelector = async function(this: Interact } const configuredScope = configuredPatterns?.length - ? await resolveModelScopeWithDiagnostics(configuredPatterns, this.session.modelRegistry) + ? await resolveModelScopeWithDiagnostics(configuredPatterns, this.session.modelRuntime) : undefined; const hasSessionScope = sessionScopedModels.length > 0; let currentEnabledIds: string[] | null = null; @@ -153,7 +150,7 @@ InteractiveModeBase.prototype.showModelsSelector = async function(this: Interact const hasEnabledAvailableModel = enabledIds?.some((id) => allModelIds.has(id)) ?? false; const allAvailableModelsEnabled = enabledIds !== null && [...allModelIds].every((id) => enabledIds.includes(id)); if (enabledIds && hasEnabledAvailableModel && !allAvailableModelsEnabled) { - const newScopedModels = await resolveModelScope(enabledIds, this.session.modelRegistry); + const newScopedModels = await resolveModelScope(enabledIds, this.session.modelRuntime); this.session.setScopedModels(newScopedModels.map((sm) => ({ model: sm.model, thinkingLevel: sm.thinkingLevel }))); } else { this.session.setScopedModels([]); diff --git a/packages/coding-agent/src/modes/interactive/interactive-render-chat.ts b/packages/coding-agent/src/modes/interactive/interactive-render-chat.ts index 074bc4b9d..2ecc9e78b 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-render-chat.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-render-chat.ts @@ -319,7 +319,7 @@ InteractiveModeBase.prototype.renderSessionEntries = function( } flushMessages(); if (this.settingsManager.getShowCacheMissNotices()) { - for (const miss of collectCacheMisses(sessionEntries, { getModel: (provider, model) => this.session.modelRegistry.find(provider, model) }).values()) { + for (const miss of collectCacheMisses(sessionEntries, { getModel: (provider, model) => this.session.modelRuntime.getModel(provider, model) }).values()) { const cause = miss.modelChanged ? " after model switch" : miss.idleMs >= CACHE_TTL_MS ? " after cache TTL expiry" : ""; this.chatContainer.addChild(new Text(theme.fg("warning", `Prompt cache miss${cause}: ${miss.missedTokens.toLocaleString()} tokens re-billed ($${miss.missedCost.toFixed(3)})`), 1, 0)); } diff --git a/packages/coding-agent/src/modes/interactive/interactive-slash-commands.ts b/packages/coding-agent/src/modes/interactive/interactive-slash-commands.ts index 8e437dfb5..4c5eaf77b 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-slash-commands.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-slash-commands.ts @@ -86,7 +86,7 @@ InteractiveModeBase.prototype.handleReloadCommand = async function(this: Interac if (savedImplicitProjectTrust) { this.showStatus("Saved project trust for future sessions"); } - const modelsJsonError = this.session.modelRegistry.getError(); + const modelsJsonError = this.session.modelRuntime.getError(); if (modelsJsonError) { this.showError(`models.json error: ${modelsJsonError}`); } @@ -388,7 +388,7 @@ InteractiveModeBase.prototype.handleSessionCommand = function(this: InteractiveM const promptTokens = assistantEntries.reduce((sum, entry) => sum + (entry.type === "message" && entry.message.role === "assistant" ? entry.message.usage.input + entry.message.usage.cacheRead + entry.message.usage.cacheWrite : 0), 0); const cacheRead = assistantEntries.reduce((sum, entry) => sum + (entry.type === "message" && entry.message.role === "assistant" ? entry.message.usage.cacheRead : 0), 0); if (promptTokens > 0) info += `${theme.fg("dim", "Cache Hit Rate:")} ${((cacheRead / promptTokens) * 100).toFixed(1)}%\n`; - const waste = computeCacheWaste(entries, { getModel: (provider, model) => this.session.modelRegistry.find(provider, model) }); + const waste = computeCacheWaste(entries, { getModel: (provider, model) => this.session.modelRuntime.getModel(provider, model) }); if (waste.missCount > 0) info += `${theme.fg("dim", "Wasted Cache Cost:")} $${waste.missedCost.toFixed(4)} (${waste.missCount} misses)\n`; if (stats.cost > 0) { diff --git a/packages/coding-agent/src/modes/interactive/interactive-startup.ts b/packages/coding-agent/src/modes/interactive/interactive-startup.ts index c1ee085b1..b5b42a3d5 100644 --- a/packages/coding-agent/src/modes/interactive/interactive-startup.ts +++ b/packages/coding-agent/src/modes/interactive/interactive-startup.ts @@ -253,7 +253,7 @@ InteractiveModeBase.prototype.run = async function(this: InteractiveModeBase): P ); } - const modelsJsonError = this.session.modelRegistry.getError(); + const modelsJsonError = this.session.modelRuntime.getError(); if (modelsJsonError) { this.showError(`models.json error: ${modelsJsonError}`); } diff --git a/packages/coding-agent/src/modes/rpc/rpc-client-api.ts b/packages/coding-agent/src/modes/rpc/rpc-client-api.ts index c017cd065..b830054f8 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-client-api.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-client-api.ts @@ -3,7 +3,7 @@ import type { AgentMessage, ThinkingLevel } from "@earendil-works/pi-agent-core" import type { Api, ImageContent, Model } from "@earendil-works/pi-ai/compat"; import type { AtomicProviderCompat } from "../../core/model-capabilities.ts"; import type { SessionStats } from "../../core/agent-session.ts"; -import type { AuthCredential } from "../../core/auth-storage.ts"; +import type { Credential } from "@earendil-works/pi-ai"; import type { BashResult } from "../../core/bash-executor.ts"; import type { VerbatimCompactionResult } from "../../core/compaction/index.ts"; import type { SessionEntry, SessionTreeNode } from "../../core/session-manager.ts"; @@ -67,7 +67,7 @@ export abstract class RpcClientApi { signal?.removeEventListener("abort", cancel); } } - async saveProviderCredential(provider: string, credential: AuthCredential): Promise { + async saveProviderCredential(provider: string, credential: Credential): Promise { return this.data(await this.request({ type: "save_provider_credential", provider, credential })); } async cancelLoginProvider(provider: string, loginId?: string): Promise { diff --git a/packages/coding-agent/src/modes/rpc/rpc-command-handler.ts b/packages/coding-agent/src/modes/rpc/rpc-command-handler.ts index 7bc3e6f2a..465ad7de9 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-command-handler.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-command-handler.ts @@ -3,7 +3,6 @@ import { runCallback } from "../../core/callback-activity.ts"; import { KeybindingsManager } from "../../core/keybindings.ts"; import type { AgentSession } from "../../core/agent-session.ts"; import type { AgentSessionRuntime } from "../../core/agent-session-runtime.ts"; -import { getOAuthProviderMetadata } from "../../core/oauth-provider-bridge.ts"; import { createRpcErrorResponse, createRpcSuccessResponse, @@ -127,7 +126,7 @@ export function createRpcCommandHandler({ return createRpcSuccessResponse(id, "get_state", state); } case "set_model": { - const models = await session.modelRegistry.getAvailable(); + const models = await session.modelRuntime.getAvailableSnapshot(); const model = models.find((candidate) => candidate.provider === command.provider && candidate.id === command.modelId); if (!model) { return createRpcErrorResponse(id, "set_model", `Model not found: ${command.provider}/${command.modelId}`); @@ -143,12 +142,12 @@ export function createRpcCommandHandler({ return createRpcSuccessResponse(id, "cycle_model", result ?? null); } case "get_available_models": { - const models = await session.modelRegistry.getAvailable(); + const models = await session.modelRuntime.getAvailableSnapshot(); return createRpcSuccessResponse(id, "get_available_models", { models, scopedModels: session.scopedModels, - customAuthProviders: session.modelRegistry.getCustomApiKeyAuthProviders(), - oauthProviders: getOAuthProviderMetadata(), + customAuthProviders: [], + oauthProviders: session.modelRuntime.getOAuthProviderMetadata(), }); } @@ -175,20 +174,26 @@ export function createRpcCommandHandler({ return createRpcSuccessResponse(id, "logout_provider", result); } case "refresh_models": { - session.modelRegistry.authStorage.reload(); - const result = await session.modelRegistry.refresh({ - timeoutMs: command.timeoutMs, - force: command.force, - allowNetwork: command.allowNetwork, - }); - return createRpcSuccessResponse(id, "refresh_models", { - aborted: result.aborted, - errors: [...result.errors].map(([provider, error]) => ({ provider, message: error.message })), - models: session.modelRegistry.getAvailable(), - scopedModels: session.scopedModels, - customAuthProviders: session.modelRegistry.getCustomApiKeyAuthProviders(), - oauthProviders: getOAuthProviderMetadata(), - }); + await session.modelRuntime.reloadCredentials(); + const controller = command.timeoutMs === undefined ? undefined : new AbortController(); + const timeout = controller ? setTimeout(() => controller.abort(), command.timeoutMs) : undefined; + try { + const result = await session.modelRuntime.refresh({ + allowNetwork: command.allowNetwork, + force: command.force, + signal: controller?.signal, + }); + return createRpcSuccessResponse(id, "refresh_models", { + aborted: result.aborted, + errors: [...result.errors].map(([provider, error]) => ({ provider, message: error.message })), + models: session.modelRuntime.getAvailableSnapshot(), + scopedModels: session.scopedModels, + customAuthProviders: [], + oauthProviders: session.modelRuntime.getOAuthProviderMetadata(), + }); + } finally { + if (timeout) clearTimeout(timeout); + } } case "set_thinking_level": { diff --git a/packages/coding-agent/src/modes/rpc/rpc-oauth-client.ts b/packages/coding-agent/src/modes/rpc/rpc-oauth-client.ts index 908713780..aa3c3f6bd 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-oauth-client.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-oauth-client.ts @@ -1,6 +1,6 @@ import { randomUUID } from "node:crypto"; -import type { AtomicOAuthLoginCallbacks } from "../../core/oauth-provider-bridge.ts"; -import { normalizeOAuthLoginError } from "../../core/oauth-provider-bridge.ts"; +import type { AtomicOAuthLoginCallbacks } from "../../core/oauth-login.ts"; +import { normalizeOAuthLoginError } from "../../core/oauth-login.ts"; import type { RpcExtensionUIRequest, RpcExtensionUIResponse } from "./rpc-types.ts"; import type { RpcOAuthLoginProviderResult } from "./rpc-types.ts"; diff --git a/packages/coding-agent/src/modes/rpc/rpc-oauth-interaction.ts b/packages/coding-agent/src/modes/rpc/rpc-oauth-interaction.ts index 55e19b395..89265713d 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-oauth-interaction.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-oauth-interaction.ts @@ -1,6 +1,6 @@ import { randomUUID } from "node:crypto"; -import type { AtomicOAuthLoginCallbacks } from "../../core/oauth-provider-bridge.ts"; -import { normalizeOAuthLoginError } from "../../core/oauth-provider-bridge.ts"; +import type { AtomicOAuthLoginCallbacks } from "../../core/oauth-login.ts"; +import { normalizeOAuthLoginError } from "../../core/oauth-login.ts"; import type { RpcPendingExtensionRequests } from "./rpc-extension-ui.ts"; import type { RpcOutput } from "./rpc-responses.ts"; import type { RpcExtensionUIRequest, RpcExtensionUIResponse } from "./rpc-types.ts"; diff --git a/packages/coding-agent/src/modes/rpc/rpc-provider-auth.ts b/packages/coding-agent/src/modes/rpc/rpc-provider-auth.ts index db63a4241..1667f963e 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-provider-auth.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-provider-auth.ts @@ -1,148 +1,79 @@ +import type { Credential } from "@earendil-works/pi-ai"; import type { AgentSession } from "../../core/agent-session.ts"; -import type { AuthCredential } from "../../core/auth-storage.ts"; import type { HostInputFormRequest } from "../../core/extensions/ui-types.ts"; -import { - getOAuthProviderMetadata, - isOAuthLoginCancelled, - loginOAuthProvider, - normalizeOAuthLoginError, -} from "../../core/oauth-provider-bridge.ts"; +import { createAuthInteraction, isOAuthLoginCancelled } from "../../core/oauth-login.ts"; import { createRpcOAuthCallbacks, type OAuthInteractionTransport } from "./rpc-oauth-interaction.ts"; import type { RpcLoginProviderResult, RpcModelCatalog, RpcOAuthLoginProviderResult } from "./rpc-types.ts"; export interface ProviderLoginInput { open(request: HostInputFormRequest, signal?: AbortSignal): Promise | undefined>; } - -interface ActiveLogin { - provider: string; - controller: AbortController; -} +interface ActiveLogin { provider: string; controller: AbortController; } export class RpcProviderAuth { private readonly controllers = new Map(); - private readonly inputForm?: ProviderLoginInput; - private readonly oauthTransport?: OAuthInteractionTransport; - + private readonly inputForm: ProviderLoginInput | undefined; + private readonly oauthTransport: OAuthInteractionTransport | undefined; constructor(inputForm?: ProviderLoginInput, oauthTransport?: OAuthInteractionTransport) { this.inputForm = inputForm; this.oauthTransport = oauthTransport; } async login(session: AgentSession, provider: string, loginId = provider): Promise { - const customAuth = session.modelRegistry.getCustomApiKeyAuth(provider); - if (!customAuth) throw new Error(`Provider does not support custom API-key login: ${provider}`); + if (!session.modelRuntime.getProvider(provider)?.auth.apiKey) throw new Error(`Provider does not support API-key login: ${provider}`); if (!this.inputForm) throw new Error("Provider login requires an interactive input host"); const controller = this.begin(provider, loginId); try { - const credential = await customAuth.login({ + const credential = await session.modelRuntime.login(provider, "api_key", { signal: controller.signal, prompt: async (prompt) => { - const values = await this.inputForm!.open({ - title: prompt.message, - heading: "PROVIDER LOGIN", - submitLabel: "[ Submit ]", - fields: [{ - name: "value", type: "string", required: false, initialValue: "", - placeholder: prompt.placeholder, - }], - }, controller.signal); + const values = await this.inputForm!.open({ title: prompt.message, heading: "PROVIDER LOGIN", submitLabel: "[ Submit ]", fields: [{ name: "value", type: "string", required: false, initialValue: "", placeholder: "placeholder" in prompt ? prompt.placeholder : undefined }] }, controller.signal); if (!values || controller.signal.aborted) throw new Error("Login cancelled"); return values.value ?? ""; }, + notify: () => {}, }); if (controller.signal.aborted) return { provider, cancelled: true }; - session.modelRegistry.authStorage.set(provider, credential); - await session.modelRegistry.refresh(); + if (credential.type !== "api_key") throw new Error(`Provider returned an unexpected ${credential.type} credential`); return { provider, cancelled: false, credential, ...this.catalog(session) }; } catch (error) { - if (controller.signal.aborted || (error instanceof Error && error.message === "Login cancelled")) { - return { provider, cancelled: true }; - } + if (controller.signal.aborted || (error instanceof Error && error.message === "Login cancelled")) return { provider, cancelled: true }; throw error; - } finally { - this.finish(loginId, controller); - } + } finally { this.finish(loginId, controller); } } async loginOAuth(session: AgentSession, provider: string, loginId = provider): Promise { - if (!getOAuthProviderMetadata().some(({ id }) => id === provider)) { - throw new Error(`Unknown OAuth provider: ${provider}`); - } + if (!session.modelRuntime.getProvider(provider)?.auth.oauth) throw new Error(`Unknown OAuth provider: ${provider}`); if (!this.oauthTransport) throw new Error("OAuth login requires an interactive host"); const controller = this.begin(provider, loginId); try { const callbacks = createRpcOAuthCallbacks(provider, loginId, controller.signal, this.oauthTransport); - let credential: AuthCredential; - try { - credential = await loginOAuthProvider(provider, callbacks); - } catch (error) { - if (isOAuthLoginCancelled(error, controller.signal)) return { provider, cancelled: true }; - throw error; - } + try { await session.modelRuntime.login(provider, "oauth", createAuthInteraction(callbacks)); } + catch (error) { if (isOAuthLoginCancelled(error, controller.signal)) return { provider, cancelled: true }; throw error; } if (controller.signal.aborted) return { provider, cancelled: true }; - await this.persistOAuthAndRefresh(session, provider, credential, controller.signal); + session.refreshCurrentModelFromRegistry(); return { provider, cancelled: false, ...this.catalog(session) }; - } finally { - this.finish(loginId, controller); - } + } finally { this.finish(loginId, controller); } } - async save(session: AgentSession, provider: string, credential: AuthCredential): Promise { - await session.modelRegistry.authStorage.asCredentialStore().modify(provider, async () => credential); - await session.modelRegistry.refresh(); + async save(session: AgentSession, provider: string, credential: Credential): Promise { + await session.modelRuntime.saveCredential(provider, credential); session.refreshCurrentModelFromRegistry(); return this.catalog(session); } - cancel(provider: string, loginId?: string): void { - if (loginId !== undefined) { - const active = this.controllers.get(loginId); - if (active?.provider === provider) active.controller.abort(); - return; - } - for (const active of this.controllers.values()) { - if (active.provider === provider) active.controller.abort(); - } + if (loginId !== undefined) { const active = this.controllers.get(loginId); if (active?.provider === provider) active.controller.abort(); return; } + for (const active of this.controllers.values()) if (active.provider === provider) active.controller.abort(); } - private begin(provider: string, loginId: string): AbortController { - if ([...this.controllers.values()].some((active) => active.provider === provider)) { - throw new Error(`Login already in progress: ${provider}`); - } - const controller = new AbortController(); - this.controllers.set(loginId, { provider, controller }); - return controller; + if ([...this.controllers.values()].some((active) => active.provider === provider)) throw new Error(`Login already in progress: ${provider}`); + const controller = new AbortController(); this.controllers.set(loginId, { provider, controller }); return controller; } - - private finish(loginId: string, controller: AbortController): void { - if (this.controllers.get(loginId)?.controller === controller) this.controllers.delete(loginId); - } - - private async persistOAuthAndRefresh( - session: AgentSession, - provider: string, - credential: AuthCredential, - signal: AbortSignal, - ): Promise { - const credentialStore = session.modelRegistry.authStorage.asCredentialStore(); - if (signal.aborted) throw normalizeOAuthLoginError(signal.reason, signal); - await credentialStore.modify(provider, async () => credential); - // Upstream pi parity: once fresh OAuth tokens are persisted, the catalog - // refresh outcome (per-provider errors or a timeout-aborted result) must - // not fail the login or touch the stored credential. Refresh tokens are - // rotated by providers, so rolling back here would re-install a - // server-side-invalidated credential (permanent invalid_grant). - await session.modelRegistry.refresh(); - session.refreshCurrentModelFromRegistry(); - } - + private finish(loginId: string, controller: AbortController): void { if (this.controllers.get(loginId)?.controller === controller) this.controllers.delete(loginId); } private catalog(session: AgentSession): RpcModelCatalog { return { - models: session.modelRegistry.getAvailable(), - scopedModels: [...session.scopedModels], - customAuthProviders: session.modelRegistry.getCustomApiKeyAuthProviders(), - oauthProviders: getOAuthProviderMetadata(), + models: [...session.modelRuntime.getAvailableSnapshot()], scopedModels: [...session.scopedModels], customAuthProviders: [], + oauthProviders: session.modelRuntime.getOAuthProviderMetadata(), }; } } diff --git a/packages/coding-agent/src/modes/rpc/rpc-session-binding.ts b/packages/coding-agent/src/modes/rpc/rpc-session-binding.ts index 81730b966..12f03c472 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-session-binding.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-session-binding.ts @@ -66,7 +66,7 @@ export class RpcSessionBinding { // sessions instead of defaulting to 0. const models = session.scopedModels.length > 0 ? session.scopedModels.map((scoped) => scoped.model) - : session.modelRegistry.getAvailable(); + : session.modelRuntime.getAvailableSnapshot(); this.footerDataProvider.setAvailableProviderCount(new Set(models.map((model) => model.provider)).size); try { diff --git a/packages/coding-agent/src/modes/rpc/rpc-types.ts b/packages/coding-agent/src/modes/rpc/rpc-types.ts index 554bb91e9..2487f5b6b 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-types.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-types.ts @@ -9,14 +9,15 @@ import type { AgentMessage, ThinkingLevel } from "@earendil-works/pi-agent-core" import type { AuthInfoLink, OAuthAuthInfo, OAuthDeviceCodeInfo, OAuthPrompt, OAuthSelectPrompt } from "@earendil-works/pi-ai"; import type { Api, ImageContent, Model } from "@earendil-works/pi-ai/compat"; import type { AgentSessionEvent, SessionStats } from "../../core/agent-session.ts"; -import type { AuthCredential, AuthStatus } from "../../core/auth-storage.ts"; +import type { AuthStatus } from "../../core/provider-composer.ts"; +import type { Credential } from "@earendil-works/pi-ai"; import type { BashResult } from "../../core/bash-executor.ts"; import type { VerbatimCompactionResult } from "../../core/compaction/index.ts"; import type { SessionEntry, SessionTreeNode } from "../../core/session-manager.ts"; import type { SourceInfo } from "../../core/source-info.ts"; import type { ResourceOverlap } from "../../core/diagnostics.ts"; import type { ModelFallbackReason } from "../../core/model-resolver-types.ts"; -import type { OAuthProviderMetadata } from "../../core/oauth-provider-bridge.ts"; +import type { OAuthProviderMetadata } from "../../core/oauth-login.ts"; // ============================================================================ // RPC Commands (stdin) @@ -37,7 +38,7 @@ export type RpcLoginProviderResult = | (RpcModelCatalog & { provider: string; cancelled: false; - credential: import("../../core/auth-storage.ts").ApiKeyCredential; + credential: Extract; }) | { provider: string; cancelled: true }; @@ -61,7 +62,7 @@ export type RpcCommand = | { id?: string; type: "cycle_model"; direction?: "forward" | "backward" } | { id?: string; type: "get_available_models" } | { id?: string; type: "login_provider"; provider: string; authType?: "api_key" | "oauth"; loginId?: string } - | { id?: string; type: "save_provider_credential"; provider: string; credential: AuthCredential } + | { id?: string; type: "save_provider_credential"; provider: string; credential: Credential } | { id?: string; type: "cancel_login_provider"; provider: string; loginId?: string } | { id?: string; type: "logout_provider"; provider: string } | { id?: string; type: "refresh_models"; timeoutMs?: number; force?: boolean; allowNetwork?: boolean } diff --git a/packages/coding-agent/src/package-manager-cli.ts b/packages/coding-agent/src/package-manager-cli.ts index 487791de9..2075596ea 100644 --- a/packages/coding-agent/src/package-manager-cli.ts +++ b/packages/coding-agent/src/package-manager-cli.ts @@ -22,7 +22,7 @@ import type { InlineExtension } from "./core/extensions/types.ts"; import { DefaultPackageManager } from "./core/package-manager.ts"; import { type AppMode, resolveProjectTrusted } from "./core/project-trust.ts"; import { DefaultResourceLoader } from "./core/resource-loader.ts"; -import { ModelRegistry } from "./core/model-registry.ts"; +import { ModelRuntime } from "./core/model-runtime.ts"; import { SettingsManager } from "./core/settings-manager.ts"; import { hasProjectTrustInputs, ProjectTrustStore } from "./core/trust-manager.ts"; import { spawnProcess } from "./utils/child-process.ts"; @@ -76,19 +76,17 @@ export async function refreshModelCatalogs( if (!loaded) throw new Error("Model catalog refresh timed out."); const authPaths = [join(agentDir, "auth.json"), ...getAgentConfigPaths("auth.json")] .filter((path, index, paths) => paths.indexOf(path) === index); - const modelPaths = [join(agentDir, "models.json"), ...getAgentConfigPaths("models.json")] - .filter((path, index, paths) => paths.indexOf(path) === index); - const modelRegistry = ModelRegistry.create(AuthStorage.create(authPaths), modelPaths); + const modelRuntime = await ModelRuntime.create({ credentials: AuthStorage.create(authPaths), modelsPath: join(agentDir, "models.json") }); const extensionsResult = resourceLoader.getExtensions(); if (extensionsResult.errors.length > 0) { const details = extensionsResult.errors.map(({ path, error }) => `${path}: ${error}`).join("; "); throw new Error(`Could not load extensions for model catalog refresh: ${details}`); } for (const registration of extensionsResult.runtime.pendingProviderRegistrations) { - if ("provider" in registration) modelRegistry.registerProvider(registration.provider); - else modelRegistry.registerProvider(registration.name, registration.config); + if ("provider" in registration) modelRuntime.registerNativeProvider(registration.provider); + else modelRuntime.registerProvider(registration.name, registration.config); } - const refresh = modelRegistry.refresh({ allowNetwork: true, force: true, signal: controller.signal }); + const refresh = modelRuntime.refresh({ allowNetwork: true, force: true, signal: controller.signal }); const result = await Promise.race([refresh, aborted]); if (!result || result.aborted) throw new Error("Model catalog refresh timed out."); if (result.errors.size > 0) { diff --git a/packages/coding-agent/test/agent-session-async-bash.test.ts b/packages/coding-agent/test/agent-session-async-bash.test.ts index 954e73332..e6b947806 100644 --- a/packages/coding-agent/test/agent-session-async-bash.test.ts +++ b/packages/coding-agent/test/agent-session-async-bash.test.ts @@ -8,7 +8,7 @@ import { AgentSession } from "../src/core/agent-session.ts"; import { AsyncJobManager } from "../src/core/async/job-manager.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { convertToLlm } from "../src/core/messages.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createTestResourceLoader } from "./utilities.ts"; @@ -55,7 +55,7 @@ function messageText(message: AgentMessage): string { .join("\n"); } -function createSession(tempDir: string, onTurn: (userTexts: string[], stream: MockAssistantStream) => void): AgentSession { +async function createSession(tempDir: string, onTurn: (userTexts: string[], stream: MockAssistantStream) => void): AgentSession { const model = getModel("anthropic", "claude-sonnet-4-5")!; const agent = new Agent({ convertToLlm, @@ -68,13 +68,13 @@ function createSession(tempDir: string, onTurn: (userTexts: string[], stream: Mo }, }); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); return new AgentSession({ agent, sessionManager: SessionManager.inMemory(), settingsManager: SettingsManager.create(tempDir, tempDir), cwd: tempDir, - modelRegistry: ModelRegistry.create(authStorage, tempDir), + modelRuntime: await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }), resourceLoader: createTestResourceLoader(), }); } @@ -97,7 +97,7 @@ describe("AgentSession async bash auto-delivery", () => { it("starts an idle follow-up turn from actual async bash completion", async () => { const turns: string[][] = []; - session = createSession(tempDir, (userTexts, stream) => { + session = await createSession(tempDir, (userTexts, stream) => { turns.push(userTexts); stream.push({ type: "start", partial: createAssistantMessage("") }); stream.push({ type: "done", reason: "stop", message: createAssistantMessage("done") }); @@ -111,7 +111,7 @@ describe("AgentSession async bash auto-delivery", () => { it("queues actual async bash completion as a follow-up while streaming and drains it after the turn", async () => { let finishFirstTurn: (() => void) | undefined; const turns: string[][] = []; - session = createSession(tempDir, (userTexts, stream) => { + session = await createSession(tempDir, (userTexts, stream) => { turns.push(userTexts); stream.push({ type: "start", partial: createAssistantMessage("") }); if (userTexts.some((text) => text.includes("streaming-async"))) { @@ -135,7 +135,7 @@ describe("AgentSession async bash auto-delivery", () => { it("keeps a streaming async result admitted when later polling acknowledges the job", async () => { let finishFirstTurn: (() => void) | undefined; const turns: string[][] = []; - session = createSession(tempDir, (userTexts, stream) => { + session = await createSession(tempDir, (userTexts, stream) => { turns.push(userTexts); stream.push({ type: "start", partial: createAssistantMessage("") }); if (userTexts.some((text) => text.includes("stale-async"))) { diff --git a/packages/coding-agent/test/agent-session-auth-load-failure.test.ts b/packages/coding-agent/test/agent-session-auth-load-failure.test.ts index 2aa41b402..11a8676aa 100644 --- a/packages/coding-agent/test/agent-session-auth-load-failure.test.ts +++ b/packages/coding-agent/test/agent-session-auth-load-failure.test.ts @@ -1,17 +1,4 @@ -/** - * Regression: a prompt preflight must not misreport a credential-store LOAD - * failure as "No API key found". - * - * When a fresh AuthStorage cannot read auth.json (e.g. it is briefly locked by a - * concurrent process, leaving an ELOCKED error), it ends up with an empty - * in-memory credential set and a recorded loadError. Previously the prompt - * preflight only saw `hasConfiguredAuth() === false` and threw the misleading - * "No API key found for " \u2014 even though the credentials exist on disk. - * The preflight now surfaces the real load failure instead (issue #1431). - * - * cross-ref: packages/coding-agent/src/core/agent-session.ts (prompt preflight) - * packages/coding-agent/src/core/auth-storage.ts (getLoadError) - */ +/** Pinned-pi credential stores preserve an empty snapshot when their initial read fails. */ import { mkdirSync, rmSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; @@ -20,9 +7,9 @@ import { getModel } from "@earendil-works/pi-ai/compat"; import { afterEach, beforeEach, describe, expect, it } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage, type AuthStorageBackend } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; import { createTestResourceLoader } from "./utilities.ts"; class ThrowingAuthStorageBackend implements AuthStorageBackend { @@ -38,7 +25,7 @@ class ThrowingAuthStorageBackend implements AuthStorageBackend { } } -describe("AgentSession prompt preflight \u2014 auth-storage load failure (#1431)", () => { +describe("AgentSession prompt preflight after an auth-storage load failure", () => { let session: AgentSession | undefined; let tempDir: string; let savedAnthropicKey: string | undefined; @@ -46,7 +33,6 @@ describe("AgentSession prompt preflight \u2014 auth-storage load failure (#1431) beforeEach(() => { tempDir = join(tmpdir(), `pi-auth-load-failure-${Date.now()}-${Math.random().toString(36).slice(2)}`); mkdirSync(tempDir, { recursive: true }); - // Ensure no environment key masks the load failure for this provider. savedAnthropicKey = process.env.ANTHROPIC_API_KEY; delete process.env.ANTHROPIC_API_KEY; }); @@ -58,10 +44,9 @@ describe("AgentSession prompt preflight \u2014 auth-storage load failure (#1431) if (tempDir) rmSync(tempDir, { recursive: true, force: true }); }); - it("surfaces the load failure instead of 'No API key found'", async () => { + it("reports unavailable provider authentication through the runtime", async () => { const loadError = Object.assign(new Error("Lock file is already being held"), { code: "ELOCKED" }); const model = getModel("anthropic", "claude-sonnet-4-5")!; - const agent = new Agent({ getApiKey: () => "test-key", initialState: { model, systemPrompt: "Test", tools: [] }, @@ -69,35 +54,19 @@ describe("AgentSession prompt preflight \u2014 auth-storage load failure (#1431) throw new Error("streamFn must not run when preflight fails"); }, }); - const authStorage = AuthStorage.fromStorage(new ThrowingAuthStorageBackend(loadError)); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - + const modelRegistry = await createModelRegistry(authStorage, tempDir); session = new AgentSession({ agent, sessionManager: SessionManager.inMemory(), settingsManager: SettingsManager.create(tempDir, tempDir), cwd: tempDir, - modelRegistry, + modelRuntime: getModelRuntime(modelRegistry), resourceLoader: createTestResourceLoader(), }); - // Sanity: the store genuinely failed to load and looks "empty". - expect(authStorage.getLoadError()).toBe(loadError); + expect(await authStorage.read("anthropic")).toBeUndefined(); expect(modelRegistry.hasConfiguredAuth(model)).toBe(false); - - await session.prompt("hello").then( - () => { - throw new Error("expected prompt() to reject"); - }, - (error: unknown) => { - const text = error instanceof Error ? error.message : String(error); - expect(text).toContain("Could not load stored credentials for anthropic"); - expect(text).toContain("Lock file is already being held"); - expect(text).not.toContain("No API key found"); - // The original load error is preserved as the cause. - expect((error as { cause?: unknown }).cause).toBe(loadError); - }, - ); + await expect(session.prompt("hello")).rejects.toThrow("No API key found for anthropic"); }); }); diff --git a/packages/coding-agent/test/agent-session-auto-compaction-overflow-await.suite.ts b/packages/coding-agent/test/agent-session-auto-compaction-overflow-await.suite.ts index 1e87859ed..7e6830906 100644 --- a/packages/coding-agent/test/agent-session-auto-compaction-overflow-await.suite.ts +++ b/packages/coding-agent/test/agent-session-auto-compaction-overflow-await.suite.ts @@ -6,7 +6,7 @@ import { type AssistantMessage, getModel } from "@earendil-works/pi-ai/compat"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.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 { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createTestResourceLoader } from "./utilities.ts"; @@ -39,7 +39,7 @@ describe("AgentSession overflow auto-compaction continuation", () => { let session: AgentSession; let tempDir: string; - beforeEach(() => { + beforeEach(async () => { tempDir = join(tmpdir(), `pi-overflow-continuation-${Date.now()}`); mkdirSync(tempDir, { recursive: true }); vi.useFakeTimers(); @@ -47,13 +47,13 @@ describe("AgentSession overflow auto-compaction continuation", () => { const model = getModel("anthropic", "claude-sonnet-4-5")!; const agent = new Agent({ initialState: { model, systemPrompt: "Test", tools: [] } }); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); session = new AgentSession({ agent, sessionManager: SessionManager.inMemory(), settingsManager: SettingsManager.create(tempDir, tempDir), cwd: tempDir, - modelRegistry: ModelRegistry.create(authStorage, tempDir), + modelRuntime: await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }), resourceLoader: createTestResourceLoader(), }); }); diff --git a/packages/coding-agent/test/agent-session-auto-compaction-queue-01.suite.ts b/packages/coding-agent/test/agent-session-auto-compaction-queue-01.suite.ts index 044d0a1fc..17c5c3358 100644 --- a/packages/coding-agent/test/agent-session-auto-compaction-queue-01.suite.ts +++ b/packages/coding-agent/test/agent-session-auto-compaction-queue-01.suite.ts @@ -6,7 +6,7 @@ import { type AssistantMessage, getModel } from "@earendil-works/pi-ai/compat"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.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 { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createFauxStreamFn } from "./test-harness.ts"; @@ -73,7 +73,7 @@ describe("AgentSession auto-compaction queue resume", () => { let sessionManager: SessionManager; let tempDir: string; - beforeEach(() => { + beforeEach(async () => { compactionMocks.runVerbatimCompaction.mockClear(); tempDir = join(tmpdir(), `pi-auto-compaction-queue-${Date.now()}-${Math.random().toString(36).slice(2)}`); mkdirSync(tempDir, { recursive: true }); @@ -91,15 +91,15 @@ describe("AgentSession auto-compaction queue resume", () => { sessionManager.appendMessage({ role: "user", content: [{ type: "text", text: "existing compactable context" }], timestamp: Date.now() }); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - authStorage.setRuntimeApiKey("anthropic", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); }); diff --git a/packages/coding-agent/test/agent-session-auto-compaction-queue-02.suite.ts b/packages/coding-agent/test/agent-session-auto-compaction-queue-02.suite.ts index 1f8664a0d..3837d73aa 100644 --- a/packages/coding-agent/test/agent-session-auto-compaction-queue-02.suite.ts +++ b/packages/coding-agent/test/agent-session-auto-compaction-queue-02.suite.ts @@ -6,7 +6,7 @@ import { type AssistantMessage, getModel } from "@earendil-works/pi-ai/compat"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.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 { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createTestResourceLoader } from "./utilities.ts"; @@ -72,7 +72,7 @@ describe("AgentSession auto-compaction queue resume", () => { let sessionManager: SessionManager; let tempDir: string; - beforeEach(() => { + beforeEach(async () => { compactionMocks.runVerbatimCompaction.mockClear(); tempDir = join(tmpdir(), `pi-auto-compaction-queue-${Date.now()}`); mkdirSync(tempDir, { recursive: true }); @@ -91,15 +91,15 @@ describe("AgentSession auto-compaction queue resume", () => { sessionManager.appendMessage({ role: "user", content: [{ type: "text", text: "existing compactable context" }], timestamp: Date.now() }); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - authStorage.setRuntimeApiKey("anthropic", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); }); diff --git a/packages/coding-agent/test/agent-session-auto-compaction-queue-03.suite.ts b/packages/coding-agent/test/agent-session-auto-compaction-queue-03.suite.ts index c2ea4dad9..d463db10d 100644 --- a/packages/coding-agent/test/agent-session-auto-compaction-queue-03.suite.ts +++ b/packages/coding-agent/test/agent-session-auto-compaction-queue-03.suite.ts @@ -10,7 +10,7 @@ import { MAX_OUTPUT_BUDGET_ERROR_CONTINUATION_ATTEMPTS, } from "../src/core/agent-session-auto-compaction.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 { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createTestResourceLoader } from "./utilities.ts"; @@ -51,7 +51,7 @@ describe("AgentSession auto-compaction length-stop resume", () => { let sessionManager: SessionManager; let tempDir: string; - beforeEach(() => { + beforeEach(async () => { compactionMocks.runVerbatimCompaction.mockClear(); compactionMocks.estimateContextTokens.mockReset(); compactionMocks.estimateContextTokens.mockReturnValue({ tokens: 0, usageTokens: 0, trailingTokens: 0, lastUsageIndex: null }); @@ -65,15 +65,15 @@ describe("AgentSession auto-compaction length-stop resume", () => { sessionManager.appendMessage({ role: "user", content: [{ type: "text", text: "existing compactable context" }], timestamp: Date.now() }); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - authStorage.setRuntimeApiKey("anthropic", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); }); diff --git a/packages/coding-agent/test/agent-session-branching.test.ts b/packages/coding-agent/test/agent-session-branching.test.ts index 76d355e16..7016ecca9 100644 --- a/packages/coding-agent/test/agent-session-branching.test.ts +++ b/packages/coding-agent/test/agent-session-branching.test.ts @@ -48,7 +48,7 @@ describe.skipIf(!API_KEY)("AgentSession forking", () => { const model = getModel("anthropic", "claude-sonnet-4-5")!; sessionManager = noSession ? SessionManager.inMemory(tempDir) : SessionManager.create(tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - authStorage.setRuntimeApiKey("anthropic", API_KEY!); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: API_KEY! })); const servicesOptions = { agentDir: tempDir, diff --git a/packages/coding-agent/test/agent-session-concurrent-01.suite.ts b/packages/coding-agent/test/agent-session-concurrent-01.suite.ts index fee8bc265..4de1e7b88 100644 --- a/packages/coding-agent/test/agent-session-concurrent-01.suite.ts +++ b/packages/coding-agent/test/agent-session-concurrent-01.suite.ts @@ -18,7 +18,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { convertToLlm } from "../src/core/messages.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import type { BuildSystemPromptOptions } from "../src/core/system-prompt.ts"; @@ -108,7 +108,7 @@ describe("AgentSession concurrent prompt guard", () => { let tempDir: string; beforeEach(() => { - tempDir = join(tmpdir(), `pi-concurrent-test-${Date.now()}`); + tempDir = join(tmpdir(), `pi-concurrent-test-${Date.now()}-${Math.random().toString(36).slice(2)}`); mkdirSync(tempDir, { recursive: true }); }); @@ -123,7 +123,7 @@ describe("AgentSession concurrent prompt guard", () => { } }); - function createSession() { + async function createSession() { const model = getModel("anthropic", "claude-sonnet-4-5")!; let abortSignal: AbortSignal | undefined; @@ -156,16 +156,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - // Set a runtime API key so validation passes - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -173,8 +172,7 @@ describe("AgentSession concurrent prompt guard", () => { } it("should throw when prompt() called while streaming", async () => { - createSession(); - + await createSession(); // Start first prompt (don't await, it will block until abort) const firstPrompt = session.prompt("First message"); @@ -194,8 +192,7 @@ describe("AgentSession concurrent prompt guard", () => { await firstPrompt.catch(() => {}); // Ignore abort error }); it("should allow steer() while streaming", async () => { - createSession(); - + await createSession(); // Start first prompt const firstPrompt = session.prompt("First message"); await new Promise((resolve) => setTimeout(resolve, 10)); @@ -209,8 +206,7 @@ describe("AgentSession concurrent prompt guard", () => { await firstPrompt.catch(() => {}); }); it("should allow followUp() while streaming", async () => { - createSession(); - + await createSession(); // Start first prompt const firstPrompt = session.prompt("First message"); await new Promise((resolve) => setTimeout(resolve, 10)); @@ -269,8 +265,8 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); const extensionsResult = await createTestExtensionsResult([ (pi) => { @@ -288,7 +284,7 @@ describe("AgentSession concurrent prompt guard", () => { sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader({ extensionsResult }), }); session.subscribe((event) => { @@ -382,15 +378,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); diff --git a/packages/coding-agent/test/agent-session-concurrent-02.suite.ts b/packages/coding-agent/test/agent-session-concurrent-02.suite.ts index 2d9f3101c..003caeacd 100644 --- a/packages/coding-agent/test/agent-session-concurrent-02.suite.ts +++ b/packages/coding-agent/test/agent-session-concurrent-02.suite.ts @@ -18,7 +18,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { convertToLlm } from "../src/core/messages.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import type { BuildSystemPromptOptions } from "../src/core/system-prompt.ts"; @@ -123,7 +123,7 @@ describe("AgentSession concurrent prompt guard", () => { } }); - function createSession() { + async function createSession() { const model = getModel("anthropic", "claude-sonnet-4-5")!; let abortSignal: AbortSignal | undefined; @@ -156,16 +156,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - // Set a runtime API key so validation passes - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -173,7 +172,7 @@ describe("AgentSession concurrent prompt guard", () => { } - it("should replace generic abort events for interrupt custom messages", () => { + it("should replace generic abort events for interrupt custom messages", async () => { const model = getModel("anthropic", "claude-sonnet-4-5")!; const agent = new Agent({ getApiKey: () => "test-key", @@ -186,15 +185,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -281,15 +280,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -362,15 +361,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); diff --git a/packages/coding-agent/test/agent-session-concurrent-03.suite.ts b/packages/coding-agent/test/agent-session-concurrent-03.suite.ts index 558cf98f0..3ec7dc8d2 100644 --- a/packages/coding-agent/test/agent-session-concurrent-03.suite.ts +++ b/packages/coding-agent/test/agent-session-concurrent-03.suite.ts @@ -18,7 +18,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { convertToLlm } from "../src/core/messages.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import type { BuildSystemPromptOptions } from "../src/core/system-prompt.ts"; @@ -123,7 +123,7 @@ describe("AgentSession concurrent prompt guard", () => { } }); - function createSession() { + async function createSession() { const model = getModel("anthropic", "claude-sonnet-4-5")!; let abortSignal: AbortSignal | undefined; @@ -156,16 +156,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - // Set a runtime API key so validation passes - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -217,15 +216,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -307,15 +306,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -391,15 +390,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); diff --git a/packages/coding-agent/test/agent-session-concurrent-04.suite.ts b/packages/coding-agent/test/agent-session-concurrent-04.suite.ts index e285be843..8eb47ceff 100644 --- a/packages/coding-agent/test/agent-session-concurrent-04.suite.ts +++ b/packages/coding-agent/test/agent-session-concurrent-04.suite.ts @@ -18,7 +18,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { convertToLlm } from "../src/core/messages.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import type { BuildSystemPromptOptions } from "../src/core/system-prompt.ts"; @@ -123,7 +123,7 @@ describe("AgentSession concurrent prompt guard", () => { } }); - function createSession() { + async function createSession() { const model = getModel("anthropic", "claude-sonnet-4-5")!; let abortSignal: AbortSignal | undefined; @@ -156,16 +156,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - // Set a runtime API key so validation passes - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -195,15 +194,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -300,15 +299,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), baseToolsOverride: { dummy: tool }, }); diff --git a/packages/coding-agent/test/agent-session-concurrent-05.suite.ts b/packages/coding-agent/test/agent-session-concurrent-05.suite.ts index 3de2e9f1f..b8392de84 100644 --- a/packages/coding-agent/test/agent-session-concurrent-05.suite.ts +++ b/packages/coding-agent/test/agent-session-concurrent-05.suite.ts @@ -18,7 +18,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { convertToLlm } from "../src/core/messages.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import type { BuildSystemPromptOptions } from "../src/core/system-prompt.ts"; @@ -123,7 +123,7 @@ describe("AgentSession concurrent prompt guard", () => { } }); - function createSession() { + async function createSession() { const model = getModel("anthropic", "claude-sonnet-4-5")!; let abortSignal: AbortSignal | undefined; @@ -156,16 +156,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - // Set a runtime API key so validation passes - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); @@ -257,15 +256,15 @@ describe("AgentSession concurrent prompt guard", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), baseToolsOverride: { dummy: tool }, }); diff --git a/packages/coding-agent/test/agent-session-dynamic-provider.test.ts b/packages/coding-agent/test/agent-session-dynamic-provider.test.ts index ff454ece4..c891a37b5 100644 --- a/packages/coding-agent/test/agent-session-dynamic-provider.test.ts +++ b/packages/coding-agent/test/agent-session-dynamic-provider.test.ts @@ -31,7 +31,7 @@ describe("AgentSession dynamic provider registration", () => { const settingsManager = SettingsManager.create(tempDir, agentDir); const sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(join(agentDir, "auth.json")); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); const resourceLoader = new DefaultResourceLoader({ cwd: tempDir, agentDir, diff --git a/packages/coding-agent/test/agent-session-overflow-eviction.test.ts b/packages/coding-agent/test/agent-session-overflow-eviction.test.ts index 871e8fee5..c6c89c4ab 100644 --- a/packages/coding-agent/test/agent-session-overflow-eviction.test.ts +++ b/packages/coding-agent/test/agent-session-overflow-eviction.test.ts @@ -6,10 +6,10 @@ import { afterEach, beforeEach, describe, expect, it } from "vitest"; import { registerFauxProvider } from "@earendil-works/pi-ai/compat"; import { AgentSession, type AgentSessionEvent } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createTestResourceLoader } from "./utilities.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; interface AutoCompactionSurface { _runAutoCompaction(reason: "overflow" | "threshold", willRetry: boolean): Promise; @@ -26,7 +26,7 @@ describe("AgentSession auth-missing compaction failure semantics", () => { let unregister: (() => void) | undefined; let events: AgentSessionEvent[]; - beforeEach(() => { + beforeEach(async () => { tempDir = join(tmpdir(), `atomic-overflow-eviction-${Date.now()}`); mkdirSync(tempDir, { recursive: true }); events = []; @@ -38,13 +38,13 @@ describe("AgentSession auth-missing compaction failure semantics", () => { const settingsManager = SettingsManager.create(tempDir, tempDir); settingsManager.applyOverrides({ compaction: { reserveTokens: 20 } }); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); + const modelRegistry = await createModelRegistry(authStorage, tempDir); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime: getModelRuntime(modelRegistry), resourceLoader: createTestResourceLoader(), }); session.subscribe((event) => events.push(event)); @@ -125,7 +125,7 @@ describe("AgentSession auth-missing compaction failure semantics", () => { const end = events.find((event) => event.type === "compaction_end" && event.reason === "threshold"); expect(end).toMatchObject({ type: "compaction_end", reason: "threshold", result: undefined, aborted: false, willRetry: false }); if (end?.type !== "compaction_end") throw new Error("missing compaction_end"); - expect(end.errorMessage).toContain("Compaction provider authentication is unavailable"); + expect(end.errorMessage).toContain("No API key found for faux"); expect(sessionManager.getEntries().filter((entry) => entry.type === "compaction")).toHaveLength(0); }); }); diff --git a/packages/coding-agent/test/agent-session-retry.test.ts b/packages/coding-agent/test/agent-session-retry.test.ts index 8fb2e837c..13e61cd3b 100644 --- a/packages/coding-agent/test/agent-session-retry.test.ts +++ b/packages/coding-agent/test/agent-session-retry.test.ts @@ -1,15 +1,15 @@ import { existsSync, mkdirSync, rmSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; -import { Agent, type AgentEvent, type AgentTool } from "@earendil-works/pi-agent-core"; +import { Agent, type AgentEvent, type AgentTool, type StreamFn } from "@earendil-works/pi-agent-core"; import { type AssistantMessage, type AssistantMessageEvent, EventStream, getModel } from "@earendil-works/pi-ai/compat"; import { Type } from "typebox"; import { afterEach, beforeEach, describe, expect, it } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; import { createTestResourceLoader } from "./utilities.ts"; class MockAssistantStream extends EventStream { @@ -54,7 +54,7 @@ describe("AgentSession retry", () => { let session: AgentSession; let tempDir: string; - beforeEach(() => { + beforeEach(async () => { tempDir = join(tmpdir(), `pi-retry-test-${Date.now()}`); mkdirSync(tempDir, { recursive: true }); }); @@ -68,7 +68,11 @@ describe("AgentSession retry", () => { } }); - function createSession(options?: { failCount?: number; maxRetries?: number; delayAssistantMessageEndMs?: number }) { + async function createSession(options?: { + failCount?: number; + maxRetries?: number; + delayAssistantMessageEndMs?: number; + }) { const failCount = options?.failCount ?? 1; const maxRetries = options?.maxRetries ?? 3; const delayAssistantMessageEndMs = options?.delayAssistantMessageEndMs ?? 0; @@ -102,8 +106,8 @@ describe("AgentSession retry", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRegistry = await createModelRegistry(authStorage, tempDir); settingsManager.applyOverrides({ retry: { enabled: true, maxRetries, baseDelayMs: 1 } }); session = new AgentSession({ @@ -111,7 +115,7 @@ describe("AgentSession retry", () => { sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime: getModelRuntime(modelRegistry), resourceLoader: createTestResourceLoader(), }); @@ -129,8 +133,32 @@ describe("AgentSession retry", () => { return { session, getCallCount: () => callCount }; } + + async function createSessionWithStream(streamFn: StreamFn): Promise { + const model = getModel("anthropic", "claude-sonnet-4-5")!; + const agent = new Agent({ + getApiKey: () => "test-key", + initialState: { model, systemPrompt: "Test", tools: [] }, + streamFn, + }); + const sessionManager = SessionManager.inMemory(); + const settingsManager = SettingsManager.create(tempDir, tempDir); + const authStorage = AuthStorage.create(join(tempDir, "auth.json")); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRegistry = await createModelRegistry(authStorage, tempDir); + settingsManager.applyOverrides({ retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } }); + session = new AgentSession({ + agent, + sessionManager, + settingsManager, + cwd: tempDir, + modelRuntime: getModelRuntime(modelRegistry), + resourceLoader: createTestResourceLoader(), + }); + return session; + } it("retries after a transient error and succeeds", async () => { - const created = createSession({ failCount: 1 }); + const created = await createSession({ failCount: 1 }); const events: string[] = []; created.session.subscribe((event) => { if (event.type === "auto_retry_start") events.push(`start:${event.attempt}`); @@ -145,7 +173,7 @@ describe("AgentSession retry", () => { }); it("exhausts max retries and emits failure", async () => { - const created = createSession({ failCount: 99, maxRetries: 2 }); + const created = await createSession({ failCount: 99, maxRetries: 2 }); const events: string[] = []; created.session.subscribe((event) => { if (event.type === "auto_retry_start") events.push(`start:${event.attempt}`); @@ -162,7 +190,7 @@ describe("AgentSession retry", () => { }); it("prompt waits for retry completion even when assistant message_end handling is delayed", async () => { - const created = createSession({ failCount: 1, delayAssistantMessageEndMs: 40 }); + const created = await createSession({ failCount: 1, delayAssistantMessageEndMs: 40 }); await created.session.prompt("Test"); @@ -171,7 +199,7 @@ describe("AgentSession retry", () => { }); it("retries provider network_error failures", async () => { - const created = createSession({ failCount: 0 }); + const created = await createSession({ failCount: 0 }); let callCount = 0; const streamFn = () => { callCount++; @@ -199,20 +227,20 @@ describe("AgentSession retry", () => { const agent = new Agent({ getApiKey: () => "test-key", initialState: { model, systemPrompt: "Test", tools: [] }, - streamFn, + streamFn: streamFn, }); const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRegistry = await createModelRegistry(authStorage, tempDir); settingsManager.applyOverrides({ retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } }); session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime: getModelRuntime(modelRegistry), resourceLoader: createTestResourceLoader(), }); @@ -289,8 +317,8 @@ describe("AgentSession retry", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRegistry = await createModelRegistry(authStorage, tempDir); settingsManager.applyOverrides({ retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } }); session = new AgentSession({ @@ -298,7 +326,7 @@ describe("AgentSession retry", () => { sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime: getModelRuntime(modelRegistry), resourceLoader: createTestResourceLoader(), baseToolsOverride: { echo: echoTool }, }); @@ -317,109 +345,49 @@ describe("AgentSession retry", () => { }); it("retries bare provider finish_reason: error failures", async () => { - // github-copilot Gemini models surface MALFORMED_FUNCTION_CALL / OTHER / - // UNEXPECTED_TOOL_CALL as a bare finish_reason "error" (CAPI mapping), which - // pi-ai reports as "Provider finish_reason: error". This must be retryable. let callCount = 0; - const streamFn = () => { + await createSessionWithStream(() => { callCount++; const stream = new MockAssistantStream(); queueMicrotask(() => { if (callCount === 1) { - const msg = createAssistantMessage("", { + const error = createAssistantMessage("", { stopReason: "error", errorMessage: "Provider finish_reason: error", }); - stream.push({ type: "start", partial: msg }); - stream.push({ type: "error", reason: "error", error: msg }); + stream.push({ type: "start", partial: error }); + stream.push({ type: "error", reason: "error", error }); return; } - const msg = createAssistantMessage("Recovered after retry"); - stream.push({ type: "start", partial: msg }); - stream.push({ type: "done", reason: "stop", message: msg }); + const recovered = createAssistantMessage("Recovered after retry"); + stream.push({ type: "start", partial: recovered }); + stream.push({ type: "done", reason: "stop", message: recovered }); }); return stream; - }; - - const model = getModel("anthropic", "claude-sonnet-4-5")!; - const agent = new Agent({ - getApiKey: () => "test-key", - initialState: { model, systemPrompt: "Test", tools: [] }, - streamFn, - }); - const sessionManager = SessionManager.inMemory(); - const settingsManager = SettingsManager.create(tempDir, tempDir); - const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); - settingsManager.applyOverrides({ retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } }); - - session = new AgentSession({ - agent, - sessionManager, - settingsManager, - cwd: tempDir, - modelRegistry, - resourceLoader: createTestResourceLoader(), - }); - - const events: string[] = []; - session.subscribe((event) => { - if (event.type === "auto_retry_start") events.push(`start:${event.attempt}`); - if (event.type === "auto_retry_end") events.push(`end:success=${event.success}`); }); await session.prompt("Test"); - expect(callCount).toBe(2); - expect(events).toEqual(["start:1", "end:success=true"]); }); it("retries degenerate empty completions and bounds them by maxRetries", async () => { - // github-copilot Gemini models intermittently end the stream with finish_reason - // "stop", an empty content array, and 0 output tokens. That degenerate turn must - // be retried (not silently accepted), and the empty "stop" must NOT reset the - // retry counter, so repeated empties still honor maxRetries. let callCount = 0; - const streamFn = () => { + await createSessionWithStream(() => { callCount++; const stream = new MockAssistantStream(); queueMicrotask(() => { if (callCount <= 2) { - const msg = createAssistantMessage("", { content: [], stopReason: "stop" }); - stream.push({ type: "start", partial: msg }); - stream.push({ type: "done", reason: "stop", message: msg }); + const empty = createAssistantMessage("", { content: [], stopReason: "stop" }); + stream.push({ type: "start", partial: empty }); + stream.push({ type: "done", reason: "stop", message: empty }); return; } - const msg = createAssistantMessage("Recovered after empty completions"); - stream.push({ type: "start", partial: msg }); - stream.push({ type: "done", reason: "stop", message: msg }); + const recovered = createAssistantMessage("Recovered after empty completions"); + stream.push({ type: "start", partial: recovered }); + stream.push({ type: "done", reason: "stop", message: recovered }); }); return stream; - }; - - const model = getModel("anthropic", "claude-sonnet-4-5")!; - const agent = new Agent({ - getApiKey: () => "test-key", - initialState: { model, systemPrompt: "Test", tools: [] }, - streamFn, }); - const sessionManager = SessionManager.inMemory(); - const settingsManager = SettingsManager.create(tempDir, tempDir); - const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); - settingsManager.applyOverrides({ retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } }); - - session = new AgentSession({ - agent, - sessionManager, - settingsManager, - cwd: tempDir, - modelRegistry, - resourceLoader: createTestResourceLoader(), - }); - const events: string[] = []; session.subscribe((event) => { if (event.type === "auto_retry_start") events.push(`start:${event.attempt}`); @@ -427,63 +395,41 @@ describe("AgentSession retry", () => { }); await session.prompt("Test"); - - // Two empty completions -> two retries -> third call succeeds. expect(callCount).toBe(3); expect(events).toEqual(["start:1", "start:2", "end:success=true"]); - expect(session.isRetrying).toBe(false); }); - it("retries structured safety-trigger errors (content_filter, Anthropic refusal) for all providers", () => { - const created = createSession({ failCount: 0 }); - const probe = created.session as unknown as { - _isRetryableError(message: AssistantMessage): boolean; - }; - - // OpenAI-style `finish_reason: content_filter` (any provider). - const anthropicBlocked = createAssistantMessage("", { + it("retries structured safety-trigger errors for all providers", async () => { + const created = await createSession({ failCount: 0 }); + const probe = created.session as unknown as { _isRetryableError(message: AssistantMessage): boolean }; + expect(probe._isRetryableError(createAssistantMessage("", { stopReason: "error", errorMessage: "Provider finish_reason: content_filter", - }); - expect(probe._isRetryableError(anthropicBlocked)).toBe(true); - - // github-copilot CAPI mapping of spurious Gemini RECITATION/safety blocks. - const geminiBlocked = createAssistantMessage("", { + }))).toBe(true); + expect(probe._isRetryableError(createAssistantMessage("", { stopReason: "error", errorMessage: "Provider finish_reason: content_filter", provider: "github-copilot", api: "openai-completions", model: "gemini-3.1-pro-preview", - }); - expect(probe._isRetryableError(geminiBlocked)).toBe(true); - - // pi-ai's canned mapping of Anthropic `refusal` stops. - const anthropicRefusal = createAssistantMessage("", { + }))).toBe(true); + expect(probe._isRetryableError(createAssistantMessage("", { stopReason: "error", errorMessage: "The model refused to complete the request", - }); - expect(probe._isRetryableError(anthropicRefusal)).toBe(true); - - // Non-safety, non-transient errors stay non-retryable. - const badRequest = createAssistantMessage("", { + }))).toBe(true); + expect(probe._isRetryableError(createAssistantMessage("", { stopReason: "error", errorMessage: "Invalid request: unknown parameter", - }); - expect(probe._isRetryableError(badRequest)).toBe(false); + }))).toBe(false); }); - it("does not treat a reasoning-only turn (output > 0) as an empty completion", () => { - const created = createSession({ failCount: 0 }); - const probe = created.session as unknown as { - _isEmptyCompletion(message: AssistantMessage): boolean; - }; - - const emptyZeroOutput = createAssistantMessage("", { stopReason: "stop", content: [] }); - expect(probe._isEmptyCompletion(emptyZeroOutput)).toBe(true); - - const reasoningOnly = createAssistantMessage("", { - stopReason: "stop", + it("does not classify a reasoning-only turn with output tokens as empty", async () => { + const created = await createSession({ failCount: 0 }); + const probe = created.session as unknown as { _isEmptyCompletion(message: AssistantMessage): boolean }; + expect(probe._isEmptyCompletion(createAssistantMessage("", { content: [], stopReason: "stop" }))).toBe(true); + expect(probe._isEmptyCompletion(createAssistantMessage("", { content: [], + stopReason: "stop", usage: { input: 10, output: 5, @@ -492,8 +438,6 @@ describe("AgentSession retry", () => { totalTokens: 15, cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, }, - }); - expect(probe._isEmptyCompletion(reasoningOnly)).toBe(false); + }))).toBe(false); }); - }); diff --git a/packages/coding-agent/test/agent-session-runtime-events.test.ts b/packages/coding-agent/test/agent-session-runtime-events.test.ts index 348c04513..32bb828ff 100644 --- a/packages/coding-agent/test/agent-session-runtime-events.test.ts +++ b/packages/coding-agent/test/agent-session-runtime-events.test.ts @@ -10,6 +10,7 @@ import { createAgentSessionServices, } from "../src/core/agent-session-runtime.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import type { ExtensionFactory, @@ -42,11 +43,33 @@ describe("AgentSessionRuntime session lifecycle events", () => { faux.setResponses([fauxAssistantMessage("one"), fauxAssistantMessage("two"), fauxAssistantMessage("three")]); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await authStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); + const modelRuntime = await ModelRuntime.create({ + credentials: authStorage, + modelsPath: join(tempDir, "models.json"), + }); + const model = faux.getModel(); + modelRuntime.registerProvider(model.provider, { + baseUrl: model.baseUrl, + api: model.api, + models: [ + { + id: model.id, + name: model.name, + api: model.api, + reasoning: model.reasoning, + input: model.input, + cost: model.cost, + contextWindow: model.contextWindow, + maxTokens: model.maxTokens, + baseUrl: model.baseUrl, + }, + ], + }); const runtimeOptions = { agentDir: tempDir, - authStorage, + modelRuntime, model: faux.getModel(), resourceLoaderOptions: { extensionFactories: [extensionFactory], diff --git a/packages/coding-agent/test/agent-session-safety-refusal.test.ts b/packages/coding-agent/test/agent-session-safety-refusal.test.ts index d1344b898..526346d60 100644 --- a/packages/coding-agent/test/agent-session-safety-refusal.test.ts +++ b/packages/coding-agent/test/agent-session-safety-refusal.test.ts @@ -6,7 +6,7 @@ import { type AssistantMessage, type AssistantMessageEvent, EventStream, getMode import { afterEach, beforeEach, describe, expect, it } from "vitest"; import { AgentSession } from "../src/core/agent-session.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 { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createTestResourceLoader } from "./utilities.ts"; @@ -63,7 +63,7 @@ describe("AgentSession safety-refusal retry", () => { } }); - function createSession(streamFn: () => MockAssistantStream) { + async function createSession(streamFn: () => MockAssistantStream) { const model = getModel("anthropic", "claude-sonnet-4-5")!; const agent = new Agent({ getApiKey: () => "test-key", @@ -73,8 +73,8 @@ describe("AgentSession safety-refusal retry", () => { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, allowModelNetwork: false }); settingsManager.applyOverrides({ retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } }); session = new AgentSession({ @@ -82,7 +82,7 @@ describe("AgentSession safety-refusal retry", () => { sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); return session; @@ -95,7 +95,7 @@ describe("AgentSession safety-refusal retry", () => { // re-requested (not accepted as a final answer), and it must NOT reset the // retry counter, so repeated refusals still honor maxRetries (issue #1608). let callCount = 0; - createSession(() => { + await createSession(() => { callCount++; const stream = new MockAssistantStream(); queueMicrotask(() => { @@ -139,8 +139,8 @@ describe("AgentSession safety-refusal retry", () => { expect(assistantTexts).toEqual(["Recovered after canned refusals"]); }); - it("only detects tightly-guarded canned safety refusals", () => { - createSession(() => new MockAssistantStream()); + it("only detects tightly-guarded canned safety refusals", async () => { + await createSession(() => new MockAssistantStream()); const probe = session as unknown as { _isSafetyRefusal(message: AssistantMessage): boolean; }; diff --git a/packages/coding-agent/test/agent-session-services-model-paths.test.ts b/packages/coding-agent/test/agent-session-services-model-paths.test.ts index 8b4f68eae..7df429288 100644 --- a/packages/coding-agent/test/agent-session-services-model-paths.test.ts +++ b/packages/coding-agent/test/agent-session-services-model-paths.test.ts @@ -4,9 +4,9 @@ import { dirname, join } from "node:path"; import { afterEach, describe, expect, it } from "vitest"; import { createAgentSessionServices } from "../src/core/agent-session-services.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 { SettingsManager } from "../src/core/settings-manager.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; const originalHome = process.env.HOME; const originalUserProfile = process.env.USERPROFILE; @@ -30,22 +30,6 @@ function configureTemporaryHome(home: string): void { delete process.env.PI_CODING_AGENT_DIR; } -function writeLegacyOverride(home: string): void { - const legacyAgentDir = join(home, ".pi", "agent"); - mkdirSync(legacyAgentDir, { recursive: true }); - writeFileSync( - join(legacyAgentDir, "models.json"), - JSON.stringify({ - providers: { - openrouter: { - modelOverrides: { - "anthropic/claude-sonnet-4": { name: "Legacy startup override" }, - }, - }, - }, - }), - ); -} function writeCustomModels(path: string, models: Array<{ id: string; name?: string }>): void { mkdirSync(dirname(path), { recursive: true }); @@ -76,96 +60,16 @@ afterEach(() => { }); describe("agent session service model paths", () => { - it("loads legacy and primary models.json layers during normal CLI startup", async () => { - const home = mkdtempSync(join(tmpdir(), "atomic-service-model-paths-")); - tempDirs.push(home); - configureTemporaryHome(home); - writeLegacyOverride(home); - const agentDir = join(home, ".atomic", "agent"); - mkdirSync(agentDir, { recursive: true }); - const services = await createAgentSessionServices({ - cwd: home, - agentDir, - authStorage: AuthStorage.inMemory(), - settingsManager: SettingsManager.inMemory(), - }); - - expect(services.modelRegistry.find("openrouter", "anthropic/claude-sonnet-4")?.name).toBe( - "Legacy startup override", - ); - }); - - it("loads project .atomic/models.json during normal CLI startup", async () => { + it("does not load project-scoped models for a trusted project", async () => { const home = mkdtempSync(join(tmpdir(), "atomic-service-project-models-")); tempDirs.push(home); configureTemporaryHome(home); const cwd = join(home, "project"); - mkdirSync(cwd); - writeCustomModel(join(cwd, ".atomic", "models.json"), "atomic-1995-reload-probe"); - - const services = await createAgentSessionServices({ - cwd, - agentDir: join(home, ".atomic", "agent"), - authStorage: AuthStorage.inMemory(), - settingsManager: SettingsManager.inMemory(), - resourceLoaderOptions: { noExtensions: true, noSkills: true, noPromptTemplates: true, noThemes: true }, - }); - - expect(services.modelRegistry.find("project-probe", "atomic-1995-reload-probe")?.name).toBe( - "atomic-1995-reload-probe", - ); - }); - - it("layers Pi global/project before Atomic global/project and reloads raw IDs verbatim", async () => { - const home = mkdtempSync(join(tmpdir(), "atomic-service-layered-models-")); - tempDirs.push(home); - configureTemporaryHome(home); - const cwd = join(home, "project"); - const agentDir = join(home, ".atomic", "agent"); - mkdirSync(cwd); - const sharedId = "shared[raw]-model"; - const layers = [ - [join(home, ".pi", "agent", "models.json"), "pi-global", "Pi global"], - [join(cwd, ".pi", "models.json"), "pi-project", "Pi project"], - [join(agentDir, "models.json"), "atomic-global", "Atomic global"], - [join(cwd, ".atomic", "models.json"), "atomic-project", "Atomic project"], - ] as const; - for (const [path, uniqueId, sharedName] of layers) { - writeCustomModels(path, [{ id: uniqueId }, { id: sharedId, name: sharedName }]); - } - - const services = await createAgentSessionServices({ - cwd, - agentDir, - authStorage: AuthStorage.inMemory(), - settingsManager: SettingsManager.inMemory(), - resourceLoaderOptions: { noExtensions: true, noSkills: true, noPromptTemplates: true, noThemes: true }, - }); - expect(layers.map(([, id]) => services.modelRegistry.find("project-probe", id)?.id)).toEqual( - layers.map(([, id]) => id), - ); - expect(services.modelRegistry.find("project-probe", sharedId)?.name).toBe("Atomic project"); - - writeCustomModels(join(cwd, ".atomic", "models.json"), [ - { id: "atomic-project-after" }, - { id: sharedId, name: "Atomic project after reload" }, - ]); - await services.modelRegistry.refresh({ allowNetwork: false }); - expect(services.modelRegistry.find("project-probe", "atomic-project")).toBeUndefined(); - expect(services.modelRegistry.find("project-probe", "atomic-project-after")?.id).toBe("atomic-project-after"); - expect(services.modelRegistry.find("project-probe", sharedId)?.name).toBe("Atomic project after reload"); - }); - - it("does not load project models while the project is untrusted", async () => { - const home = mkdtempSync(join(tmpdir(), "atomic-service-untrusted-models-")); - tempDirs.push(home); - configureTemporaryHome(home); - const cwd = join(home, "project"); const agentDir = join(home, ".atomic", "agent"); mkdirSync(cwd); - writeCustomModel(join(cwd, ".atomic", "models.json"), "untrusted-project-model"); - const settingsManager = SettingsManager.create(cwd, agentDir, { projectTrusted: false }); + writeCustomModel(join(cwd, ".atomic", "models.json"), "project-scoped-only"); + const settingsManager = SettingsManager.create(cwd, agentDir, { projectTrusted: true }); const services = await createAgentSessionServices({ cwd, @@ -175,31 +79,7 @@ describe("agent session service model paths", () => { resourceLoaderOptions: { noExtensions: true, noSkills: true, noPromptTemplates: true, noThemes: true }, }); - expect(services.modelRegistry.find("project-probe", "untrusted-project-model")).toBeUndefined(); - }); - - it("loads project models after startup trust resolution accepts the project", async () => { - const home = mkdtempSync(join(tmpdir(), "atomic-service-trusted-models-")); - tempDirs.push(home); - configureTemporaryHome(home); - const cwd = join(home, "project"); - const agentDir = join(home, ".atomic", "agent"); - mkdirSync(cwd); - writeCustomModel(join(cwd, ".atomic", "models.json"), "newly-trusted-project-model"); - const settingsManager = SettingsManager.create(cwd, agentDir, { projectTrusted: false }); - - const services = await createAgentSessionServices({ - cwd, - agentDir, - authStorage: AuthStorage.inMemory(), - settingsManager, - resourceLoaderOptions: { noExtensions: true, noSkills: true, noPromptTemplates: true, noThemes: true }, - resourceLoaderReloadOptions: { resolveProjectTrust: async () => true }, - }); - - expect(services.modelRegistry.find("project-probe", "newly-trusted-project-model")?.id).toBe( - "newly-trusted-project-model", - ); + expect(services.modelRuntime.getModel("project-probe", "project-scoped-only")).toBeUndefined(); }); it("preserves an explicitly supplied model registry without adding project paths", async () => { @@ -211,19 +91,19 @@ describe("agent session service model paths", () => { const explicitPath = join(home, "explicit-models.json"); writeCustomModel(explicitPath, "explicit-only"); writeCustomModel(join(cwd, ".atomic", "models.json"), "project-only"); - const explicitRegistry = ModelRegistry.create(AuthStorage.inMemory(), explicitPath); + const explicitRegistry = await createModelRegistry(AuthStorage.inMemory(), explicitPath); const services = await createAgentSessionServices({ cwd, agentDir: join(home, ".atomic", "agent"), - modelRegistry: explicitRegistry, + modelRuntime: getModelRuntime(explicitRegistry), settingsManager: SettingsManager.inMemory(), resourceLoaderOptions: { noExtensions: true, noSkills: true, noPromptTemplates: true, noThemes: true }, }); - expect(services.modelRegistry).toBe(explicitRegistry); - expect(services.modelRegistry.find("project-probe", "explicit-only")?.id).toBe("explicit-only"); - expect(services.modelRegistry.find("project-probe", "project-only")).toBeUndefined(); + expect(services.modelRuntime).toBe(getModelRuntime(explicitRegistry)); + expect(services.modelRuntime.getModel("project-probe", "explicit-only")?.id).toBe("explicit-only"); + expect(services.modelRuntime.getModel("project-probe", "project-only")).toBeUndefined(); }); it("keeps an API-supplied modelsPath isolated from default project layers", async () => { @@ -237,7 +117,7 @@ describe("agent session service model paths", () => { writeCustomModel(join(cwd, ".atomic", "models.json"), "project-default-only"); const runtime = await ModelRuntime.create({ - authStorage: AuthStorage.inMemory(), + credentials: AuthStorage.inMemory(), modelsPath: explicitPath, allowModelNetwork: false, }); @@ -246,29 +126,4 @@ describe("agent session service model paths", () => { expect(runtime.getModel("project-probe", "project-default-only")).toBeUndefined(); }); - it("keeps a custom agent directory isolated while still layering project models", async () => { - const home = mkdtempSync(join(tmpdir(), "atomic-service-custom-model-paths-")); - tempDirs.push(home); - configureTemporaryHome(home); - writeLegacyOverride(home); - const customAgentDir = join(home, "custom-agent"); - const cwd = join(home, "project"); - mkdirSync(cwd); - writeCustomModel(join(customAgentDir, "models.json"), "custom-global"); - writeCustomModel(join(cwd, ".atomic", "models.json"), "custom-project"); - - const services = await createAgentSessionServices({ - cwd, - agentDir: customAgentDir, - authStorage: AuthStorage.inMemory(), - settingsManager: SettingsManager.inMemory(), - resourceLoaderOptions: { noExtensions: true, noSkills: true, noPromptTemplates: true, noThemes: true }, - }); - - expect(services.modelRegistry.find("project-probe", "custom-global")?.id).toBe("custom-global"); - expect(services.modelRegistry.find("project-probe", "custom-project")?.id).toBe("custom-project"); - expect(services.modelRegistry.find("openrouter", "anthropic/claude-sonnet-4")?.name).not.toBe( - "Legacy startup override", - ); - }); }); diff --git a/packages/coding-agent/test/agent-session-stats.test.ts b/packages/coding-agent/test/agent-session-stats.test.ts index 3a121b15f..b8358156c 100644 --- a/packages/coding-agent/test/agent-session-stats.test.ts +++ b/packages/coding-agent/test/agent-session-stats.test.ts @@ -1,11 +1,17 @@ import { Agent } from "@earendil-works/pi-agent-core"; -import { type AssistantMessage, getModel, type Usage } from "@earendil-works/pi-ai/compat"; +import { + type AssistantMessage, + getModel, + streamSimple, + type ToolResultMessage, + type Usage, +} from "@earendil-works/pi-ai/compat"; import { describe, expect, it } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; +import { createInMemoryModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; import { createTestResourceLoader } from "./utilities.ts"; import { appendTestCompaction } from "./verbatim-compaction-test-helpers.ts"; @@ -53,15 +59,27 @@ function createUserMessage(text: string, timestamp: number) { }; } +function createToolResultMessage(usage: Usage): ToolResultMessage { + return { + role: "toolResult", + toolCallId: "tool-call-1", + toolName: "test_tool", + content: [{ type: "text", text: "tool result" }], + usage, + isError: false, + timestamp: 1, + }; +} -function createSession() { +async function createSession() { const settingsManager = SettingsManager.inMemory(); const sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); const session = new AgentSession({ agent: new Agent({ getApiKey: () => "test-key", + streamFn: streamSimple, initialState: { model, systemPrompt: "You are a helpful assistant.", @@ -72,7 +90,7 @@ function createSession() { sessionManager, settingsManager, cwd: process.cwd(), - modelRegistry: ModelRegistry.inMemory(authStorage), + modelRuntime: getModelRuntime(await createInMemoryModelRegistry(authStorage)), resourceLoader: createTestResourceLoader(), }); @@ -84,8 +102,8 @@ function syncAgentMessages(session: AgentSession, sessionManager: SessionManager } describe("AgentSession.getSessionStats", () => { - it("exposes the current context usage alongside token totals", () => { - const { session, sessionManager } = createSession(); + it("exposes the current context usage alongside token totals", async () => { + const { session, sessionManager } = await createSession(); try { sessionManager.appendMessage(createUserMessage("hello", 1)); @@ -102,8 +120,8 @@ describe("AgentSession.getSessionStats", () => { } }); - it("reports unknown current context usage immediately after context compaction", () => { - const { session, sessionManager } = createSession(); + it("reports unknown current context usage immediately after compaction", async () => { + const { session, sessionManager } = await createSession(); try { sessionManager.appendMessage(createUserMessage("first", 1)); @@ -111,10 +129,12 @@ describe("AgentSession.getSessionStats", () => { sessionManager.appendMessage(createUserMessage("second", 3)); sessionManager.appendMessage(createAssistantMessage("response2", 195_000, 4)); appendTestCompaction(sessionManager, 195_000, 50_000); - sessionManager.appendMessage(createUserMessage("third", Date.now() + 1)); + sessionManager.appendMessage(createUserMessage("third", 5)); syncAgentMessages(session, sessionManager); const stats = session.getSessionStats(); + // Totals cover ALL entries, including history compacted away (180k + 195k). + expect(stats.tokens.input).toBe(375_000); expect(stats.contextUsage).toBeDefined(); expect(stats.contextUsage?.tokens).toBeNull(); expect(stats.contextUsage?.percent).toBeNull(); @@ -123,8 +143,8 @@ describe("AgentSession.getSessionStats", () => { } }); - it("uses post-context-compaction usage for current context instead of stale kept usage (duplicate check)", () => { - const { session, sessionManager } = createSession(); + it("uses post-compaction usage for current context instead of stale kept usage", async () => { + const { session, sessionManager } = await createSession(); try { sessionManager.appendMessage(createUserMessage("first", 1)); @@ -137,6 +157,8 @@ describe("AgentSession.getSessionStats", () => { syncAgentMessages(session, sessionManager); const stats = session.getSessionStats(); + // Totals cover ALL entries, including history compacted away (180k + 195k + 25k). + expect(stats.tokens.input).toBe(400_000); expect(stats.contextUsage).toBeDefined(); expect(stats.contextUsage?.tokens).toBe(25_000); expect(stats.contextUsage?.percent).toBe((25_000 / model.contextWindow) * 100); @@ -145,53 +167,62 @@ describe("AgentSession.getSessionStats", () => { } }); - it("reports unknown current context usage immediately after context compaction", () => { - const { session, sessionManager } = createSession(); + it("includes branch summary usage in session totals", async () => { + const { session, sessionManager } = await createSession(); try { - sessionManager.appendMessage(createUserMessage("first", 1)); - sessionManager.appendMessage(createAssistantMessage("response1", 195_000, 2)); - appendTestCompaction(sessionManager, 195_000, 50_000); - sessionManager.appendMessage(createUserMessage("second", Date.now() + 1)); + sessionManager.branchWithSummary(null, "summary", undefined, false, { + input: 10, + output: 20, + cacheRead: 30, + cacheWrite: 40, + totalTokens: 100, + cost: { input: 0.1, output: 0.2, cacheRead: 0.3, cacheWrite: 0.4, total: 1 }, + }); syncAgentMessages(session, sessionManager); const stats = session.getSessionStats(); - expect(stats.contextUsage).toBeDefined(); - expect(stats.contextUsage?.tokens).toBeNull(); - expect(stats.contextUsage?.percent).toBeNull(); + expect(stats.tokens).toEqual({ input: 10, output: 20, cacheRead: 30, cacheWrite: 40, total: 100 }); + expect(stats.cost).toBe(1); } finally { session.dispose(); } }); - it("uses post-context-compaction usage for current context instead of stale kept usage", () => { - const { session, sessionManager } = createSession(); + + it("includes tool result usage in session totals", async () => { + const { session, sessionManager } = await createSession(); try { - sessionManager.appendMessage(createUserMessage("first", 1)); - sessionManager.appendMessage(createAssistantMessage("response1", 195_000, 2)); - appendTestCompaction(sessionManager, 195_000, 50_000); - sessionManager.appendMessage(createUserMessage("second", Date.now() + 1)); - sessionManager.appendMessage(createAssistantMessage("response2", 25_000, Date.now() + 2)); + sessionManager.appendMessage( + createToolResultMessage({ + input: 10, + output: 20, + cacheRead: 30, + cacheWrite: 40, + totalTokens: 100, + cost: { input: 0.1, output: 0.2, cacheRead: 0.3, cacheWrite: 0.4, total: 1 }, + }), + ); syncAgentMessages(session, sessionManager); const stats = session.getSessionStats(); - expect(stats.contextUsage).toBeDefined(); - expect(stats.contextUsage?.tokens).toBe(25_000); - expect(stats.contextUsage?.percent).toBe((25_000 / model.contextWindow) * 100); + expect(stats.tokens).toEqual({ input: 10, output: 20, cacheRead: 30, cacheWrite: 40, total: 100 }); + expect(stats.cost).toBe(1); } finally { session.dispose(); } }); - it("does not double-count mirrored cache buckets in post-compaction context usage", () => { - const { session, sessionManager } = createSession(); + it("does not double-count mirrored cache buckets in post-compaction context usage", async () => { + const { session, sessionManager } = await createSession(); try { sessionManager.appendMessage(createUserMessage("first", 1)); sessionManager.appendMessage(createAssistantMessage("response1", 195_000, 2)); appendTestCompaction(sessionManager, 216_006, 81_414); - sessionManager.appendMessage(createUserMessage("second", Date.now() + 1)); + const postCompaction = Date.now() + 1; + sessionManager.appendMessage(createUserMessage("second", postCompaction)); sessionManager.appendMessage( createAssistantMessageWithUsage( "response2", @@ -203,13 +234,12 @@ describe("AgentSession.getSessionStats", () => { totalTokens: 232_500, cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, }, - Date.now() + 2, + postCompaction + 1, ), ); syncAgentMessages(session, sessionManager); const stats = session.getSessionStats(); - expect(stats.contextUsage).toBeDefined(); expect(stats.contextUsage?.tokens).toBe(116_500); expect(stats.contextUsage?.percent).toBe((116_500 / model.contextWindow) * 100); } finally { @@ -217,18 +247,27 @@ describe("AgentSession.getSessionStats", () => { } }); - it("counts normalized Codex cache partitions in post-compaction context usage", () => { - const { session, sessionManager } = createSession(); + it("counts normalized Codex cache partitions in post-compaction context usage", async () => { + const { session, sessionManager } = await createSession(); try { sessionManager.appendMessage(createUserMessage("first", 1)); sessionManager.appendMessage(createAssistantMessage("response1", 195_000, 2)); appendTestCompaction(sessionManager, 195_000, 50_000); - sessionManager.appendMessage(createUserMessage("second", Date.now() + 1)); + const postCompaction = Date.now() + 1; + sessionManager.appendMessage(createUserMessage("second", postCompaction)); sessionManager.appendMessage({ - ...createAssistantMessageWithUsage("response2", { - input: 7_907, output: 7, cacheRead: 7_936, cacheWrite: 0, totalTokens: 15_850, - cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, - }, Date.now() + 2), + ...createAssistantMessageWithUsage( + "response2", + { + input: 7_907, + output: 7, + cacheRead: 7_936, + cacheWrite: 0, + totalTokens: 15_850, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + postCompaction + 1, + ), api: "openai-codex-responses", provider: "openai-codex", model: "gpt-5.5", diff --git a/packages/coding-agent/test/agent-session-tree-navigation.test.ts b/packages/coding-agent/test/agent-session-tree-navigation.test.ts index 258d11ae7..3678b8ce5 100644 --- a/packages/coding-agent/test/agent-session-tree-navigation.test.ts +++ b/packages/coding-agent/test/agent-session-tree-navigation.test.ts @@ -15,8 +15,8 @@ import { API_KEY, createTestSession, type TestSessionContext } from "./utilities describe.skipIf(!API_KEY)("AgentSession tree navigation e2e", () => { let ctx: TestSessionContext; - beforeEach(() => { - ctx = createTestSession({ + beforeEach(async () => { + ctx = await createTestSession({ systemPrompt: "You are a helpful assistant. Reply with just a few words.", settingsOverrides: { compaction: { reserveTokens: 1 } }, }); @@ -279,8 +279,8 @@ describe.skipIf(!API_KEY)("AgentSession tree navigation e2e", () => { describe.skipIf(!API_KEY)("AgentSession tree navigation - branch scenarios", () => { let ctx: TestSessionContext; - beforeEach(() => { - ctx = createTestSession({ + beforeEach(async () => { + ctx = await createTestSession({ systemPrompt: "You are a helpful assistant. Reply with just a few words.", }); }); diff --git a/packages/coding-agent/test/auth-storage-01.suite.ts b/packages/coding-agent/test/auth-storage-01.suite.ts index 45c6c0eae..9115cf4d0 100644 --- a/packages/coding-agent/test/auth-storage-01.suite.ts +++ b/packages/coding-agent/test/auth-storage-01.suite.ts @@ -1,498 +1,207 @@ -import { existsSync, mkdirSync, readdirSync, readFileSync, rmSync, statSync, writeFileSync } from "node:fs"; +import { existsSync, mkdirSync, readFileSync, readdirSync, rmSync, statSync, writeFileSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; -import { registerOAuthProvider } from "../src/core/oauth-provider-bridge.ts"; import lockfile from "proper-lockfile"; import { afterEach, beforeEach, describe, expect, test, vi } from "vitest"; -import { AuthStorage, type AuthStorageBackend, FileAuthStorageBackend } from "../src/core/auth-storage.ts"; -import { clearConfigValueCache } from "../src/core/resolve-config-value.ts"; - -/** - * Backend whose every access throws — used to simulate a credential-store load - * failure (e.g. ELOCKED under concurrent auth.json access) so a *fresh* - * AuthStorage ends up with an empty in-memory store and a recorded loadError - * (issue #1431). - */ -class ThrowingAuthStorageBackend implements AuthStorageBackend { - constructor(private readonly error: Error) {} - read(): string | undefined { - throw this.error; - } - withLock(): T { - throw this.error; - } - async withLockAsync(): Promise { - throw this.error; - } -} -describe("AuthStorage", () => { +import { AuthStorage, FileAuthStorageBackend } from "../src/core/auth-storage.ts"; +import { + clearConfigValueCache, + resolveConfigValue, + resolveConfigValueUncached, +} from "../src/core/resolve-config-value.ts"; +import * as shellModule from "../src/utils/shell.ts"; + +describe("resolveConfigValue", () => { let tempDir: string; - let authJsonPath: string; - let authStorage: AuthStorage; beforeEach(() => { - tempDir = join(tmpdir(), `pi-test-auth-storage-${Date.now()}-${Math.random().toString(36).slice(2)}`); + tempDir = join(tmpdir(), `pi-config-value-${Date.now()}-${Math.random().toString(36).slice(2)}`); mkdirSync(tempDir, { recursive: true }); - authJsonPath = join(tempDir, "auth.json"); + clearConfigValueCache(); }); afterEach(() => { - if (tempDir && existsSync(tempDir)) { - rmSync(tempDir, { recursive: true }); - } + if (existsSync(tempDir)) rmSync(tempDir, { recursive: true }); clearConfigValueCache(); vi.restoreAllMocks(); }); - function writeAuthJson(data: Record) { - writeFileSync(authJsonPath, JSON.stringify(data)); - } - - function toShPath(value: string): string { - let escaped = ""; - for (const char of value.replace(/\\/g, "/")) { - if (char === '"' || char === "\\" || char === "$" || char === "`") { - escaped += `\\${char}`; - } else { - escaped += char; - } + test("resolves literals, environment templates, and escapes", () => { + process.env.TEST_CONFIG_LEFT = "left"; + process.env.TEST_CONFIG_RIGHT = "right"; + try { + expect(resolveConfigValue("literal-key")).toBe("literal-key"); + expect(resolveConfigValue("$TEST_CONFIG_LEFT")).toBe("left"); + expect(resolveConfigValue("$" + "{TEST_CONFIG_LEFT}_$TEST_CONFIG_RIGHT")).toBe("left_right"); + expect(resolveConfigValue("$$TEST_CONFIG_LEFT")).toBe("$TEST_CONFIG_LEFT"); + expect(resolveConfigValue("$!literal-$TEST_CONFIG_RIGHT")).toBe("!literal-right"); + } finally { + delete process.env.TEST_CONFIG_LEFT; + delete process.env.TEST_CONFIG_RIGHT; } - return escaped; - } - - describe("API key resolution", () => { - test("literal API key is returned directly", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "sk-ant-literal-key" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBe("sk-ant-literal-key"); - }); - test("apiKey with ! prefix executes command and uses stdout", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!echo test-api-key-from-command" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBe("test-api-key-from-command"); - }); - test("apiKey with ! prefix trims whitespace from command output", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!echo ' spaced-key '" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBe("spaced-key"); - }); - test("apiKey with ! prefix handles multiline output (uses trimmed result)", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!printf 'line1\\nline2'" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBe("line1\nline2"); - }); - test("apiKey with ! prefix returns undefined on command failure", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!exit 1" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBeUndefined(); - }); - test("apiKey with ! prefix returns undefined on nonexistent command", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!nonexistent-command-12345" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBeUndefined(); - }); - test("apiKey with ! prefix returns undefined on empty output", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!printf ''" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBeUndefined(); - }); - test("apiKey with $ prefix resolves to env value", async () => { - const originalEnv = process.env.TEST_AUTH_API_KEY_12345; - process.env.TEST_AUTH_API_KEY_12345 = "env-api-key-value"; - - try { - writeAuthJson({ - anthropic: { type: "api_key", key: "$TEST_AUTH_API_KEY_12345" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBe("env-api-key-value"); - } finally { - if (originalEnv === undefined) { - delete process.env.TEST_AUTH_API_KEY_12345; - } else { - process.env.TEST_AUTH_API_KEY_12345 = originalEnv; - } - } - }); - test("apiKey as literal value is used directly when not an env var", async () => { - // Make sure this isn't an env var - delete process.env.literal_api_key_value; - - writeAuthJson({ - anthropic: { type: "api_key", key: "literal_api_key_value" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBe("literal_api_key_value"); - }); - test("apiKey command can use shell features like pipes", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!echo 'hello world' | tr ' ' '-'" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - const apiKey = await authStorage.getApiKey("anthropic"); - - expect(apiKey).toBe("hello-world"); - }); - test("command is only executed once per process", async () => { - // Use a command that writes to a file to count invocations - const counterFile = join(tempDir, "counter"); - writeFileSync(counterFile, "0"); - - const counterPath = toShPath(counterFile); - const command = `!sh -c 'count=$(cat "${counterPath}"); echo $((count + 1)) > "${counterPath}"; echo "key-value"'`; - writeAuthJson({ - anthropic: { type: "api_key", key: command }, - }); - - authStorage = AuthStorage.create(authJsonPath); - - // Call multiple times - await authStorage.getApiKey("anthropic"); - await authStorage.getApiKey("anthropic"); - await authStorage.getApiKey("anthropic"); - - // Command should have only run once - const count = parseInt(readFileSync(counterFile, "utf-8").trim(), 10); - expect(count).toBe(1); - }); - test("cache persists across AuthStorage instances", async () => { - const counterFile = join(tempDir, "counter"); - writeFileSync(counterFile, "0"); - - const counterPath = toShPath(counterFile); - const command = `!sh -c 'count=$(cat "${counterPath}"); echo $((count + 1)) > "${counterPath}"; echo "key-value"'`; - writeAuthJson({ - anthropic: { type: "api_key", key: command }, - }); - - // Create multiple AuthStorage instances - const storage1 = AuthStorage.create(authJsonPath); - await storage1.getApiKey("anthropic"); - - const storage2 = AuthStorage.create(authJsonPath); - await storage2.getApiKey("anthropic"); - - // Command should still have only run once - const count = parseInt(readFileSync(counterFile, "utf-8").trim(), 10); - expect(count).toBe(1); - }); - test("clearConfigValueCache allows command to run again", async () => { - const counterFile = join(tempDir, "counter"); - writeFileSync(counterFile, "0"); - - const counterPath = toShPath(counterFile); - const command = `!sh -c 'count=$(cat "${counterPath}"); echo $((count + 1)) > "${counterPath}"; echo "key-value"'`; - writeAuthJson({ - anthropic: { type: "api_key", key: command }, - }); - - authStorage = AuthStorage.create(authJsonPath); - await authStorage.getApiKey("anthropic"); - - // Clear cache and call again - clearConfigValueCache(); - await authStorage.getApiKey("anthropic"); - - // Command should have run twice - const count = parseInt(readFileSync(counterFile, "utf-8").trim(), 10); - expect(count).toBe(2); - }); - test("different commands are cached separately", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!echo key-anthropic" }, - openai: { type: "api_key", key: "!echo key-openai" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - - const keyA = await authStorage.getApiKey("anthropic"); - const keyB = await authStorage.getApiKey("openai"); - - expect(keyA).toBe("key-anthropic"); - expect(keyB).toBe("key-openai"); - }); - test("failed commands are cached (not retried)", async () => { - const counterFile = join(tempDir, "counter"); - writeFileSync(counterFile, "0"); - - const counterPath = toShPath(counterFile); - const command = `!sh -c 'count=$(cat "${counterPath}"); echo $((count + 1)) > "${counterPath}"; exit 1'`; - writeAuthJson({ - anthropic: { type: "api_key", key: command }, - }); - - authStorage = AuthStorage.create(authJsonPath); - - // Call multiple times - all should return undefined - const key1 = await authStorage.getApiKey("anthropic"); - const key2 = await authStorage.getApiKey("anthropic"); + }); - expect(key1).toBeUndefined(); - expect(key2).toBeUndefined(); + test("uses credential-scoped environment before process.env", () => { + process.env.TEST_CONFIG_SCOPED = "process"; + try { + expect(resolveConfigValue("$TEST_CONFIG_SCOPED", { TEST_CONFIG_SCOPED: "credential" })).toBe("credential"); + } finally { + delete process.env.TEST_CONFIG_SCOPED; + } + }); - // Command should have only run once despite failures - const count = parseInt(readFileSync(counterFile, "utf-8").trim(), 10); - expect(count).toBe(1); - }); - test("environment variables are not cached (changes are picked up)", async () => { - const envVarName = "TEST_AUTH_KEY_CACHE_TEST_98765"; - const originalEnv = process.env[envVarName]; + test("executes shell commands and trims their output", () => { + expect(resolveConfigValue("!echo ' spaced-key '")).toBe("spaced-key"); + expect(resolveConfigValue("!printf 'line1\\nline2'")).toBe("line1\nline2"); + expect(resolveConfigValue("!echo 'hello world' | tr ' ' '-'")).toBe("hello-world"); + }); - try { - process.env[envVarName] = "first-value"; + test.each(["!exit 1", "!nonexistent-command-12345", "!printf ''"])( + "returns undefined when command resolution fails: %s", + (command) => { + expect(resolveConfigValue(command)).toBeUndefined(); + }, + ); - writeAuthJson({ - anthropic: { type: "api_key", key: `$${envVarName}` }, - }); + test("caches successful and failed commands until explicitly cleared", () => { + const counterFile = join(tempDir, "counter"); + writeFileSync(counterFile, "0"); + const escapedPath = counterFile.replace(/\\/g, "/").replace(/"/g, '\\"'); + const success = `!sh -c 'count=$(cat "${escapedPath}"); echo $((count + 1)) > "${escapedPath}"; echo value'`; - authStorage = AuthStorage.create(authJsonPath); + expect(resolveConfigValue(success)).toBe("value"); + expect(resolveConfigValue(success)).toBe("value"); + expect(readFileSync(counterFile, "utf-8").trim()).toBe("1"); - const key1 = await authStorage.getApiKey("anthropic"); - expect(key1).toBe("first-value"); + clearConfigValueCache(); + expect(resolveConfigValue(success)).toBe("value"); + expect(readFileSync(counterFile, "utf-8").trim()).toBe("2"); - // Change env var - process.env[envVarName] = "second-value"; + const failure = `!sh -c 'count=$(cat "${escapedPath}"); echo $((count + 1)) > "${escapedPath}"; exit 1'`; + expect(resolveConfigValue(failure)).toBeUndefined(); + expect(resolveConfigValue(failure)).toBeUndefined(); + expect(readFileSync(counterFile, "utf-8").trim()).toBe("3"); + }); - const key2 = await authStorage.getApiKey("anthropic"); - expect(key2).toBe("second-value"); - } finally { - if (originalEnv === undefined) { - delete process.env[envVarName]; - } else { - process.env[envVarName] = originalEnv; - } - } - }); - test("returns undefined on compromised lock and allows a later retry", async () => { - const providerId = `test-oauth-provider-${Date.now()}-${Math.random().toString(36).slice(2)}`; - registerOAuthProvider({ - id: providerId, - name: "Test OAuth Provider", - async login() { - throw new Error("Not used in this test"); - }, - async refreshToken(credentials) { - return { - ...credentials, - access: "refreshed-access-token", - expires: Date.now() + 60_000, - }; - }, - getApiKey(credentials) { - return `Bearer ${credentials.access}`; - }, - }); + test("does not cache environment values", () => { + process.env.TEST_CONFIG_DYNAMIC = "first"; + try { + expect(resolveConfigValue("$TEST_CONFIG_DYNAMIC")).toBe("first"); + process.env.TEST_CONFIG_DYNAMIC = "second"; + expect(resolveConfigValue("$TEST_CONFIG_DYNAMIC")).toBe("second"); + } finally { + delete process.env.TEST_CONFIG_DYNAMIC; + } + }); - writeAuthJson({ - [providerId]: { - type: "oauth", - refresh: "refresh-token", - access: "expired-access-token", - expires: Date.now() - 10_000, - }, - }); + test("uncached resolution executes a command on every call", () => { + const counterFile = join(tempDir, "uncached-counter"); + writeFileSync(counterFile, "0"); + const escapedPath = counterFile.replace(/\\/g, "/").replace(/"/g, '\\"'); + const command = `!sh -c 'count=$(cat "${escapedPath}"); echo $((count + 1)) > "${escapedPath}"; echo value'`; + expect(resolveConfigValueUncached(command)).toBe("value"); + expect(resolveConfigValueUncached(command)).toBe("value"); + expect(readFileSync(counterFile, "utf-8").trim()).toBe("2"); + }); - authStorage = AuthStorage.create(authJsonPath); + test("uses stdin when the configured Windows shell requires it", () => { + if (process.platform === "win32") return; + const platformDescriptor = Object.getOwnPropertyDescriptor(process, "platform"); + vi.spyOn(shellModule, "getShellConfig").mockReturnValue({ + shell: "/bin/bash", + args: ["-s"], + commandTransport: "stdin", + }); + try { + Object.defineProperty(process, "platform", { configurable: true, value: "win32" }); + const expansion = "$" + "{name}"; + expect(resolveConfigValueUncached(`!name='World'; echo "Hello, ${expansion}!"`)).toBe("Hello, World!"); + } finally { + if (platformDescriptor) Object.defineProperty(process, "platform", platformDescriptor); + } + }); +}); - const realLock = lockfile.lock.bind(lockfile); - const lockSpy = vi.spyOn(lockfile, "lock"); - lockSpy.mockImplementationOnce(async (file, options) => { - options?.onCompromised?.(new Error("Unable to update lock within the stale threshold")); - return realLock(file, options); - }); +describe("AuthStorage file backend regressions", () => { + let tempDir: string; + let authPath: string; - const firstTry = await authStorage.getApiKey(providerId); - expect(firstTry).toBeUndefined(); + beforeEach(() => { + tempDir = join(tmpdir(), `atomic-auth-backend-${Date.now()}-${Math.random().toString(36).slice(2)}`); + mkdirSync(tempDir, { recursive: true }); + authPath = join(tempDir, "auth.json"); + writeFileSync(authPath, JSON.stringify({ anthropic: { type: "api_key", key: "anthropic-key" } })); + }); - lockSpy.mockRestore(); + afterEach(() => { + if (existsSync(tempDir)) rmSync(tempDir, { recursive: true }); + }); - const secondTry = await authStorage.getApiKey(providerId); - expect(secondTry).toBe("Bearer refreshed-access-token"); - }); - test("reload reads credentials without acquiring the write lock", () => { - // A configured provider must remain readable even while another - // process/stage holds the exclusive auth.json lock. Pure reads are - // lock-free, so a fresh AuthStorage never reports ELOCKED here. - writeAuthJson({ anthropic: { type: "api_key", key: "anthropic-key" } }); + test("fresh storage reads credentials while another process holds the write lock", async () => { + const release = lockfile.lockSync(authPath, { realpath: false }); + try { + const storage = AuthStorage.create(authPath); + await expect(storage.read("anthropic")).resolves.toEqual({ type: "api_key", key: "anthropic-key" }); + } finally { + release(); + } + }); - const release = lockfile.lockSync(authJsonPath, { realpath: false }); - try { - const storage = AuthStorage.create(authJsonPath); - expect(storage.getLoadError()).toBeNull(); - expect(storage.hasAuth("anthropic")).toBe(true); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "anthropic-key" }); - } finally { - release(); + test("many fresh storage reads succeed while the write lock is held", async () => { + const release = lockfile.lockSync(authPath, { realpath: false }); + try { + for (let index = 0; index < 12; index++) { + const storage = AuthStorage.create(authPath); + await expect(storage.read("anthropic")).resolves.toEqual({ type: "api_key", key: "anthropic-key" }); } - }); - test("many fresh AuthStorage reads succeed while the lock is held", () => { - // Mirrors the fallback-auth-stress harness: many parallel stages each - // build their own AuthStorage. None should be starved by the held lock. - writeAuthJson({ anthropic: { type: "api_key", key: "anthropic-key" } }); + } finally { + release(); + } + }); - const release = lockfile.lockSync(authJsonPath, { realpath: false }); - try { - for (let i = 0; i < 12; i++) { - const storage = AuthStorage.create(authJsonPath); - expect(storage.getLoadError()).toBeNull(); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "anthropic-key" }); - } - } finally { - release(); - } + test("reads resolve while an async writer holds the lock across an await", async () => { + const backend = new FileAuthStorageBackend(authPath, [authPath]); + let releaseHeld!: () => void; + const held = new Promise((resolve) => { + releaseHeld = resolve; }); - test("writes are atomic and leave no temp files behind", () => { - writeAuthJson({ anthropic: { type: "api_key", key: "a" } }); - authStorage = AuthStorage.create(authJsonPath); - - authStorage.set("openai", { type: "api_key", key: "b" }); - - const leftovers = readdirSync(tempDir).filter((entry) => entry.endsWith(".tmp")); - expect(leftovers).toEqual([]); - - const parsed = JSON.parse(readFileSync(authJsonPath, "utf-8")) as Record; - expect(parsed.anthropic).toEqual({ type: "api_key", key: "a" }); - expect(parsed.openai).toEqual({ type: "api_key", key: "b" }); + let markEntered!: () => void; + const entered = new Promise((resolve) => { + markEntered = resolve; }); - test("atomic write preserves 0600 permissions", () => { - if (process.platform === "win32") return; - writeAuthJson({ anthropic: { type: "api_key", key: "a" } }); - authStorage = AuthStorage.create(authJsonPath); - - authStorage.set("openai", { type: "api_key", key: "b" }); - - expect(statSync(authJsonPath).mode & 0o777).toBe(0o600); + const writer = backend.withLockAsync(async () => { + markEntered(); + await held; + return { result: undefined }; }); - test("reads resolve while an async writer holds the lock across an await", async () => { - // Reproduces the fallback-auth-stress interaction: an in-flight OAuth - // refresh holds the exclusive lock across a network await while sibling - // stages create fresh AuthStorages. With a blocking, lock-taking read - // path these reads would busy-wait and fail ELOCKED; lock-free reads must - // resolve immediately and keep the event loop free so the writer can - // finish (issue #1431). - writeAuthJson({ anthropic: { type: "api_key", key: "anthropic-key" } }); - const backend = new FileAuthStorageBackend(authJsonPath, [authJsonPath]); - - let releaseHeld!: () => void; - const held = new Promise((resolve) => { - releaseHeld = resolve; - }); - const writer = backend.withLockAsync(async () => { - await held; - return { result: undefined }; - }); + await entered; - try { - for (let i = 0; i < 10; i++) { - const storage = AuthStorage.create(authJsonPath); - expect(storage.getLoadError()).toBeNull(); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "anthropic-key" }); - } - } finally { - releaseHeld(); - await writer; + try { + for (let index = 0; index < 10; index++) { + const storage = AuthStorage.create(authPath); + await expect(storage.read("anthropic")).resolves.toEqual({ type: "api_key", key: "anthropic-key" }); } - }); - test("set preserves unrelated external edits", () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "old-anthropic" }, - openai: { type: "api_key", key: "openai-key" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - - // Simulate external edit while process is running - writeAuthJson({ - anthropic: { type: "api_key", key: "old-anthropic" }, - openai: { type: "api_key", key: "openai-key" }, - google: { type: "api_key", key: "google-key" }, - }); - - authStorage.set("anthropic", { type: "api_key", key: "new-anthropic" }); - - const updated = JSON.parse(readFileSync(authJsonPath, "utf-8")) as Record; - expect(updated.anthropic.key).toBe("new-anthropic"); - expect(updated.openai.key).toBe("openai-key"); - expect(updated.google.key).toBe("google-key"); - }); - test("remove preserves unrelated external edits", () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "anthropic-key" }, - openai: { type: "api_key", key: "openai-key" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - - // Simulate external edit while process is running - writeAuthJson({ - anthropic: { type: "api_key", key: "anthropic-key" }, - openai: { type: "api_key", key: "openai-key" }, - google: { type: "api_key", key: "google-key" }, - }); + } finally { + releaseHeld(); + await writer; + } + }); - authStorage.remove("anthropic"); + test("atomic writes leave no sibling temporary files", async () => { + const storage = AuthStorage.create(authPath); + await storage.modify("openai", async () => ({ type: "api_key", key: "openai-key" })); - const updated = JSON.parse(readFileSync(authJsonPath, "utf-8")) as Record; - expect(updated.anthropic).toBeUndefined(); - expect(updated.openai.key).toBe("openai-key"); - expect(updated.google.key).toBe("google-key"); + expect(readdirSync(tempDir).filter((entry) => entry.endsWith(".tmp"))).toEqual([]); + expect(JSON.parse(readFileSync(authPath, "utf8"))).toEqual({ + anthropic: { type: "api_key", key: "anthropic-key" }, + openai: { type: "api_key", key: "openai-key" }, }); - test("surfaces a malformed auth file without overwriting it", () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "anthropic-key" }, - }); - - authStorage = AuthStorage.create(authJsonPath); - writeFileSync(authJsonPath, "{invalid-json", "utf-8"); - - authStorage.reload(); - expect(() => authStorage.set("openai", { type: "api_key", key: "openai-key" })).toThrow(); + }); - const raw = readFileSync(authJsonPath, "utf-8"); - expect(raw).toBe("{invalid-json"); - expect(authStorage.get("openai")).toBeUndefined(); - }); -}); + const posixTest = process.platform === "win32" ? test.skip : test; + posixTest("atomic writes preserve 0600 file permissions", async () => { + const storage = AuthStorage.create(authPath); + await storage.modify("openai", async () => ({ type: "api_key", key: "openai-key" })); + expect(statSync(authPath).mode & 0o777).toBe(0o600); + }); }); diff --git a/packages/coding-agent/test/auth-storage-02.suite.ts b/packages/coding-agent/test/auth-storage-02.suite.ts index 593cc2f42..28c49ba94 100644 --- a/packages/coding-agent/test/auth-storage-02.suite.ts +++ b/packages/coding-agent/test/auth-storage-02.suite.ts @@ -1,33 +1,14 @@ -import { existsSync, mkdirSync, readdirSync, readFileSync, rmSync, statSync, writeFileSync } from "node:fs"; +import { existsSync, mkdirSync, readFileSync, rmSync, writeFileSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; +import { createModels, type Provider } from "@earendil-works/pi-ai"; import lockfile from "proper-lockfile"; import { afterEach, beforeEach, describe, expect, test, vi } from "vitest"; -import { AuthStorage, type AuthStorageBackend, FileAuthStorageBackend } from "../src/core/auth-storage.ts"; -import { clearConfigValueCache } from "../src/core/resolve-config-value.ts"; - -/** - * Backend whose every access throws — used to simulate a credential-store load - * failure (e.g. ELOCKED under concurrent auth.json access) so a *fresh* - * AuthStorage ends up with an empty in-memory store and a recorded loadError - * (issue #1431). - */ -class ThrowingAuthStorageBackend implements AuthStorageBackend { - constructor(private readonly error: Error) {} - read(): string | undefined { - throw this.error; - } - withLock(): T { - throw this.error; - } - async withLockAsync(): Promise { - throw this.error; - } -} +import { AuthStorage } from "../src/core/auth-storage.ts"; + describe("AuthStorage", () => { let tempDir: string; let authJsonPath: string; - let authStorage: AuthStorage; beforeEach(() => { tempDir = join(tmpdir(), `pi-test-auth-storage-${Date.now()}-${Math.random().toString(36).slice(2)}`); @@ -36,121 +17,201 @@ describe("AuthStorage", () => { }); afterEach(() => { - if (tempDir && existsSync(tempDir)) { - rmSync(tempDir, { recursive: true }); - } - clearConfigValueCache(); + if (existsSync(tempDir)) rmSync(tempDir, { recursive: true }); vi.restoreAllMocks(); }); - function writeAuthJson(data: Record) { + function writeAuthJson(data: Record): void { writeFileSync(authJsonPath, JSON.stringify(data)); } - function toShPath(value: string): string { - let escaped = ""; - for (const char of value.replace(/\\/g, "/")) { - if (char === '"' || char === "\\" || char === "$" || char === "`") { - escaped += `\\${char}`; - } else { - escaped += char; - } + test("reads and resolves stored API-key credentials", async () => { + const original = process.env.TEST_AUTH_STORAGE_KEY; + process.env.TEST_AUTH_STORAGE_KEY = "environment-key"; + try { + writeAuthJson({ anthropic: { type: "api_key", key: "$TEST_AUTH_STORAGE_KEY" } }); + const storage = AuthStorage.create(authJsonPath); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "environment-key" }); + } finally { + if (original === undefined) delete process.env.TEST_AUTH_STORAGE_KEY; + else process.env.TEST_AUTH_STORAGE_KEY = original; } - return escaped; - } + }); - describe("API key resolution", () => { - test("reload records parse errors and drainErrors clears buffer", () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "anthropic-key" }, - }); + test("resolves command-backed API-key credentials", async () => { + writeAuthJson({ anthropic: { type: "api_key", key: "!printf 'command-key'" } }); + const storage = AuthStorage.create(authJsonPath); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "command-key" }); + }); - authStorage = AuthStorage.create(authJsonPath); - writeFileSync(authJsonPath, "{invalid-json", "utf-8"); + test("returns OAuth credentials unchanged", async () => { + const credential = { + type: "oauth" as const, + access: "access-token", + refresh: "refresh-token", + expires: Date.now() + 60_000, + }; + const storage = AuthStorage.inMemory({ anthropic: credential }); + expect(await storage.read("anthropic")).toEqual(credential); + }); - authStorage.reload(); + test("credential-scoped env takes precedence and remains inspectable", async () => { + writeAuthJson({ + anthropic: { + type: "api_key", + key: "$SCOPED_KEY", + env: { SCOPED_KEY: "scoped-value", REGION: "test-region" }, + }, + }); + const storage = AuthStorage.create(authJsonPath); + expect(await storage.read("anthropic")).toMatchObject({ + key: "scoped-value", + env: { SCOPED_KEY: "scoped-value", REGION: "test-region" }, + }); + }); - // Keeps previous in-memory data on reload failure - expect(authStorage.get("anthropic")).toEqual({ type: "api_key", key: "anthropic-key" }); + test("modify persists a credential while preserving unrelated external edits", async () => { + writeAuthJson({ anthropic: { type: "api_key", key: "old" } }); + const storage = AuthStorage.create(authJsonPath); + writeAuthJson({ + anthropic: { type: "api_key", key: "old" }, + openai: { type: "api_key", key: "external" }, + }); - const firstDrain = authStorage.drainErrors(); - expect(firstDrain.length).toBeGreaterThan(0); - expect(firstDrain[0]).toBeInstanceOf(Error); + await storage.modify("anthropic", async () => ({ type: "api_key", key: "new" })); - const secondDrain = authStorage.drainErrors(); - expect(secondDrain).toHaveLength(0); + expect(JSON.parse(readFileSync(authJsonPath, "utf8"))).toEqual({ + anthropic: { type: "api_key", key: "new" }, + openai: { type: "api_key", key: "external" }, }); - test("getLoadError is null after a successful load", () => { - writeAuthJson({ anthropic: { type: "api_key", key: "anthropic-key" } }); - authStorage = AuthStorage.create(authJsonPath); - expect(authStorage.getLoadError()).toBeNull(); + }); + + test("modify with undefined leaves the current credential unchanged", async () => { + writeAuthJson({ anthropic: { type: "api_key", key: "stored" } }); + const storage = AuthStorage.create(authJsonPath); + expect(await storage.modify("anthropic", async () => undefined)).toEqual({ type: "api_key", key: "stored" }); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "stored" }); + }); + + test("serializes concurrent modifications", async () => { + writeAuthJson({}); + const first = AuthStorage.create(authJsonPath); + const second = AuthStorage.create(authJsonPath); + await Promise.all([ + first.modify("anthropic", async () => ({ type: "api_key", key: "anthropic-key" })), + second.modify("openai", async () => ({ type: "api_key", key: "openai-key" })), + ]); + expect(JSON.parse(readFileSync(authJsonPath, "utf8"))).toEqual({ + anthropic: { type: "api_key", key: "anthropic-key" }, + openai: { type: "api_key", key: "openai-key" }, }); - test("a fresh-load failure is surfaced via getLoadError and leaves credentials empty", () => { - const loadError = Object.assign(new Error("Lock file is already being held"), { code: "ELOCKED" }); - const storage = AuthStorage.fromStorage(new ThrowingAuthStorageBackend(loadError)); - - // The failure is preserved, not swallowed: an empty store is NOT - // authoritative — the provider is not "absent", the store could not be read. - expect(storage.getLoadError()).toBe(loadError); - expect(storage.hasAuth("some-locked-provider")).toBe(false); - expect(storage.list()).toEqual([]); + }); + + test("delete removes one credential while preserving others", async () => { + writeAuthJson({ + anthropic: { type: "api_key", key: "anthropic-key" }, + openai: { type: "api_key", key: "openai-key" }, }); - test("getLoadError clears after a subsequent successful reload", () => { - writeAuthJson({ anthropic: { type: "api_key", key: "anthropic-key" } }); - authStorage = AuthStorage.create(authJsonPath); - expect(authStorage.getLoadError()).toBeNull(); - - // Corrupt the file and reload -> load error recorded. - writeFileSync(authJsonPath, "{invalid-json", "utf-8"); - authStorage.reload(); - expect(authStorage.getLoadError()).toBeInstanceOf(Error); - - // Repair and reload -> error cleared. - writeAuthJson({ anthropic: { type: "api_key", key: "anthropic-key" } }); - authStorage.reload(); - expect(authStorage.getLoadError()).toBeNull(); + const storage = AuthStorage.create(authJsonPath); + writeAuthJson({ + anthropic: { type: "api_key", key: "anthropic-key" }, + openai: { type: "api_key", key: "openai-key" }, + google: { type: "api_key", key: "external-key" }, }); - test("does not expose stored API keys or OAuth tokens", () => { - authStorage = AuthStorage.inMemory({ - anthropic: { type: "api_key", key: "secret-api-key" }, - openai: { - type: "oauth", - access: "secret-access-token", - refresh: "secret-refresh-token", - expires: Date.now() + 1000, - }, - }); + await storage.delete("anthropic"); + await expect(storage.list()).resolves.toEqual([ + { providerId: "openai", type: "api_key" }, + { providerId: "google", type: "api_key" }, + ]); + expect(await storage.read("anthropic")).toBeUndefined(); + expect(await storage.read("openai")).toEqual({ type: "api_key", key: "openai-key" }); + expect(await storage.read("google")).toEqual({ type: "api_key", key: "external-key" }); + }); - expect(authStorage.getAuthStatus("anthropic")).toEqual({ configured: true, source: "stored" }); - expect(authStorage.getAuthStatus("openai")).toEqual({ configured: true, source: "stored" }); - expect(JSON.stringify(authStorage.getAuthStatus("anthropic"))).not.toContain("secret-api-key"); - expect(JSON.stringify(authStorage.getAuthStatus("openai"))).not.toContain("secret-access-token"); - expect(JSON.stringify(authStorage.getAuthStatus("openai"))).not.toContain("secret-refresh-token"); - }); - test("runtime override takes priority over auth.json", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!echo stored-key" }, - }); + test("in-memory storage implements the same credential-store behavior", async () => { + const storage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "initial" } }); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "initial" }); + await storage.modify("anthropic", async () => ({ type: "api_key", key: "updated" })); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "updated" }); + await storage.delete("anthropic"); + await expect(storage.list()).resolves.toEqual([]); + }); - authStorage = AuthStorage.create(authJsonPath); - authStorage.setRuntimeApiKey("anthropic", "runtime-key"); + test("does not write after lock acquisition failure and recovers on retry", async () => { + writeAuthJson({ anthropic: { type: "api_key", key: "stored" } }); + const storage = AuthStorage.create(authJsonPath); + const lockSpy = vi.spyOn(lockfile, "lock").mockRejectedValueOnce(new Error("lock unavailable")); - const apiKey = await authStorage.getApiKey("anthropic"); + await expect(storage.modify("openai", async () => ({ type: "api_key", key: "new" }))).rejects.toThrow( + "lock unavailable", + ); + expect(JSON.parse(readFileSync(authJsonPath, "utf8"))).toEqual({ + anthropic: { type: "api_key", key: "stored" }, + }); - expect(apiKey).toBe("runtime-key"); + lockSpy.mockRestore(); + await storage.modify("openai", async () => ({ type: "api_key", key: "new" })); + expect(JSON.parse(readFileSync(authJsonPath, "utf8"))).toEqual({ + anthropic: { type: "api_key", key: "stored" }, + openai: { type: "api_key", key: "new" }, }); - test("removing runtime override falls back to auth.json", async () => { - writeAuthJson({ - anthropic: { type: "api_key", key: "!echo stored-key" }, - }); + }); - authStorage = AuthStorage.create(authJsonPath); - authStorage.setRuntimeApiKey("anthropic", "runtime-key"); - authStorage.removeRuntimeApiKey("anthropic"); + test("surfaces a compromised OAuth refresh lock and allows a later retry", async () => { + const providerId = "oauth-provider"; + writeAuthJson({ + [providerId]: { + type: "oauth", + access: "expired-access", + refresh: "refresh-token", + expires: 0, + }, + }); + const storage = AuthStorage.create(authJsonPath); + const provider: Provider = { + id: providerId, + name: "OAuth Provider", + auth: { + oauth: { + name: "OAuth", + login: async () => { + throw new Error("not used"); + }, + refresh: async (credential) => ({ + ...credential, + access: "refreshed-access", + expires: Date.now() + 60_000, + }), + toAuth: async (credential) => ({ apiKey: credential.access }), + }, + }, + getModels: () => [], + stream: () => { + throw new Error("not used"); + }, + streamSimple: () => { + throw new Error("not used"); + }, + }; + const models = createModels({ credentials: storage }); + models.setProvider(provider); + + const realLock = lockfile.lock.bind(lockfile); + const lockSpy = vi.spyOn(lockfile, "lock").mockImplementationOnce(async (file, options) => { + options?.onCompromised?.(new Error("lock compromised")); + return realLock(file, options); + }); + await expect(models.getAuth(providerId)).rejects.toMatchObject({ code: "auth" }); - const apiKey = await authStorage.getApiKey("anthropic"); + lockSpy.mockRestore(); + await expect(models.getAuth(providerId)).resolves.toMatchObject({ auth: { apiKey: "refreshed-access" } }); + }); - expect(apiKey).toBe("stored-key"); - }); -}); + test("does not overwrite malformed auth files", async () => { + writeAuthJson({ anthropic: { type: "api_key", key: "stored" } }); + const storage = AuthStorage.create(authJsonPath); + writeFileSync(authJsonPath, "{invalid-json", "utf8"); + await expect(storage.modify("openai", async () => ({ type: "api_key", key: "new" }))).rejects.toThrow(); + expect(readFileSync(authJsonPath, "utf8")).toBe("{invalid-json"); + }); }); diff --git a/packages/coding-agent/test/auth-storage-persistence.test.ts b/packages/coding-agent/test/auth-storage-persistence.test.ts index e91779236..e8b2a469d 100644 --- a/packages/coding-agent/test/auth-storage-persistence.test.ts +++ b/packages/coding-agent/test/auth-storage-persistence.test.ts @@ -1,6 +1,5 @@ -import { afterEach, describe, expect, test } from "vitest"; +import { describe, expect, test } from "vitest"; import { AuthStorage, type AuthStorageBackend } from "../src/core/auth-storage.ts"; -import { registerLegacyOAuthProvider, resetLegacyOAuthProviders } from "../src/core/oauth-provider-bridge.ts"; class ControllableAuthBackend implements AuthStorageBackend { value: string | undefined; @@ -33,10 +32,9 @@ class ControllableAuthBackend implements AuthStorageBackend { } } -afterEach(() => resetLegacyOAuthProviders()); - describe("AuthStorage persistence failures", () => { - test("surfaces malformed storage and preserves in-memory credentials", () => { + + test("surfaces malformed storage without overwriting the last valid snapshot", async () => { const backend = new ControllableAuthBackend( JSON.stringify({ anthropic: { type: "api_key", key: "existing" } }), ); @@ -44,203 +42,64 @@ describe("AuthStorage persistence failures", () => { backend.value = "{invalid-json"; storage.reload(); - expect(() => storage.set("openai", { type: "api_key", key: "new" })).toThrow(); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); - expect(storage.get("openai")).toBeUndefined(); + await expect(storage.modify("openai", async () => ({ type: "api_key", key: "new" }))).rejects.toThrow(); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "existing" }); + expect(await storage.read("openai")).toBeUndefined(); expect(backend.value).toBe("{invalid-json"); }); - test("surfaces write failures without mutating in-memory credentials", () => { + test("failed modify remains transactional in storage and memory", async () => { const backend = new ControllableAuthBackend( JSON.stringify({ anthropic: { type: "api_key", key: "existing" } }), ); const storage = AuthStorage.fromStorage(backend); backend.writeError = new Error("disk full"); - expect(() => storage.set("anthropic", { type: "api_key", key: "replacement" })).toThrow("disk full"); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); - expect(() => storage.remove("anthropic")).toThrow("disk full"); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); + await expect( + storage.modify("anthropic", async () => ({ type: "api_key", key: "replacement" })), + ).rejects.toThrow("disk full"); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "existing" }); expect(JSON.parse(backend.value ?? "{}")).toEqual({ anthropic: { type: "api_key", key: "existing" } }); - expect(() => storage.logout("anthropic")).toThrow("disk full"); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); }); - test("credential adapter keeps failed writes transactional", async () => { + test("failed delete remains transactional in storage and memory", async () => { const backend = new ControllableAuthBackend( JSON.stringify({ anthropic: { type: "api_key", key: "existing" } }), ); const storage = AuthStorage.fromStorage(backend); backend.writeError = new Error("disk full"); - await expect( - storage.asCredentialStore().modify("anthropic", async () => ({ type: "api_key", key: "replacement" })), - ).rejects.toThrow("disk full"); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); + await expect(storage.delete("anthropic")).rejects.toThrow("disk full"); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "existing" }); + expect(JSON.parse(backend.value ?? "{}")).toEqual({ anthropic: { type: "api_key", key: "existing" } }); }); - test("OAuth login persistence failure preserves the previously committed credential", async () => { - const backend = new ControllableAuthBackend(JSON.stringify({ - openrouter: { type: "api_key", key: "previous-key" }, - })); + test("a repaired credential snapshot accepts the next modification", async () => { + const backend = new ControllableAuthBackend("{invalid-json"); const storage = AuthStorage.fromStorage(backend); - registerLegacyOAuthProvider("openrouter", { - name: "OpenRouter test", - login: async () => ({ refresh: "", access: "minted-key", expires: Number.MAX_SAFE_INTEGER }), - refreshToken: async (credentials) => credentials, - getApiKey: (credentials) => credentials.access, - }); - backend.writeError = new Error("auth.json is read-only"); - - await expect(storage.login("openrouter", { - onAuth: () => {}, - onDeviceCode: () => {}, - onPrompt: async () => "", - onSelect: async () => undefined, - })).rejects.toThrow("auth.json is read-only"); - expect(storage.get("openrouter")).toEqual({ type: "api_key", key: "previous-key" }); - }); - - test("cancelled OAuth login preserves the previously committed credential", async () => { - const storage = AuthStorage.inMemory({ - "kimi-coding": { type: "api_key", key: "previous-key" }, - }); - registerLegacyOAuthProvider("kimi-coding", { - name: "Kimi Code test", - login: async () => { throw new Error("Login cancelled"); }, - refreshToken: async (credentials) => credentials, - getApiKey: (credentials) => credentials.access, - }); - - await expect(storage.login("kimi-coding", { - onAuth: () => {}, - onDeviceCode: () => {}, - onPrompt: async () => "", - onSelect: async () => undefined, - })).rejects.toThrow("Login cancelled"); - expect(storage.get("kimi-coding")).toEqual({ type: "api_key", key: "previous-key" }); - }); + backend.value = JSON.stringify({ anthropic: { type: "api_key", key: "existing" } }); - test("failed async logout keeps the persisted credential in memory", async () => { - const backend = new ControllableAuthBackend( - JSON.stringify({ anthropic: { type: "api_key", key: "existing" } }), - ); - const storage = AuthStorage.fromStorage(backend); - backend.writeError = new Error("disk full"); + await storage.modify("openai", async () => ({ type: "api_key", key: "new" })); - await expect(storage.logoutAsync("anthropic")).rejects.toThrow("disk full"); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "existing" }); + expect(await storage.read("openai")).toEqual({ type: "api_key", key: "new" }); }); - test("credential adapter serializes delete behind an in-flight modify", async () => { + test("delete is serialized behind an in-flight modification", async () => { const storage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "existing" } }); - const credentials = storage.asCredentialStore(); let release!: () => void; - const gate = new Promise((resolve) => { release = resolve; }); - const modify = credentials.modify("anthropic", async () => { + const gate = new Promise((resolve) => { + release = resolve; + }); + const modification = storage.modify("anthropic", async () => { await gate; return { type: "api_key", key: "replacement" }; }); - const deletion = credentials.delete("anthropic"); - - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); - release(); - await Promise.all([modify, deletion]); - expect(storage.get("anthropic")).toBeUndefined(); - }); - - test("serialized logout wins over an in-flight legacy OAuth refresh", async () => { - let release!: () => void; - const gate = new Promise((resolve) => { release = resolve; }); - registerLegacyOAuthProvider("legacy", { - name: "Legacy", - login: async () => ({ refresh: "r", access: "a", expires: 1 }), - refreshToken: async () => { - await gate; - return { refresh: "r2", access: "a2", expires: Date.now() + 60_000 }; - }, - getApiKey: (credentials) => credentials.access, - }); - const storage = AuthStorage.inMemory({ - legacy: { type: "oauth", refresh: "r", access: "a", expires: 0 }, - }); - const refresh = storage.getModelAuth("legacy"); - const logout = storage.logoutAsync("legacy"); - - release(); - await Promise.all([refresh, logout]); - - expect(storage.get("legacy")).toBeUndefined(); - }); - - test("login is serialized after an in-flight legacy OAuth refresh", async () => { - let release!: () => void; - const gate = new Promise((resolve) => { release = resolve; }); - registerLegacyOAuthProvider("login-race", { - name: "Login Race", - login: async () => ({ refresh: "login-r", access: "login-a", expires: Date.now() + 60_000 }), - refreshToken: async () => { - await gate; - return { refresh: "refresh-r", access: "refresh-a", expires: Date.now() + 60_000 }; - }, - getApiKey: (credentials) => credentials.access, - }); - const storage = AuthStorage.inMemory({ - "login-race": { type: "oauth", refresh: "old-r", access: "old-a", expires: 0 }, - }); - const refresh = storage.getModelAuth("login-race"); - const login = storage.login("login-race", { - onAuth: () => {}, - onDeviceCode: () => {}, - onPrompt: async () => "", - onSelect: async () => undefined, - }); + const deletion = storage.delete("anthropic"); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "existing" }); release(); - await Promise.all([refresh, login]); - - expect(storage.get("login-race")).toMatchObject({ access: "login-a", refresh: "login-r" }); - }); - - test("legacy synchronous logout is serialized behind an in-flight OAuth refresh", async () => { - let release!: () => void; - const gate = new Promise((resolve) => { release = resolve; }); - let markEntered!: () => void; - const entered = new Promise((resolve) => { markEntered = resolve; }); - registerLegacyOAuthProvider("legacy-sync", { - name: "Legacy Sync", - login: async () => ({ refresh: "r", access: "a", expires: 1 }), - refreshToken: async () => { - markEntered(); - await gate; - return { refresh: "r2", access: "a2", expires: Date.now() + 60_000 }; - }, - getApiKey: (credentials) => credentials.access, - }); - const storage = AuthStorage.inMemory({ - "legacy-sync": { type: "oauth", refresh: "r", access: "a", expires: 0 }, - }); - const refresh = storage.getModelAuth("legacy-sync"); - await entered; - - storage.logout("legacy-sync"); - expect(storage.get("legacy-sync")).toBeUndefined(); - release(); - await refresh; - - expect(storage.get("legacy-sync")).toBeUndefined(); - }); - - test("recovers after the credential snapshot is repaired", () => { - const backend = new ControllableAuthBackend("{invalid-json"); - const storage = AuthStorage.fromStorage(backend); - expect(storage.getLoadError()).toBeInstanceOf(Error); - - backend.value = JSON.stringify({ anthropic: { type: "api_key", key: "existing" } }); - storage.set("openai", { type: "api_key", key: "new" }); - - expect(storage.getLoadError()).toBeNull(); - expect(storage.get("anthropic")).toEqual({ type: "api_key", key: "existing" }); - expect(storage.get("openai")).toEqual({ type: "api_key", key: "new" }); + await Promise.all([modification, deletion]); + expect(await storage.read("anthropic")).toBeUndefined(); }); }); diff --git a/packages/coding-agent/test/compaction-extensions.test.ts b/packages/coding-agent/test/compaction-extensions.test.ts index 2b54f27e6..b60dfc86a 100644 --- a/packages/coding-agent/test/compaction-extensions.test.ts +++ b/packages/coding-agent/test/compaction-extensions.test.ts @@ -5,10 +5,10 @@ import { afterEach, beforeEach, describe, expect, it } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { createExtensionRuntime, type Extension, type SessionBeforeCompactEvent, type SessionBeforeCompactResult, type SessionCompactEvent, type SessionEvent } from "../src/core/extensions/index.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { createSyntheticSourceInfo } from "../src/core/source-info.ts"; +import { createInMemoryModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; import { createTestResourceLoader } from "./utilities.ts"; import { createFauxStreamFn } from "./test-harness.ts"; @@ -40,13 +40,14 @@ describe("verbatim compaction extension hooks", () => { return { path: "test", resolvedPath: "/test.ts", sourceInfo: createSyntheticSourceInfo("", { source: "test" }), handlers, tools: new Map(), messageRenderers: new Map(), commands: new Map(), flags: new Map(), shortcuts: new Map() }; } - function create(ext: Extension, streamFn?: StreamFn): void { + async function create(ext: Extension, streamFn?: StreamFn): Promise { const model = getModel("anthropic", "claude-sonnet-4-5")!; const manager = SessionManager.inMemory(); const agent = new Agent({ getApiKey: () => undefined, initialState: { model, systemPrompt: "test", tools: [] }, streamFn }); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("anthropic", "test-key"); - session = new AgentSession({ agent, sessionManager: manager, settingsManager: SettingsManager.inMemory(), cwd: process.cwd(), modelRegistry: ModelRegistry.create(authStorage), resourceLoader: { ...createTestResourceLoader(), getExtensions: () => ({ extensions: [ext], errors: [], runtime: createExtensionRuntime() }) } }); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); + const modelRegistry = await createInMemoryModelRegistry(authStorage); + session = new AgentSession({ agent, sessionManager: manager, settingsManager: SettingsManager.inMemory(), cwd: process.cwd(), modelRuntime: getModelRuntime(modelRegistry), resourceLoader: { ...createTestResourceLoader(), getExtensions: () => ({ extensions: [ext], errors: [], runtime: createExtensionRuntime() }) } }); const now = Date.now(); for (let turn = 0; turn < 5; turn++) { manager.appendMessage({ role: "user", content: `task ${turn}\nline a\nline b`, timestamp: now + turn * 2 }); @@ -56,7 +57,7 @@ describe("verbatim compaction extension hooks", () => { } it("accepts a non-empty offline compactedText override and emits observe-only result", async () => { - create(extension((event) => { + await create(extension((event) => { const headers = event.preparation.region.headerLineNumbers as Set; const markers = event.preparation.region.priorMarkerNs as Map; expect(() => headers.add(999)).toThrow("Cannot mutate frozen compaction preparation"); @@ -75,7 +76,7 @@ describe("verbatim compaction extension hooks", () => { it("persists zero retention and includes that durable summary on repeated compaction", async () => { const compactedText = "[User]: retained exactly\n(filtered 30 lines)"; - create(extension(() => ({ compactedText }))); + await create(extension(() => ({ compactedText }))); const first = await session.compact({ preserve_recent: 0 }); expect(first.firstKeptEntryId).toBeNull(); @@ -95,7 +96,7 @@ describe("verbatim compaction extension hooks", () => { it("sends the prior durable summary through the planner on repeated compaction", async () => { const faux = createFauxStreamFn(["2,2\n", "2,2\n"]); - create(extension(() => undefined), faux.streamFn); + await create(extension(() => undefined), faux.streamFn); await session.compact({ preserve_recent: 0 }); const long = Array.from({ length: 20 }, (_, index) => `planner line ${index}`).join("\n"); @@ -113,13 +114,13 @@ describe("verbatim compaction extension hooks", () => { }); it("cancels without persistence", async () => { - create(extension(() => ({ cancel: true }))); + await create(extension(() => ({ cancel: true }))); await expect(session.compact()).rejects.toThrow("Compaction cancelled"); expect(session.sessionManager.getEntries().some((entry) => entry.type === "compaction")).toBe(false); }); it("rejects whitespace extension text before persistence", async () => { - create(extension(() => ({ compactedText: " \n" }))); + await create(extension(() => ({ compactedText: " \n" }))); await expect(session.compact()).rejects.toThrow("No compacted text provided by extension"); expect(session.sessionManager.getEntries().some((entry) => entry.type === "compaction")).toBe(false); }); @@ -129,7 +130,7 @@ describe("verbatim compaction extension hooks", () => { ["empty", ""], ])("does not persist a compaction entry after one %s planner response", async (_label, response) => { const faux = createFauxStreamFn([response]); - create(extension(() => undefined), faux.streamFn); + await create(extension(() => undefined), faux.streamFn); await expect(session.compact()).rejects.toThrow(/Compaction range planning/); expect(faux.state.callCount).toBe(1); expect(session.sessionManager.getEntries().some((entry) => entry.type === "compaction")).toBe(false); @@ -141,13 +142,13 @@ describe("verbatim compaction extension hooks", () => { calls++; throw new Error("provider unavailable"); }; - create(extension(() => undefined), failingStream); + await create(extension(() => undefined), failingStream); await expect(session.compact()).rejects.toThrow("provider unavailable"); expect(calls).toBe(1); expect(session.sessionManager.getEntries().some((entry) => entry.type === "compaction")).toBe(false); }); it("isolates errors from the post-commit observer", async () => { - create(extension(() => ({ compactedText: "[User]: retained" }), () => { throw new Error("observer failed"); })); + await create(extension(() => ({ compactedText: "[User]: retained" }), () => { throw new Error("observer failed"); })); await expect(session.compact()).resolves.toMatchObject({ rung: "extension" }); expect(session.sessionManager.getEntries().some((entry) => entry.type === "compaction")).toBe(true); }); diff --git a/packages/coding-agent/test/config-value-migration.test.ts b/packages/coding-agent/test/config-value-migration.test.ts index 95f22d2fc..0e4790b5a 100644 --- a/packages/coding-agent/test/config-value-migration.test.ts +++ b/packages/coding-agent/test/config-value-migration.test.ts @@ -4,9 +4,10 @@ import * as path from "node:path"; import { afterEach, describe, expect, it, vi } from "vitest"; import { ENV_AGENT_DIR } from "../src/config.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { runMigrations } from "../src/migrations.ts"; +import { createModelRegistry } from "./model-runtime-test-utils.ts"; + describe("config value env var syntax migration", () => { const tempDirs: string[] = []; @@ -18,7 +19,7 @@ describe("config value env var syntax migration", () => { }); function createAgentDir(): string { - const agentDir = fs.mkdtempSync(path.join(os.tmpdir(), "atomic-config-value-migration-test-")); + const agentDir = fs.mkdtempSync(path.join(os.tmpdir(), "pi-config-value-migration-test-")); tempDirs.push(agentDir); return agentDir; } @@ -37,64 +38,19 @@ describe("config value env var syntax migration", () => { } } - it("rewrites legacy uppercase auth.json API key values in legacy .pi agent config when the env var exists", () => { - const homeDir = fs.mkdtempSync(path.join(os.tmpdir(), "atomic-legacy-pi-config-value-migration-test-")); - tempDirs.push(homeDir); - const previousHome = process.env.HOME; - const previousUserProfile = process.env.USERPROFILE; - const previousAgentDir = process.env[ENV_AGENT_DIR]; - const previousAnthropicKey = process.env.ANTHROPIC_API_KEY; - delete process.env[ENV_AGENT_DIR]; - process.env.HOME = homeDir; - process.env.USERPROFILE = homeDir; - process.env.ANTHROPIC_API_KEY = "secret"; + function withEnv(name: string, value: string | undefined, fn: () => void): void { + const previous = process.env[name]; + if (value === undefined) delete process.env[name]; + else process.env[name] = value; try { - const legacyAgentDir = path.join(homeDir, ".pi", "agent"); - fs.mkdirSync(legacyAgentDir, { recursive: true }); - fs.writeFileSync( - path.join(legacyAgentDir, "auth.json"), - `${JSON.stringify({ anthropic: { type: "api_key", key: "ANTHROPIC_API_KEY" } }, null, 2)}\n`, - "utf-8", - ); - const logSpy = vi.spyOn(console, "log").mockImplementation(() => {}); - - runMigrations(homeDir); - - const migrated = JSON.parse(fs.readFileSync(path.join(legacyAgentDir, "auth.json"), "utf-8")) as Record< - string, - Record - >; - expect(migrated.anthropic.key).toBe("$ANTHROPIC_API_KEY"); - const logMessage = String(logSpy.mock.calls[0]?.[0] ?? ""); - expect(logMessage).toContain('auth.json["anthropic"].key: ANTHROPIC_API_KEY -> $ANTHROPIC_API_KEY'); + fn(); } finally { - if (previousHome === undefined) { - delete process.env.HOME; - } else { - process.env.HOME = previousHome; - } - if (previousUserProfile === undefined) { - delete process.env.USERPROFILE; - } else { - process.env.USERPROFILE = previousUserProfile; - } - if (previousAgentDir === undefined) { - delete process.env[ENV_AGENT_DIR]; - } else { - process.env[ENV_AGENT_DIR] = previousAgentDir; - } - if (previousAnthropicKey === undefined) { - delete process.env.ANTHROPIC_API_KEY; - } else { - process.env.ANTHROPIC_API_KEY = previousAnthropicKey; - } + if (previous === undefined) delete process.env[name]; + else process.env[name] = previous; } - }); + } - it("rewrites legacy uppercase auth.json API key values to explicit env references when the env var exists", () => { - const agentDir = createAgentDir(); - const previousAnthropicKey = process.env.ANTHROPIC_API_KEY; - process.env.ANTHROPIC_API_KEY = "secret"; + function writeAuthFixture(agentDir: string): void { fs.writeFileSync( path.join(agentDir, "auth.json"), `${JSON.stringify( @@ -109,35 +65,118 @@ describe("config value env var syntax migration", () => { )}\n`, "utf-8", ); + } + + it("rewrites an implicit auth.json environment reference when that variable exists", () => { + const agentDir = createAgentDir(); + writeAuthFixture(agentDir); const logSpy = vi.spyOn(console, "log").mockImplementation(() => {}); - try { - withAgentDir(agentDir, () => runMigrations(agentDir)); + withEnv("ANTHROPIC_API_KEY", "secret", () => withAgentDir(agentDir, () => runMigrations(agentDir))); - const migrated = JSON.parse(fs.readFileSync(path.join(agentDir, "auth.json"), "utf-8")) as Record< + const migrated = JSON.parse(fs.readFileSync(path.join(agentDir, "auth.json"), "utf-8")) as Record< + string, + Record + >; + expect(migrated.anthropic.key).toBe("$ANTHROPIC_API_KEY"); + expect(migrated.openai.key).toBe("$OPENAI_API_KEY"); + expect(migrated.opencode.key).toBe("public"); + expect(migrated.github.access).toBe("ACCESS_TOKEN"); + expect(String(logSpy.mock.calls[0]?.[0] ?? "")).toContain( + 'auth.json["anthropic"].key: ANTHROPIC_API_KEY -> $ANTHROPIC_API_KEY', + ); + }); + + it("preserves an uppercase auth.json literal when no matching variable exists", () => { + const agentDir = createAgentDir(); + writeAuthFixture(agentDir); + const logSpy = vi.spyOn(console, "log").mockImplementation(() => {}); + + withEnv("ANTHROPIC_API_KEY", undefined, () => withAgentDir(agentDir, () => runMigrations(agentDir))); + + const migrated = JSON.parse(fs.readFileSync(path.join(agentDir, "auth.json"), "utf-8")) as Record< + string, + Record + >; + expect(migrated.anthropic.key).toBe("ANTHROPIC_API_KEY"); + expect(logSpy).not.toHaveBeenCalled(); + }); + + it("migrates implicit auth references in the legacy .pi agent directory", () => { + const homeDir = fs.mkdtempSync(path.join(os.tmpdir(), "atomic-legacy-pi-config-migration-test-")); + tempDirs.push(homeDir); + const legacyAgentDir = path.join(homeDir, ".pi", "agent"); + fs.mkdirSync(legacyAgentDir, { recursive: true }); + fs.writeFileSync( + path.join(legacyAgentDir, "auth.json"), + `${JSON.stringify({ anthropic: { type: "api_key", key: "ANTHROPIC_API_KEY" } }, null, 2)}\n`, + "utf-8", + ); + const previousHome = process.env.HOME; + const previousUserProfile = process.env.USERPROFILE; + const previousAgentDir = process.env[ENV_AGENT_DIR]; + delete process.env[ENV_AGENT_DIR]; + process.env.HOME = homeDir; + process.env.USERPROFILE = homeDir; + try { + withEnv("ANTHROPIC_API_KEY", "secret", () => runMigrations(homeDir)); + const migrated = JSON.parse(fs.readFileSync(path.join(legacyAgentDir, "auth.json"), "utf-8")) as Record< string, Record >; expect(migrated.anthropic.key).toBe("$ANTHROPIC_API_KEY"); - expect(migrated.openai.key).toBe("$OPENAI_API_KEY"); - expect(migrated.opencode.key).toBe("public"); - expect(migrated.github.access).toBe("ACCESS_TOKEN"); - const logMessage = String(logSpy.mock.calls[0]?.[0] ?? ""); - expect(logMessage).toContain("explicit $ENV_VAR syntax"); - expect(logMessage).toContain('auth.json["anthropic"].key: ANTHROPIC_API_KEY -> $ANTHROPIC_API_KEY'); } finally { - if (previousAnthropicKey === undefined) { - delete process.env.ANTHROPIC_API_KEY; - } else { - process.env.ANTHROPIC_API_KEY = previousAnthropicKey; - } + if (previousHome === undefined) delete process.env.HOME; + else process.env.HOME = previousHome; + if (previousUserProfile === undefined) delete process.env.USERPROFILE; + else process.env.USERPROFILE = previousUserProfile; + if (previousAgentDir === undefined) delete process.env[ENV_AGENT_DIR]; + else process.env[ENV_AGENT_DIR] = previousAgentDir; } }); + it("preserves models.json comments and formatting while migrating environment references", () => { + const agentDir = createAgentDir(); + const modelsPath = path.join(agentDir, "models.json"); + fs.writeFileSync( + modelsPath, + `{ + // keep provider notes + "providers": { + "CUSTOM_API_KEY": { + "metadata": { + "apiKey": "CUSTOM_API_KEY", + "headers": { "x-api-key": "HEADER_API_KEY" }, + }, + "baseUrl": "https://example.com/v1", + "apiKey": "CUSTOM_API_KEY", // migrate this value, not the key + "api": "openai-completions", + "headers": { "x-api-key": "HEADER_API_KEY" }, + "models": [{ "id": "CUSTOM_API_KEY", "name": "CUSTOM_API_KEY" }], + }, + }, +}\n`, + "utf-8", + ); + + withEnv("CUSTOM_API_KEY", "secret", () => + withEnv("HEADER_API_KEY", "secret", () => withAgentDir(agentDir, () => runMigrations(agentDir))), + ); + + const migrated = fs.readFileSync(modelsPath, "utf-8"); + expect(migrated).toContain("// keep provider notes"); + expect(migrated).toContain('"CUSTOM_API_KEY": {'); + expect(migrated).toContain('"metadata": {\n "apiKey": "CUSTOM_API_KEY"'); + expect(migrated).toContain('"apiKey": "$CUSTOM_API_KEY", // migrate this value, not the key'); + expect(migrated).toContain('"x-api-key": "$HEADER_API_KEY"'); + expect(migrated).toContain('"id": "CUSTOM_API_KEY"'); + expect(migrated).toContain('"name": "CUSTOM_API_KEY"'); + }); + it.each([ ["malformed", '{\n "providers": {\n'], ["blank", ""], - ])("does not throw on %s models.json during config migration", (_name, content) => { + ])("does not throw on %s models.json during migrations", async (_name, content) => { const agentDir = createAgentDir(); const modelsPath = path.join(agentDir, "models.json"); fs.writeFileSync(modelsPath, content, "utf-8"); @@ -145,52 +184,54 @@ describe("config value env var syntax migration", () => { withAgentDir(agentDir, () => expect(() => runMigrations(agentDir)).not.toThrow()); expect(fs.readFileSync(modelsPath, "utf-8")).toBe(content); - const registry = ModelRegistry.create(AuthStorage.create(path.join(agentDir, "auth.json")), modelsPath); + const registry = await createModelRegistry(AuthStorage.create(path.join(agentDir, "auth.json")), modelsPath); const loadError = registry.getError(); expect(loadError).toContain("Failed to parse models.json"); expect(loadError).toContain(`File: ${modelsPath}`); }); - it("rewrites legacy uppercase models.json API key and header values when env vars exist", () => { + it("migrates implicit models.json environment references to explicit dollar syntax", async () => { const agentDir = createAgentDir(); - const envVarNames = ["CUSTOM_API_KEY", "HEADER_API_KEY", "MODEL_API_KEY", "OVERRIDE_API_KEY"]; - const previousEnv = new Map(envVarNames.map((name) => [name, process.env[name]])); - for (const name of envVarNames) { - process.env[name] = "secret"; + const envKeys = ["CUSTOM_API_KEY", "HEADER_API_KEY", "MODEL_API_KEY", "OVERRIDE_API_KEY"]; + const savedEnv: Record = {}; + for (const key of envKeys) { + savedEnv[key] = process.env[key]; + process.env[key] = `env-${key}`; } - fs.writeFileSync( - path.join(agentDir, "models.json"), - `${JSON.stringify( - { - providers: { - "custom-provider": { - baseUrl: "https://example.com/v1", - apiKey: "CUSTOM_API_KEY", - api: "openai-completions", - headers: { - "x-api-key": "HEADER_API_KEY", - "x-literal": "literal", - }, - models: [ - { - id: "model-a", - headers: { "x-model-key": "MODEL_API_KEY" }, + + try { + fs.writeFileSync( + path.join(agentDir, "models.json"), + `${JSON.stringify( + { + providers: { + "custom-provider": { + baseUrl: "https://example.com/v1", + apiKey: "CUSTOM_API_KEY", + api: "openai-completions", + headers: { + "x-api-key": "HEADER_API_KEY", + "x-literal": "literal", + }, + models: [ + { + id: "model-a", + headers: { "x-model-key": "MODEL_API_KEY" }, + }, + ], + modelOverrides: { + "model-b": { headers: { "x-override-key": "OVERRIDE_API_KEY" } }, }, - ], - modelOverrides: { - "model-b": { headers: { "x-override-key": "OVERRIDE_API_KEY" } }, }, }, }, - }, - null, - 2, - )}\n`, - "utf-8", - ); - const logSpy = vi.spyOn(console, "log").mockImplementation(() => {}); + null, + 2, + )}\n`, + "utf-8", + ); + const logSpy = vi.spyOn(console, "log").mockImplementation(() => {}); - try { withAgentDir(agentDir, () => runMigrations(agentDir)); const migrated = JSON.parse(fs.readFileSync(path.join(agentDir, "models.json"), "utf-8")) as { @@ -210,122 +251,32 @@ describe("config value env var syntax migration", () => { expect(provider.headers?.["x-literal"]).toBe("literal"); expect(provider.models?.[0]?.headers?.["x-model-key"]).toBe("$MODEL_API_KEY"); expect(provider.modelOverrides?.["model-b"]?.headers?.["x-override-key"]).toBe("$OVERRIDE_API_KEY"); - const logMessage = String(logSpy.mock.calls[0]?.[0] ?? ""); - expect(logMessage).toContain( - 'models.json.providers["custom-provider"].apiKey: CUSTOM_API_KEY -> $CUSTOM_API_KEY', - ); - expect(logMessage).toContain( - 'models.json.providers["custom-provider"].headers["x-api-key"]: HEADER_API_KEY -> $HEADER_API_KEY', - ); - expect(logMessage).toContain( - 'models.json.providers["custom-provider"].models["model-a"].headers["x-model-key"]: MODEL_API_KEY -> $MODEL_API_KEY', - ); - expect(logMessage).toContain( - 'models.json.providers["custom-provider"].modelOverrides["model-b"].headers["x-override-key"]: OVERRIDE_API_KEY -> $OVERRIDE_API_KEY', - ); - } finally { - for (const [name, value] of previousEnv) { - if (value === undefined) { - delete process.env[name]; - } else { - process.env[name] = value; - } - } - } - }); - - it("preserves models.json comments and formatting while migrating env references", () => { - const agentDir = createAgentDir(); - const envVarNames = ["CUSTOM_API_KEY", "HEADER_API_KEY"]; - const previousEnv = new Map(envVarNames.map((name) => [name, process.env[name]])); - for (const name of envVarNames) { - process.env[name] = "secret"; - } - const modelsPath = path.join(agentDir, "models.json"); - fs.writeFileSync( - modelsPath, - `{ - // keep provider notes - "providers": { - "CUSTOM_API_KEY": { - "metadata": { - "apiKey": "CUSTOM_API_KEY", - "headers": { - "x-api-key": "HEADER_API_KEY", - }, - }, - "baseUrl": "https://example.com/v1", - "apiKey": "CUSTOM_API_KEY", // migrate this value, not the key - "api": "openai-completions", - "headers": { - "x-api-key": "HEADER_API_KEY", - }, - "models": [ - { - "id": "CUSTOM_API_KEY", - "name": "CUSTOM_API_KEY", - }, - ], - }, - }, -} -`, - "utf-8", - ); - const logSpy = vi.spyOn(console, "log").mockImplementation(() => {}); + expect(logSpy).toHaveBeenCalledWith(expect.stringContaining("Migrated API key/header environment references")); - try { - withAgentDir(agentDir, () => runMigrations(agentDir)); - - const migrated = fs.readFileSync(modelsPath, "utf-8"); - expect(migrated).toContain("// keep provider notes"); - expect(migrated).toContain('"CUSTOM_API_KEY": {'); - expect(migrated).toContain('"metadata": {\n "apiKey": "CUSTOM_API_KEY"'); - expect(migrated).toContain('"metadata": {\n "apiKey": "CUSTOM_API_KEY",\n "headers": {\n "x-api-key": "HEADER_API_KEY"'); - expect(migrated).toContain('"apiKey": "$CUSTOM_API_KEY", // migrate this value, not the key'); - expect(migrated).toContain('"x-api-key": "$HEADER_API_KEY",'); - expect(migrated).toContain('"id": "CUSTOM_API_KEY"'); - expect(migrated).toContain('"name": "CUSTOM_API_KEY"'); - expect(migrated).toContain(' },\n "models": ['); - expect(logSpy).toHaveBeenCalled(); + const registry = await createModelRegistry( + AuthStorage.create(path.join(agentDir, "auth.json")), + path.join(agentDir, "models.json"), + ); + const model = registry.find("custom-provider", "model-a"); + expect(model).toBeDefined(); + expect(await registry.getApiKeyForProvider("custom-provider")).toBe("env-CUSTOM_API_KEY"); + expect(await registry.getApiKeyAndHeaders(model!)).toMatchObject({ + ok: true, + apiKey: "env-CUSTOM_API_KEY", + headers: { + "x-api-key": "env-HEADER_API_KEY", + "x-literal": "literal", + "x-model-key": "env-MODEL_API_KEY", + }, + }); } finally { - for (const [name, value] of previousEnv) { - if (value === undefined) { - delete process.env[name]; + for (const key of envKeys) { + if (savedEnv[key] === undefined) { + delete process.env[key]; } else { - process.env[name] = value; + process.env[key] = savedEnv[key]; } } } }); - - it("preserves uppercase literal credentials when no matching env var exists", () => { - const agentDir = createAgentDir(); - const literalCredential = "AKIAIOSFODNN7EXAMPLE"; - const previousLiteralEnv = process.env[literalCredential]; - delete process.env[literalCredential]; - fs.writeFileSync( - path.join(agentDir, "auth.json"), - `${JSON.stringify({ aws: { type: "api_key", key: literalCredential } }, null, 2)}\n`, - "utf-8", - ); - const logSpy = vi.spyOn(console, "log").mockImplementation(() => {}); - - try { - withAgentDir(agentDir, () => runMigrations(agentDir)); - - const migrated = JSON.parse(fs.readFileSync(path.join(agentDir, "auth.json"), "utf-8")) as Record< - string, - Record - >; - expect(migrated.aws.key).toBe(literalCredential); - expect(logSpy).not.toHaveBeenCalled(); - } finally { - if (previousLiteralEnv === undefined) { - delete process.env[literalCredential]; - } else { - process.env[literalCredential] = previousLiteralEnv; - } - } - }); }); diff --git a/packages/coding-agent/test/constrained-sampling-capabilities.test.ts b/packages/coding-agent/test/constrained-sampling-capabilities.test.ts index 3eb6142dc..fc5ab522c 100644 --- a/packages/coding-agent/test/constrained-sampling-capabilities.test.ts +++ b/packages/coding-agent/test/constrained-sampling-capabilities.test.ts @@ -2,34 +2,13 @@ import type { Api, Model } from "@earendil-works/pi-ai/compat"; import { describe, expect, test } from "vitest"; import type { AtomicProviderCompat } from "../src/index.ts"; import { normalizeGrammarToolCapability } from "../src/core/model-capabilities.ts"; -import { loadBuiltInModels, mergeCompat } from "../src/core/model-registry-builtins.ts"; -import { validateModelsConfig } from "../src/core/model-registry-schemas.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; function compatOf(model: Model): AtomicProviderCompat | undefined { return model.compat as AtomicProviderCompat | undefined; } describe("constrained-sampling model capabilities", () => { - test("accepts strict and grammar capability metadata in layered model configuration", () => { - const config = { - providers: { - responses: { - baseUrl: "https://example.test/v1", - compat: { - supportsStrictMode: true, - supportsOpenAIGrammarTools: true, - supportsGrammarTools: true, - }, - }, - anthropic: { - baseUrl: "https://example.test", - compat: { supportsStrictTools: true, supportsTemperature: false, allowEmptySignature: true }, - }, - }, - }; - expect(validateModelsConfig.Check(config)).toBe(true); - }); - test("maps the Atomic alias to pi-ai without changing unsupported or unknown metadata", () => { const unknown = { supportsStrictMode: false } satisfies AtomicProviderCompat; expect(normalizeGrammarToolCapability(undefined)).toBeUndefined(); @@ -53,28 +32,27 @@ describe("constrained-sampling model capabilities", () => { ).toEqual({ supportsOpenAIGrammarTools: false, supportsGrammarTools: false }); }); - test("preserves strict capability fields while merging model overrides", () => { - const merged = mergeCompat( - { supportsStrictMode: true, supportsOpenAIGrammarTools: true }, - { supportsStrictMode: false, supportsGrammarTools: true }, - ) as AtomicProviderCompat; - expect(merged).toEqual({ + test("preserves unrelated strict capability fields while normalizing the grammar alias", () => { + expect( + normalizeGrammarToolCapability({ supportsStrictMode: false, supportsGrammarTools: true }), + ).toEqual({ supportsStrictMode: false, - supportsOpenAIGrammarTools: true, supportsGrammarTools: true, + supportsOpenAIGrammarTools: true, }); }); - test("mirrors verified generated grammar capability without enabling unknown models", () => { - const models = loadBuiltInModels(new Map(), new Map()); + test("preserves pinned generated grammar capabilities without inventing the Atomic alias", async () => { + const runtime = await ModelRuntime.create({ modelsPath: null, allowModelNetwork: false }); + const models = runtime.getModels(); const capable = models.find((model) => compatOf(model)?.supportsOpenAIGrammarTools === true); expect(capable).toBeDefined(); - expect(compatOf(capable!)?.supportsGrammarTools).toBe(true); - + expect(compatOf(capable!)?.supportsGrammarTools).toBeUndefined(); const unknown = models.find((model) => compatOf(model)?.supportsOpenAIGrammarTools === undefined); expect(unknown).toBeDefined(); expect(compatOf(unknown!)?.supportsGrammarTools).toBeUndefined(); }); + test("public Atomic model compatibility type exposes the alias", () => { const model = { compat: normalizeGrammarToolCapability({ supportsGrammarTools: true }), diff --git a/packages/coding-agent/test/extensions-input-event.test.ts b/packages/coding-agent/test/extensions-input-event.test.ts index 357fb6b89..f90f84907 100644 --- a/packages/coding-agent/test/extensions-input-event.test.ts +++ b/packages/coding-agent/test/extensions-input-event.test.ts @@ -5,9 +5,10 @@ import { afterEach, beforeEach, describe, expect, it } from "vitest"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { discoverAndLoadExtensions } from "../src/core/extensions/loader.ts"; import { ExtensionRunner } from "../src/core/extensions/runner.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { SessionManager } from "../src/core/session-manager.ts"; +import { createModelRegistry } from "./model-runtime-test-utils.ts"; + describe("Input Event", () => { let tempDir: string; let extensionsDir: string; @@ -29,7 +30,7 @@ describe("Input Event", () => { for (let i = 0; i < extensions.length; i++) fs.writeFileSync(path.join(extensionsDir, `e${i}.ts`), extensions[i]); const result = await discoverAndLoadExtensions([], tempDir, tempDir); const sm = SessionManager.inMemory(); - const mr = ModelRegistry.create(AuthStorage.create(path.join(tempDir, "auth.json"))); + const mr = await createModelRegistry(AuthStorage.create(path.join(tempDir, "auth.json"))); return new ExtensionRunner(result.extensions, result.runtime, tempDir, sm, mr); } diff --git a/packages/coding-agent/test/extensions-loader-virtual-modules.test.ts b/packages/coding-agent/test/extensions-loader-virtual-modules.test.ts index e00bd2531..23eff2bd2 100644 --- a/packages/coding-agent/test/extensions-loader-virtual-modules.test.ts +++ b/packages/coding-agent/test/extensions-loader-virtual-modules.test.ts @@ -2,7 +2,7 @@ import * as fs from "node:fs"; import * as os from "node:os"; import * as path from "node:path"; import { describe, expect, it } from "vitest"; -import { extensionLoaderTestHooks, loadExtensionModule } from "../src/core/extensions/loader-virtual-modules.ts"; +import { extensionLoaderTestHooks } from "../src/core/extensions/loader-virtual-modules.ts"; type PiAiExports = { complete?: object; @@ -10,13 +10,8 @@ type PiAiExports = { StringEnum?: object; }; -type OAuthCompatExports = { - getOAuthApiKey?: object; - getOAuthProvider?: object; - getOAuthProviders?: object; - registerOAuthProvider?: object; - resetOAuthProviders?: object; -}; +// Provider-owned OAuth is exposed through provider metadata; the removed global +// OAuth registration bridge is intentionally not part of extension aliases. describe("extension loader pi-ai compat aliases", () => { it("keys root and compat specifiers to the same virtual module object", async () => { @@ -32,22 +27,6 @@ describe("extension loader pi-ai compat aliases", () => { expect(typeof compat.StringEnum).toBe("function"); }); - it("provides legacy OAuth runtime helpers through both extension aliases", async () => { - const modules = await extensionLoaderTestHooks.loadVirtualModules(); - const earendil = modules["@earendil-works/pi-ai/oauth"] as OAuthCompatExports; - const mario = modules["@mariozechner/pi-ai/oauth"] as OAuthCompatExports; - - expect(mario).toBe(earendil); - for (const name of [ - "getOAuthApiKey", - "getOAuthProvider", - "getOAuthProviders", - "registerOAuthProvider", - "resetOAuthProviders", - ] as const) { - expect(typeof earendil[name]).toBe("function"); - } - }); it("maps root and compat specifiers to the same jiti alias path", () => { const aliases = extensionLoaderTestHooks.getAliases(); @@ -57,26 +36,6 @@ describe("extension loader pi-ai compat aliases", () => { expect(aliases["@mariozechner/pi-ai"]).toBe(aliases["@earendil-works/pi-ai/compat"]); }); - it("maps both OAuth specifiers to Atomic's populated compatibility entry", () => { - const aliases = extensionLoaderTestHooks.getAliases(); - expect(aliases["@mariozechner/pi-ai/oauth"]).toBe(aliases["@earendil-works/pi-ai/oauth"]); - expect(aliases["@earendil-works/pi-ai/oauth"]).toMatch(/oauth-compat\.js$/); - }); - - it("loads the populated OAuth bridge through the real Jiti extension path", async () => { - const tmp = fs.mkdtempSync(path.join(os.tmpdir(), "atomic-oauth-extension-")); - const extensionPath = path.join(tmp, "extension.ts"); - fs.writeFileSync( - extensionPath, - `import { getOAuthProviders } from "@earendil-works/pi-ai/oauth";\nexport default () => typeof getOAuthProviders;\n`, - ); - try { - const factory = await loadExtensionModule(extensionPath); - expect((factory as (() => string) | undefined)?.()).toBe("function"); - } finally { - fs.rmSync(tmp, { recursive: true, force: true }); - } - }); it("confirms compat is the legacy API surface while root stays core-only", async () => { const root = (await import("@earendil-works/pi-ai")) as PiAiExports; diff --git a/packages/coding-agent/test/extensions-runner/context-error-renderer-flags.suite.ts b/packages/coding-agent/test/extensions-runner/context-error-renderer-flags.suite.ts index bc12fe4e8..319c1d2f0 100644 --- a/packages/coding-agent/test/extensions-runner/context-error-renderer-flags.suite.ts +++ b/packages/coding-agent/test/extensions-runner/context-error-renderer-flags.suite.ts @@ -17,6 +17,7 @@ import type { } from "../../src/core/extensions/types.ts"; import { KeybindingsManager, type KeyId } from "../../src/core/keybindings.ts"; import { ModelRegistry } from "../../src/core/model-registry.ts"; +import { ModelRuntime } from "../../src/core/model-runtime.ts"; import { SessionManager } from "../../src/core/session-manager.ts"; describe("ExtensionRunner", () => { @@ -26,13 +27,13 @@ describe("ExtensionRunner", () => { let modelRegistry: ModelRegistry; const defaultKeybindings = new KeybindingsManager().getEffectiveConfig(); - beforeEach(() => { + beforeEach(async () => { tempDir = fs.mkdtempSync(path.join(os.tmpdir(), "pi-runner-test-")); extensionsDir = path.join(tempDir, "extensions"); fs.mkdirSync(extensionsDir); sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(path.join(tempDir, "auth.json")); - modelRegistry = ModelRegistry.create(authStorage); + modelRegistry = new ModelRegistry(await ModelRuntime.create({ credentials: authStorage, modelsPath: null })); }); afterEach(() => { diff --git a/packages/coding-agent/test/extensions-runner/lifecycle-tool-result.suite.ts b/packages/coding-agent/test/extensions-runner/lifecycle-tool-result.suite.ts index 3a4ecc4f8..1e00eb634 100644 --- a/packages/coding-agent/test/extensions-runner/lifecycle-tool-result.suite.ts +++ b/packages/coding-agent/test/extensions-runner/lifecycle-tool-result.suite.ts @@ -17,6 +17,7 @@ import type { } from "../../src/core/extensions/types.ts"; import { KeybindingsManager, type KeyId } from "../../src/core/keybindings.ts"; import { ModelRegistry } from "../../src/core/model-registry.ts"; +import { ModelRuntime } from "../../src/core/model-runtime.ts"; import { SessionManager } from "../../src/core/session-manager.ts"; describe("ExtensionRunner", () => { @@ -26,13 +27,13 @@ describe("ExtensionRunner", () => { let modelRegistry: ModelRegistry; const defaultKeybindings = new KeybindingsManager().getEffectiveConfig(); - beforeEach(() => { + beforeEach(async () => { tempDir = fs.mkdtempSync(path.join(os.tmpdir(), "pi-runner-test-")); extensionsDir = path.join(tempDir, "extensions"); fs.mkdirSync(extensionsDir); sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(path.join(tempDir, "auth.json")); - modelRegistry = ModelRegistry.create(authStorage); + modelRegistry = new ModelRegistry(await ModelRuntime.create({ credentials: authStorage, modelsPath: null })); }); afterEach(() => { diff --git a/packages/coding-agent/test/extensions-runner/project-trust.suite.ts b/packages/coding-agent/test/extensions-runner/project-trust.suite.ts index e6291a0d3..68aa379f5 100644 --- a/packages/coding-agent/test/extensions-runner/project-trust.suite.ts +++ b/packages/coding-agent/test/extensions-runner/project-trust.suite.ts @@ -17,6 +17,7 @@ import type { } from "../../src/core/extensions/types.ts"; import { KeybindingsManager, type KeyId } from "../../src/core/keybindings.ts"; import { ModelRegistry } from "../../src/core/model-registry.ts"; +import { ModelRuntime } from "../../src/core/model-runtime.ts"; import { SessionManager } from "../../src/core/session-manager.ts"; describe("ExtensionRunner", () => { @@ -26,13 +27,13 @@ describe("ExtensionRunner", () => { let modelRegistry: ModelRegistry; const defaultKeybindings = new KeybindingsManager().getEffectiveConfig(); - beforeEach(() => { + beforeEach(async () => { tempDir = fs.mkdtempSync(path.join(os.tmpdir(), "pi-runner-test-")); extensionsDir = path.join(tempDir, "extensions"); fs.mkdirSync(extensionsDir); sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(path.join(tempDir, "auth.json")); - modelRegistry = ModelRegistry.create(authStorage); + modelRegistry = new ModelRegistry(await ModelRuntime.create({ credentials: authStorage, modelsPath: null })); }); afterEach(() => { diff --git a/packages/coding-agent/test/extensions-runner/provider-command-handlers.suite.ts b/packages/coding-agent/test/extensions-runner/provider-command-handlers.suite.ts index 731c99f44..16c793bdd 100644 --- a/packages/coding-agent/test/extensions-runner/provider-command-handlers.suite.ts +++ b/packages/coding-agent/test/extensions-runner/provider-command-handlers.suite.ts @@ -17,6 +17,7 @@ import type { } from "../../src/core/extensions/types.ts"; import { KeybindingsManager, type KeyId } from "../../src/core/keybindings.ts"; import { ModelRegistry } from "../../src/core/model-registry.ts"; +import { ModelRuntime } from "../../src/core/model-runtime.ts"; import { SessionManager } from "../../src/core/session-manager.ts"; describe("ExtensionRunner", () => { @@ -24,15 +25,17 @@ describe("ExtensionRunner", () => { let extensionsDir: string; let sessionManager: SessionManager; let modelRegistry: ModelRegistry; + let modelRuntime: ModelRuntime; const defaultKeybindings = new KeybindingsManager().getEffectiveConfig(); - beforeEach(() => { + beforeEach(async () => { tempDir = fs.mkdtempSync(path.join(os.tmpdir(), "pi-runner-test-")); extensionsDir = path.join(tempDir, "extensions"); fs.mkdirSync(extensionsDir); sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(path.join(tempDir, "auth.json")); - modelRegistry = ModelRegistry.create(authStorage); + modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); + modelRegistry = new ModelRegistry(modelRuntime); }); afterEach(() => { @@ -115,7 +118,7 @@ describe("ExtensionRunner", () => { expect(errors).toEqual([ '/tmp/broken-extension.ts: Provider broken-provider: "api" is required when registering streamSimple.', ]); - await expect(modelRegistry.refresh()).resolves.toMatchObject({ aborted: false }); + await expect(modelRuntime.refresh()).resolves.toMatchObject({ aborted: false, errors: new Map() }); }); it("pre-bind unregister removes all queued registrations for a provider", () => { diff --git a/packages/coding-agent/test/extensions-runner/shortcut-conflicts.suite.ts b/packages/coding-agent/test/extensions-runner/shortcut-conflicts.suite.ts index 77d67ba1b..01f9f5340 100644 --- a/packages/coding-agent/test/extensions-runner/shortcut-conflicts.suite.ts +++ b/packages/coding-agent/test/extensions-runner/shortcut-conflicts.suite.ts @@ -17,6 +17,7 @@ import type { } from "../../src/core/extensions/types.ts"; import { KeybindingsManager, type KeyId } from "../../src/core/keybindings.ts"; import { ModelRegistry } from "../../src/core/model-registry.ts"; +import { ModelRuntime } from "../../src/core/model-runtime.ts"; import { SessionManager } from "../../src/core/session-manager.ts"; describe("ExtensionRunner", () => { @@ -26,13 +27,13 @@ describe("ExtensionRunner", () => { let modelRegistry: ModelRegistry; const defaultKeybindings = new KeybindingsManager().getEffectiveConfig(); - beforeEach(() => { + beforeEach(async () => { tempDir = fs.mkdtempSync(path.join(os.tmpdir(), "pi-runner-test-")); extensionsDir = path.join(tempDir, "extensions"); fs.mkdirSync(extensionsDir); sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(path.join(tempDir, "auth.json")); - modelRegistry = ModelRegistry.create(authStorage); + modelRegistry = new ModelRegistry(await ModelRuntime.create({ credentials: authStorage, modelsPath: null })); }); afterEach(() => { diff --git a/packages/coding-agent/test/extensions-runner/tool-command-collection.suite.ts b/packages/coding-agent/test/extensions-runner/tool-command-collection.suite.ts index a374eb4ae..73350a05e 100644 --- a/packages/coding-agent/test/extensions-runner/tool-command-collection.suite.ts +++ b/packages/coding-agent/test/extensions-runner/tool-command-collection.suite.ts @@ -17,6 +17,7 @@ import type { } from "../../src/core/extensions/types.ts"; import { KeybindingsManager, type KeyId } from "../../src/core/keybindings.ts"; import { ModelRegistry } from "../../src/core/model-registry.ts"; +import { ModelRuntime } from "../../src/core/model-runtime.ts"; import { SessionManager } from "../../src/core/session-manager.ts"; describe("ExtensionRunner", () => { @@ -26,13 +27,13 @@ describe("ExtensionRunner", () => { let modelRegistry: ModelRegistry; const defaultKeybindings = new KeybindingsManager().getEffectiveConfig(); - beforeEach(() => { + beforeEach(async () => { tempDir = fs.mkdtempSync(path.join(os.tmpdir(), "pi-runner-test-")); extensionsDir = path.join(tempDir, "extensions"); fs.mkdirSync(extensionsDir); sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(path.join(tempDir, "auth.json")); - modelRegistry = ModelRegistry.create(authStorage); + modelRegistry = new ModelRegistry(await ModelRuntime.create({ credentials: authStorage, modelsPath: null })); }); afterEach(() => { diff --git a/packages/coding-agent/test/footer-width.test.ts b/packages/coding-agent/test/footer-width.test.ts index d7d5c730f..53877dc98 100644 --- a/packages/coding-agent/test/footer-width.test.ts +++ b/packages/coding-agent/test/footer-width.test.ts @@ -1,9 +1,8 @@ -import { sep } from "node:path"; import { visibleWidth } from "@earendil-works/pi-tui"; import { beforeAll, describe, expect, it } from "vitest"; import type { AgentSession } from "../src/core/agent-session.ts"; import type { ReadonlyFooterDataProvider } from "../src/core/footer-data-provider.ts"; -import { FooterComponent, UsageMeterComponent, formatCwdForFooter } from "../src/modes/interactive/components/footer.ts"; +import { FooterComponent, formatCwdForFooter, UsageMeterComponent } from "../src/modes/interactive/components/footer.ts"; import { initTheme, theme } from "../src/modes/interactive/theme/theme.ts"; import { stripAnsi } from "../src/utils/ansi.ts"; @@ -22,22 +21,48 @@ function createSession(options: { reasoning?: boolean; thinkingLevel?: string; usage?: AssistantUsage; + branchUsage?: AssistantUsage; + compactionUsage?: AssistantUsage; + toolUsage?: AssistantUsage; contextPercent?: number; contextWindow?: number; }): AgentSession { const usage = options.usage; - const entries = - usage === undefined - ? [] - : [ - { - type: "message", - message: { - role: "assistant", - usage, - }, - }, - ]; + const entries: Array> = []; + + if (usage !== undefined) { + entries.push({ + type: "message", + message: { + role: "assistant", + usage, + }, + }); + } + + if (options.branchUsage !== undefined) { + entries.push({ + type: "branch_summary", + usage: options.branchUsage, + }); + } + + if (options.compactionUsage !== undefined) { + entries.push({ + type: "compaction", + usage: options.compactionUsage, + }); + } + + if (options.toolUsage !== undefined) { + entries.push({ + type: "message", + message: { + role: "toolResult", + usage: options.toolUsage, + }, + }); + } const session = { state: { @@ -58,11 +83,11 @@ function createSession(options: { contextWindow: options.contextWindow ?? 200_000, percent: options.contextPercent ?? 12.3, }), - modelRegistry: { + modelRuntime: { isUsingOAuth: () => false, }, settingsManager: { - getCodexFastModeSettings: () => ({ chat: false, workflow: false }), + getCodexFastModeSettings: () => ({ enabled: false }), }, }; @@ -90,59 +115,47 @@ describe("formatCwdForFooter", () => { it("abbreviates the home directory and descendants", () => { expect(formatCwdForFooter("/home/user", "/home/user")).toBe("~"); - expect(formatCwdForFooter("/home/user/project", "/home/user")).toBe(`~${sep}project`); + expect(formatCwdForFooter("/home/user/project", "/home/user")).toBe("~/project"); }); }); + describe("UsageMeterComponent context color", () => { beforeAll(() => { initTheme(undefined, false); }); it("renders over-limit auto-compacted context usage as a warning", () => { - const session = createSession({ - sessionName: "", - contextPercent: 101.2, - contextWindow: 200_000, - }); - const usageMeter = new UsageMeterComponent(session); + const usageMeter = new UsageMeterComponent( + createSession({ sessionName: "", contextPercent: 101.2, contextWindow: 200_000 }), + ); const [line] = usageMeter.render(120); - expect(line).toContain(theme.fg("warning", "101.2%/200k (auto)")); expect(line).not.toContain(theme.fg("error", "101.2%/200k (auto)")); }); it("renders near-limit auto-compacted context usage as a warning", () => { - const session = createSession({ - sessionName: "", - contextPercent: 95.4, - contextWindow: 200_000, - }); - const usageMeter = new UsageMeterComponent(session); + const usageMeter = new UsageMeterComponent( + createSession({ sessionName: "", contextPercent: 95.4, contextWindow: 200_000 }), + ); const [line] = usageMeter.render(120); - expect(line).toContain(theme.fg("warning", "95.4%/200k (auto)")); expect(line).not.toContain(theme.fg("error", "95.4%/200k (auto)")); }); it("keeps over-limit context usage red when auto-compaction is disabled", () => { - const session = createSession({ - sessionName: "", - contextPercent: 101.2, - contextWindow: 200_000, - }); - const usageMeter = new UsageMeterComponent(session); + const usageMeter = new UsageMeterComponent( + createSession({ sessionName: "", contextPercent: 101.2, contextWindow: 200_000 }), + ); usageMeter.setAutoCompactEnabled(false); const [line] = usageMeter.render(120); - expect(line).toContain(theme.fg("error", "101.2%/200k")); expect(line).not.toContain(theme.fg("warning", "101.2%/200k")); }); }); - describe("FooterComponent width handling", () => { beforeAll(() => { initTheme(undefined, false); @@ -159,61 +172,99 @@ describe("FooterComponent width handling", () => { } }); - it("shows the latest cache hit rate when cache usage is present", () => { + it("keeps stats line within width for wide model and provider names", () => { + const width = 60; + const session = createSession({ + sessionName: "", + modelId: "模".repeat(30), + provider: "공급자", + reasoning: true, + thinkingLevel: "high", + usage: { + input: 12_345, + output: 6_789, + cacheRead: 0, + cacheWrite: 0, + cost: { total: 1.234 }, + }, + }); + const footer = new FooterComponent(session, createFooterData(2)); + + const lines = footer.render(width); + for (const line of lines) { + expect(visibleWidth(line)).toBeLessThanOrEqual(width); + } + }); + + it("includes branch summary and tool result usage in the total cost", () => { const session = createSession({ sessionName: "", usage: { input: 100, output: 10, - cacheRead: 50, - cacheWrite: 50, - cost: { total: 0.001 }, + cacheRead: 0, + cacheWrite: 0, + cost: { total: 0.5 }, + }, + branchUsage: { + input: 20, + output: 5, + cacheRead: 0, + cacheWrite: 0, + cost: { total: 0.25 }, + }, + compactionUsage: { + input: 5, + output: 2, + cacheRead: 0, + cacheWrite: 0, + cost: { total: 0.125 }, + }, + toolUsage: { + input: 15, + output: 3, + cacheRead: 0, + cacheWrite: 0, + cost: { total: 0.375 }, }, }); const usageMeter = new UsageMeterComponent(session); - const statsText = stripAnsi(usageMeter.render(120).join("\n")); - expect(statsText).toContain("CH25.0%"); + const statsLine = usageMeter.render(120).map(stripAnsi).join("\n"); + expect(statsLine).toContain("$1.125"); }); - - it("marks Kimi Coding costs as subscription estimates", () => { + it("shows the latest cache hit rate when cache usage is present", () => { const session = createSession({ sessionName: "", - provider: "kimi-coding", usage: { input: 100, output: 10, - cacheRead: 0, - cacheWrite: 0, - cost: { total: 1.234 }, + cacheRead: 50, + cacheWrite: 50, + cost: { total: 0.001 }, }, }); const usageMeter = new UsageMeterComponent(session); - expect(stripAnsi(usageMeter.render(120).join("\n"))).toContain("$1.234 (sub)"); + const statsLine = usageMeter.render(120).map(stripAnsi).join("\n"); + expect(statsLine).toContain("CH25.0%"); }); - it("keeps stats line within width for wide model and provider names", () => { - const width = 60; + + it("marks Kimi Coding costs as subscription estimates", () => { const session = createSession({ sessionName: "", - modelId: "模".repeat(30), - provider: "공급자", - reasoning: true, - thinkingLevel: "high", + provider: "kimi-coding", usage: { - input: 12_345, - output: 6_789, + input: 100, + output: 10, cacheRead: 0, cacheWrite: 0, cost: { total: 1.234 }, }, }); - const footer = new FooterComponent(session, createFooterData(2)); + const usageMeter = new UsageMeterComponent(session); - const lines = footer.render(width); - for (const line of lines) { - expect(visibleWidth(line)).toBeLessThanOrEqual(width); - } + expect(usageMeter.render(120).map(stripAnsi).join("\n")).toContain("$1.234 (sub)"); }); }); diff --git a/packages/coding-agent/test/interactive-auth-login.test.ts b/packages/coding-agent/test/interactive-auth-login.test.ts index a691e4540..3b6821992 100644 --- a/packages/coding-agent/test/interactive-auth-login.test.ts +++ b/packages/coding-agent/test/interactive-auth-login.test.ts @@ -21,14 +21,9 @@ describe("interactive API-key login persistence failures", () => { const showStatus = vi.fn(); const completeProviderAuthentication = vi.fn(); const editor = {}; + const login = vi.fn(async () => { throw saveError; }); const harness = { - session: { - model: undefined, - modelRegistry: { - authStorage: { set: vi.fn(() => { throw saveError; }) }, - getCustomApiKeyAuth: () => undefined, - }, - }, + session: { model: undefined, modelRuntime: { login } }, ui: { setFocus: vi.fn(), requestRender: vi.fn() }, editorContainer: { clear: vi.fn(), addChild: vi.fn() }, editor, @@ -44,10 +39,10 @@ describe("interactive API-key login persistence failures", () => { ) => Promise; await showApiKeyLoginDialog.call(harness, "example", "Example Provider"); - expect(harness.session.modelRegistry.authStorage.set).toHaveBeenCalledWith("example", { - type: "api_key", - key: "secret-key", - }); + expect(login).toHaveBeenCalledWith("example", "api_key", expect.objectContaining({ + prompt: expect.any(Function), + notify: expect.any(Function), + })); expect(completeProviderAuthentication).not.toHaveBeenCalled(); expect(showStatus).not.toHaveBeenCalled(); expect(showError).toHaveBeenCalledWith( @@ -64,10 +59,7 @@ describe("interactive OAuth cancellation", () => { const completeProviderAuthentication = vi.fn(); const editor = {}; const harness = { - session: { - model: undefined, - modelRegistry: { authStorage: { getOAuthProviders: () => [{ id: "kimi-coding", usesCallbackServer: false }] } }, - }, + session: { model: undefined }, runtimeHost: { loginOAuthProvider: async () => { throw new DOMException("The operation was aborted.", "AbortError"); } }, ui: { setFocus: vi.fn(), requestRender: vi.fn() }, editorContainer: { clear: vi.fn(), addChild: vi.fn() }, @@ -91,10 +83,7 @@ describe("interactive OAuth cancellation", () => { const completeProviderAuthentication = vi.fn(async () => {}); const editor = {}; const harness = { - session: { - model: undefined, - modelRegistry: { authStorage: { getOAuthProviders: () => [{ id: "corp-oauth", usesCallbackServer: false }] } }, - }, + session: { model: undefined }, runtimeHost: { loginOAuthProvider: async () => ({ modelsRefreshed: true }) }, ui: { setFocus: vi.fn(), requestRender: vi.fn() }, editorContainer: { clear: vi.fn(), addChild: vi.fn() }, @@ -114,15 +103,64 @@ describe("interactive OAuth cancellation", () => { ); }); - it("keeps a post-login refresh AbortError visible", async () => { - const refreshFailure = new DOMException("catalog refresh aborted", "AbortError"); - const showError = vi.fn(); + it("honors transported callback metadata and resolves manual redirect input", async () => { + const showManualInput = vi + .spyOn(LoginDialogComponent.prototype, "showManualInput") + .mockResolvedValue("https://localhost/callback?code=manual"); + const completeProviderAuthentication = vi.fn(async () => {}); const editor = {}; + const addedChildren: object[] = []; + const loginOAuthProvider = vi.fn(async (_provider: string, callbacks: { + onAuth(info: { url: string; instructions?: string }): void; + onManualCodeInput?(): Promise; + }) => { + callbacks.onAuth({ url: "https://corp.invalid/login" }); + expect(await callbacks.onManualCodeInput?.()).toBe("https://localhost/callback?code=manual"); + return { modelsRefreshed: true }; + }); const harness = { session: { model: undefined, - modelRegistry: { authStorage: { getOAuthProviders: () => [{ id: "corp-oauth", usesCallbackServer: false }] } }, + modelRuntime: { + getOAuthProviderMetadata: () => [{ + id: "corp-oauth", + name: "Corp OAuth", + loginLabel: "Sign in to Corp", + usesCallbackServer: true, + }], + }, }, + runtimeHost: { loginOAuthProvider }, + ui: { setFocus: vi.fn(), requestRender: vi.fn() }, + editorContainer: { + clear: vi.fn(), + addChild: vi.fn((child: object) => addedChildren.push(child)), + }, + editor, + showError: vi.fn(), + completeProviderAuthentication, + showOAuthLoginSelect: vi.fn(), + }; + const showLoginDialog = InteractiveModeBase.prototype.showLoginDialog as ( + this: typeof harness, providerId: string, providerName: string, + ) => Promise; + + await showLoginDialog.call(harness, "corp-oauth", "Corp OAuth"); + + expect(showManualInput).toHaveBeenCalledWith( + "Paste redirect URL below, or complete login in browser:", + ); + const dialog = addedChildren[0] as LoginDialogComponent; + expect(dialog.render(100).join("\n")).toContain("Sign in to Corp"); + expect(completeProviderAuthentication).toHaveBeenCalledOnce(); + }, 1_000); + + it("keeps a post-login refresh AbortError visible", async () => { + const refreshFailure = new DOMException("catalog refresh aborted", "AbortError"); + const showError = vi.fn(); + const editor = {}; + const harness = { + session: { model: undefined }, runtimeHost: { loginOAuthProvider: async () => ({ modelsRefreshed: false }) }, ui: { setFocus: vi.fn(), requestRender: vi.fn() }, editorContainer: { clear: vi.fn(), addChild: vi.fn() }, @@ -144,10 +182,7 @@ describe("interactive OAuth cancellation", () => { const showError = vi.fn(); const editor = {}; const harness = { - session: { - model: undefined, - modelRegistry: { authStorage: { getOAuthProviders: () => [{ id: "kimi-coding", usesCallbackServer: false }] } }, - }, + session: { model: undefined }, runtimeHost: { loginOAuthProvider: async () => { throw new Error("Kimi Code login was denied."); } }, ui: { setFocus: vi.fn(), requestRender: vi.fn() }, editorContainer: { clear: vi.fn(), addChild: vi.fn() }, @@ -180,7 +215,7 @@ describe("post-login model refresh", () => { const setupAutocompleteProvider = vi.fn(); const showStatus = vi.fn(); const harness = { - session: { modelRegistry: { refresh, getAvailable }, setModel }, + session: { modelRuntime: { refresh, getAvailableSnapshot: getAvailable }, setModel }, updateAvailableProviderCount, setupAutocompleteProvider, footer: { invalidate: vi.fn() }, @@ -202,6 +237,7 @@ describe("post-login model refresh", () => { await complete.call(harness, scenario.provider, scenario.name, scenario.authType, loggedOutModel); expect(refresh).toHaveBeenCalledOnce(); + expect(refresh).toHaveBeenCalledWith(); expect(setModel).toHaveBeenCalledWith(model); expect(refresh.mock.invocationCallOrder[0]).toBeLessThan(getAvailable.mock.invocationCallOrder[0]!); expect(setModel.mock.invocationCallOrder[0]).toBeLessThan(updateAvailableProviderCount.mock.invocationCallOrder[0]!); @@ -215,12 +251,10 @@ describe("post-login model refresh", () => { ]) { it(`completes login when the post-login model refresh ${outcome.label}`, async () => { const showStatus = vi.fn(); + const refresh = vi.fn(async () => outcome.result); const harness = { session: { - modelRegistry: { - refresh: async () => outcome.result, - getAvailable: () => [], - }, + modelRuntime: { refresh, getAvailableSnapshot: () => [] }, }, updateAvailableProviderCount: vi.fn(), setupAutocompleteProvider: vi.fn(), diff --git a/packages/coding-agent/test/interactive-catalog-startup-refresh.test.ts b/packages/coding-agent/test/interactive-catalog-startup-refresh.test.ts index fb61a422d..635cbdc3a 100644 --- a/packages/coding-agent/test/interactive-catalog-startup-refresh.test.ts +++ b/packages/coding-agent/test/interactive-catalog-startup-refresh.test.ts @@ -15,13 +15,13 @@ function fakeMode(overrides?: { const mode = { session: { scopedModels: [], - modelRegistry: { + modelRuntime: { refresh: async (options: { allowNetwork?: boolean } = {}) => { calls.refreshOptions.push(options); if (overrides?.refreshRejects) throw new Error("network refresh failed"); return { aborted: false, errors: new Map() }; }, - getAvailable: () => [ + getAvailableSnapshot: () => [ { provider: "anthropic" }, { provider: "openai" }, { provider: "openai" }, diff --git a/packages/coding-agent/test/interactive-deferred-startup.test.ts b/packages/coding-agent/test/interactive-deferred-startup.test.ts index 445159927..7b0ff8374 100644 --- a/packages/coding-agent/test/interactive-deferred-startup.test.ts +++ b/packages/coding-agent/test/interactive-deferred-startup.test.ts @@ -4,6 +4,7 @@ import { AuthStorage } from "../src/core/auth-storage.ts"; import { ModelRegistry } from "../src/core/model-registry.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { InteractiveMode } from "../src/modes/interactive/interactive-mode.ts"; +import { createInMemoryModelRegistry, createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; const claudeModel = { provider: "anthropic", @@ -38,9 +39,9 @@ describe("applyDeferredModelScope", () => { const mode = { options: { deferredModelScopePatterns: ["claude-*", "extension-only-*"] }, session: { - modelRegistry: { - getAvailable: vi.fn(async () => [claudeModel]), - find: vi.fn(), + modelRuntime: { + getAvailableSnapshot: vi.fn(() => [claudeModel]), + getModel: vi.fn(), hasConfiguredAuth: vi.fn(() => true), }, setScopedModels, @@ -61,9 +62,9 @@ describe("applyDeferredModelScope", () => { const mode = { options: { deferredModelScopePatterns: ["claude-*:high"], deferredModelScopePreserveThinking: true }, session: { - modelRegistry: { - getAvailable: vi.fn(async () => [claudeModel]), - find: vi.fn(), + modelRuntime: { + getAvailableSnapshot: vi.fn(() => [claudeModel]), + getModel: vi.fn(), hasConfiguredAuth: vi.fn(() => true), }, setScopedModels: vi.fn(), @@ -89,7 +90,7 @@ describe("retryDeferredModelRestore", () => { sessionManager: { buildSessionContext: () => ({ model: undefined }) }, session: { model: claudeModel, - modelRegistry: { hasConfiguredAuth: vi.fn(() => true) }, + modelRuntime: { hasConfiguredAuth: vi.fn(() => true) }, setModel: vi.fn(), }, showWarning: vi.fn(), @@ -97,12 +98,12 @@ describe("retryDeferredModelRestore", () => { await InteractiveMode.prototype.retryDeferredModelRestore.call(mode as never); - expect(mode.session.modelRegistry.hasConfiguredAuth).toHaveBeenCalledWith(claudeModel); + expect(mode.session.modelRuntime.hasConfiguredAuth).toHaveBeenCalledWith(claudeModel.provider); expect(mode.showWarning).not.toHaveBeenCalled(); }); it("selects an exact settings default registered during deferred extension loading", async () => { - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const registry = await createInMemoryModelRegistry(AuthStorage.inMemory()); registerExtensionModel(registry, "deferred-extension", "deferred-model"); const settingsManager = SettingsManager.inMemory({ defaultProvider: "deferred-extension", @@ -116,7 +117,7 @@ describe("retryDeferredModelRestore", () => { options: { modelFallbackMessage: genericUnsupportedWarning }, settingsManager, sessionManager: { buildSessionContext: () => ({ model: undefined }) }, - session: { model: undefined, modelRegistry: registry, setModel, setThinkingLevel }, + session: { model: undefined, modelRuntime: getModelRuntime(registry), setModel, setThinkingLevel }, showWarning, }; @@ -129,8 +130,8 @@ describe("retryDeferredModelRestore", () => { it("uses normal fallback after a model-less extension provider registers", async () => { const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const registry = ModelRegistry.inMemory(authStorage); + await authStorage.modify("openai", async () => ({ type: "api_key", key: "test-key" })); + const registry = await createInMemoryModelRegistry(authStorage); registry.registerProvider("deferred-extension", { api: "openai-completions", streamSimple: () => { @@ -147,7 +148,7 @@ describe("retryDeferredModelRestore", () => { options: { modelFallbackMessage: genericUnsupportedWarning }, settingsManager, sessionManager: { buildSessionContext: () => ({ model: undefined }) }, - session: { model: undefined, modelRegistry: registry, setModel, setThinkingLevel: vi.fn() }, + session: { model: undefined, modelRuntime: getModelRuntime(registry), setModel, setThinkingLevel: vi.fn() }, showWarning, }; @@ -160,8 +161,8 @@ describe("retryDeferredModelRestore", () => { it("publishes the final generic warning when the settings provider remains unsupported", async () => { const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const registry = ModelRegistry.inMemory(authStorage); + await authStorage.modify("openai", async () => ({ type: "api_key", key: "test-key" })); + const registry = await createInMemoryModelRegistry(authStorage); const settingsManager = SettingsManager.inMemory({ defaultProvider: "absent-extension", defaultModel: "missing-model", @@ -172,7 +173,7 @@ describe("retryDeferredModelRestore", () => { options: { modelFallbackMessage: genericUnsupportedWarning }, settingsManager, sessionManager: { buildSessionContext: () => ({ model: undefined }) }, - session: { model: undefined, modelRegistry: registry, setModel, setThinkingLevel: vi.fn() }, + session: { model: undefined, modelRuntime: getModelRuntime(registry), setModel, setThinkingLevel: vi.fn() }, showWarning, }; @@ -183,7 +184,7 @@ describe("retryDeferredModelRestore", () => { }); it("uses ordinary no-model guidance for a supported provider when nothing is available", async () => { - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const registry = await createInMemoryModelRegistry(AuthStorage.inMemory()); registry.registerProvider("deferred-extension", { api: "openai-completions", streamSimple: () => { @@ -199,7 +200,7 @@ describe("retryDeferredModelRestore", () => { options: { modelFallbackMessage: genericUnsupportedWarning }, settingsManager, sessionManager: { buildSessionContext: () => ({ model: undefined }) }, - session: { model: undefined, modelRegistry: registry, setModel: vi.fn(), setThinkingLevel: vi.fn() }, + session: { model: undefined, modelRuntime: getModelRuntime(registry), setModel: vi.fn(), setThinkingLevel: vi.fn() }, showWarning, }; @@ -227,10 +228,10 @@ describe("retryDeferredModelRestore", () => { }, session: { model: sameProviderTemplate, - modelRegistry: { - find: vi.fn(() => exactModel), - getAvailable: vi.fn(async () => [sameProviderTemplate]), - hasConfiguredAuth: vi.fn((model) => model !== exactModel), + modelRuntime: { + getModel: vi.fn(() => exactModel), + getAvailableSnapshot: vi.fn(() => [sameProviderTemplate]), + hasConfiguredAuth: vi.fn(() => false), }, setModel, }, @@ -239,7 +240,7 @@ describe("retryDeferredModelRestore", () => { await InteractiveMode.prototype.retryDeferredModelRestore.call(mode as never); - expect(mode.session.modelRegistry.getAvailable).not.toHaveBeenCalled(); + expect(mode.session.modelRuntime.getAvailableSnapshot).not.toHaveBeenCalled(); expect(setModel).not.toHaveBeenCalled(); expect(showWarning).toHaveBeenCalledWith("Could not restore saved model", undefined); }); diff --git a/packages/coding-agent/test/interactive-mode-anthropic-warning.test.ts b/packages/coding-agent/test/interactive-mode-anthropic-warning.test.ts index 10e351928..5213561b3 100644 --- a/packages/coding-agent/test/interactive-mode-anthropic-warning.test.ts +++ b/packages/coding-agent/test/interactive-mode-anthropic-warning.test.ts @@ -7,19 +7,21 @@ function createSettingsManager(warnings: { anthropicExtraUsage?: boolean } = {}) }; } +function createModelRuntime(credential: { type: "oauth" } | undefined, apiKey?: string) { + return { + checkAuth: vi.fn().mockResolvedValue(credential), + isUsingOAuth: vi.fn().mockReturnValue(credential?.type === "oauth"), + getAuth: vi.fn().mockResolvedValue(apiKey ? { auth: { apiKey } } : credential ? { auth: credential } : undefined), + }; +} + describe("InteractiveMode.maybeWarnAboutAnthropicSubscriptionAuth", () => { test("warns once when Anthropic subscription auth is detected", async () => { + const modelRuntime = createModelRuntime(undefined, "sk-ant-oat01-test"); const fakeThis: any = { anthropicSubscriptionWarningShown: false, settingsManager: createSettingsManager(), - session: { - modelRegistry: { - authStorage: { - get: vi.fn().mockReturnValue(undefined), - }, - getApiKeyForProvider: vi.fn().mockResolvedValue("sk-ant-oat01-test"), - }, - }, + session: { modelRuntime }, showWarning: vi.fn(), }; @@ -31,44 +33,38 @@ describe("InteractiveMode.maybeWarnAboutAnthropicSubscriptionAuth", () => { }); expect(fakeThis.showWarning).toHaveBeenCalledTimes(1); - expect(fakeThis.session.modelRegistry.getApiKeyForProvider).toHaveBeenCalledTimes(1); + expect(modelRuntime.getAuth).toHaveBeenCalledTimes(1); }); test("warns when Anthropic OAuth is stored even if token refresh lookup would fail", async () => { + const modelRuntime = createModelRuntime({ type: "oauth" }); + modelRuntime.getAuth.mockRejectedValue(new Error("stale Anthropic OAuth credential")); const fakeThis: any = { anthropicSubscriptionWarningShown: false, settingsManager: createSettingsManager(), - session: { - modelRegistry: { - authStorage: { - get: vi.fn().mockReturnValue({ type: "oauth" }), - }, - getApiKeyForProvider: vi.fn().mockResolvedValue(undefined), - }, - }, + session: { modelRuntime }, + showError: vi.fn(), showWarning: vi.fn(), }; + await (InteractiveMode as any).prototype.maybeWarnAboutAnthropicSubscriptionAuth.call(fakeThis, { + provider: "anthropic", + }); await (InteractiveMode as any).prototype.maybeWarnAboutAnthropicSubscriptionAuth.call(fakeThis, { provider: "anthropic", }); expect(fakeThis.showWarning).toHaveBeenCalledTimes(1); - expect(fakeThis.session.modelRegistry.getApiKeyForProvider).not.toHaveBeenCalled(); + expect(fakeThis.showError).not.toHaveBeenCalled(); + expect(modelRuntime.getAuth).not.toHaveBeenCalled(); }); test("does not warn for non-Anthropic models", async () => { + const modelRuntime = createModelRuntime(undefined); const fakeThis: any = { anthropicSubscriptionWarningShown: false, settingsManager: createSettingsManager(), - session: { - modelRegistry: { - authStorage: { - get: vi.fn(), - }, - getApiKeyForProvider: vi.fn(), - }, - }, + session: { modelRuntime }, showWarning: vi.fn(), }; @@ -77,21 +73,15 @@ describe("InteractiveMode.maybeWarnAboutAnthropicSubscriptionAuth", () => { }); expect(fakeThis.showWarning).not.toHaveBeenCalled(); - expect(fakeThis.session.modelRegistry.getApiKeyForProvider).not.toHaveBeenCalled(); + expect(modelRuntime.getAuth).not.toHaveBeenCalled(); }); test("does not warn when Anthropic extra usage warning is disabled", async () => { + const modelRuntime = createModelRuntime(undefined); const fakeThis: any = { anthropicSubscriptionWarningShown: false, settingsManager: createSettingsManager({ anthropicExtraUsage: false }), - session: { - modelRegistry: { - authStorage: { - get: vi.fn(), - }, - getApiKeyForProvider: vi.fn(), - }, - }, + session: { modelRuntime }, showWarning: vi.fn(), }; @@ -100,7 +90,7 @@ describe("InteractiveMode.maybeWarnAboutAnthropicSubscriptionAuth", () => { }); expect(fakeThis.showWarning).not.toHaveBeenCalled(); - expect(fakeThis.session.modelRegistry.authStorage.get).not.toHaveBeenCalled(); - expect(fakeThis.session.modelRegistry.getApiKeyForProvider).not.toHaveBeenCalled(); + expect(modelRuntime.checkAuth).not.toHaveBeenCalled(); + expect(modelRuntime.getAuth).not.toHaveBeenCalled(); }); }); diff --git a/packages/coding-agent/test/interactive-mode-status-autocomplete.suite.ts b/packages/coding-agent/test/interactive-mode-status-autocomplete.suite.ts index 43e1da924..29c39a7e8 100644 --- a/packages/coding-agent/test/interactive-mode-status-autocomplete.suite.ts +++ b/packages/coding-agent/test/interactive-mode-status-autocomplete.suite.ts @@ -148,15 +148,15 @@ describe("InteractiveMode /fast autocomplete", () => { models: Model[], scopedModels: Model[] = [], options: { - hasConfiguredAuth?: (model: Model) => boolean; + hasConfiguredAuth?: (provider: string) => boolean; extensionCommands?: ExtensionCommandFixture[]; } = {}, ): AutocompleteProvider { const fakeThis: any = { session: { scopedModels: scopedModels.map((model) => ({ model })), - modelRegistry: { - getAvailable: vi.fn(() => models), + modelRuntime: { + getAvailableSnapshot: vi.fn(() => models), hasConfiguredAuth: vi.fn(options.hasConfiguredAuth ?? (() => true)), }, promptTemplates: [], @@ -219,7 +219,7 @@ describe("InteractiveMode /fast autocomplete", () => { const scopedModel = createModel(scopedProvider, `${scopedProvider}-unauthenticated`); const labels = await slashLabels( createProvider([createModel("openai", "available-openai")], [scopedModel], { - hasConfiguredAuth: (model) => model !== scopedModel, + hasConfiguredAuth: (provider) => provider !== scopedProvider, }), ); @@ -271,7 +271,7 @@ describe("InteractiveMode.createBaseAutocompleteProvider", () => { type AutocompleteHost = { session: { scopedModels: []; - modelRegistry: { getAvailable: () => [] }; + modelRuntime: { getAvailableSnapshot: () => [] }; promptTemplates: []; extensionRunner: { getRegisteredCommands: () => [] }; resourceLoader: { getSkills: () => { skills: [] } }; @@ -289,7 +289,7 @@ describe("InteractiveMode.createBaseAutocompleteProvider", () => { const fakeThis: AutocompleteHost = { session: { scopedModels: [], - modelRegistry: { getAvailable: () => [] }, + modelRuntime: { getAvailableSnapshot: () => [] }, promptTemplates: [], extensionRunner: { getRegisteredCommands: () => [] }, resourceLoader: { getSkills: () => ({ skills: [] }) }, @@ -338,7 +338,7 @@ describe("InteractiveMode.createBaseAutocompleteProvider", () => { type FakeInteractiveMode = { session: { scopedModels: Array<{ model: TestModel }>; - modelRegistry: { getAvailable: () => TestModel[] }; + modelRuntime: { getAvailableSnapshot: () => TestModel[] }; promptTemplates: []; extensionRunner: { getRegisteredCommands: () => [] }; resourceLoader: { getSkills: () => { skills: [] } }; @@ -361,7 +361,7 @@ describe("InteractiveMode.createBaseAutocompleteProvider", () => { const fakeThis: FakeInteractiveMode = { session: { scopedModels: [], - modelRegistry: { getAvailable: () => models }, + modelRuntime: { getAvailableSnapshot: () => models }, promptTemplates: [], extensionRunner: { getRegisteredCommands: () => [] }, resourceLoader: { getSkills: () => ({ skills: [] }) }, @@ -393,7 +393,7 @@ describe("InteractiveMode deferred workflow autocomplete", () => { deferredStartupPending: true, session: { scopedModels: [], - modelRegistry: { getAvailable: () => [] }, + modelRuntime: { getAvailableSnapshot: () => [] }, promptTemplates: [], extensionRunner: { getRegisteredCommands: () => [] }, resourceLoader: { getSkills: () => ({ skills: [] }) }, diff --git a/packages/coding-agent/test/interactive-model-routing-auth-failure.test.ts b/packages/coding-agent/test/interactive-model-routing-auth-failure.test.ts new file mode 100644 index 000000000..fb90e4486 --- /dev/null +++ b/packages/coding-agent/test/interactive-model-routing-auth-failure.test.ts @@ -0,0 +1,35 @@ +import { describe, expect, it, vi } from "vitest"; +import { InteractiveMode } from "../src/modes/interactive/interactive-mode.ts"; + +describe("Anthropic subscription warning auth failures", () => { + it("ignores a stale non-OAuth lookup without rejecting the advisory warning path", async () => { + const authFailure = new Error("stale Anthropic credential"); + const fakeThis = { + anthropicSubscriptionWarningShown: false, + settingsManager: { getWarnings: () => ({}) }, + session: { + modelRuntime: { + getAuth: vi.fn().mockRejectedValue(authFailure), + isUsingOAuth: vi.fn().mockReturnValue(false), + }, + }, + showError: vi.fn(), + showWarning: vi.fn(), + }; + + const maybeWarn = (InteractiveMode as never as { + prototype: { + maybeWarnAboutAnthropicSubscriptionAuth: ( + this: typeof fakeThis, + model: { provider: string }, + ) => Promise; + }; + }).prototype.maybeWarnAboutAnthropicSubscriptionAuth; + await expect(maybeWarn.call(fakeThis, { provider: "anthropic" })).resolves.toBeUndefined(); + + expect(fakeThis.session.modelRuntime.getAuth).toHaveBeenCalledTimes(1); + expect(fakeThis.showError).not.toHaveBeenCalled(); + expect(fakeThis.showWarning).not.toHaveBeenCalled(); + expect(fakeThis.anthropicSubscriptionWarningShown).toBe(false); + }); +}); diff --git a/packages/coding-agent/test/interactive-model-routing-offline.test.ts b/packages/coding-agent/test/interactive-model-routing-offline.test.ts index addf36d30..124159b43 100644 --- a/packages/coding-agent/test/interactive-model-routing-offline.test.ts +++ b/packages/coding-agent/test/interactive-model-routing-offline.test.ts @@ -23,9 +23,9 @@ test("offline model candidate startup restores caches without catalog network re const mode = { session: { scopedModels: [], - modelRegistry: { + modelRuntime: { refresh, - getAvailable: () => [], + getAvailableSnapshot: () => [], }, }, }; @@ -42,7 +42,7 @@ test("footer provider count uses the current snapshot without refreshing catalog const mode = { session: { scopedModels: [], - modelRegistry: { refresh, getAvailable: () => [{ provider: "one" }, { provider: "one" }, { provider: "two" }] }, + modelRuntime: { refresh, getAvailableSnapshot: () => [{ provider: "one" }, { provider: "one" }, { provider: "two" }] }, }, footerDataProvider: { setAvailableProviderCount }, }; @@ -57,7 +57,7 @@ test("offline scoped-model selector refresh stays cache-only", async () => { const refresh = vi.fn(async () => ({ aborted: false, errors: new Map() })); const showStatus = vi.fn(); const mode = { - session: { scopedModels: [], modelRegistry: { refresh, getAvailable: () => [] } }, + session: { scopedModels: [], modelRuntime: { refresh, getAvailableSnapshot: () => [] } }, settingsManager: { getEnabledModels: () => undefined }, showStatus, }; diff --git a/packages/coding-agent/test/interactive-startup-resource-gate.suite.ts b/packages/coding-agent/test/interactive-startup-resource-gate.suite.ts index 22ae4b526..893a7d547 100644 --- a/packages/coding-agent/test/interactive-startup-resource-gate.suite.ts +++ b/packages/coding-agent/test/interactive-startup-resource-gate.suite.ts @@ -43,7 +43,7 @@ function configureDeferredGateMode(mode: InteractiveMode): void { reload: async () => {}, resourceLoader: { getThemes: () => ({ themes: [] }) }, extensionRunner: {}, - modelRegistry: { getError: () => undefined }, + modelRuntime: { getError: () => undefined }, }, }, options: { configurable: true, value: {} }, diff --git a/packages/coding-agent/test/interactive-startup-resource-ordering.test.ts b/packages/coding-agent/test/interactive-startup-resource-ordering.test.ts index dc1aee408..ef502fa9e 100644 --- a/packages/coding-agent/test/interactive-startup-resource-ordering.test.ts +++ b/packages/coding-agent/test/interactive-startup-resource-ordering.test.ts @@ -62,7 +62,7 @@ async function renderStartupWithEarlyNotify( reload: async () => {}, resourceLoader: { getThemes: () => ({ themes: [] }) }, extensionRunner: {}, - modelRegistry: { getError: () => undefined }, + modelRuntime: { getError: () => undefined }, }, }, options: { value: {} }, @@ -119,7 +119,7 @@ function configureDeferredMode(mode: InteractiveMode): void { reload: async () => {}, resourceLoader: { getThemes: () => ({ themes: [] }) }, extensionRunner: {}, - modelRegistry: { getError: () => undefined }, + modelRuntime: { getError: () => undefined }, }, }, options: { configurable: true, value: {} }, @@ -228,8 +228,12 @@ const anthropicSubscriptionNotice: StartupNotice = { configurable: true, value: { model: { provider: "anthropic" }, - modelRegistry: { - authStorage: { get: () => ({ type: "oauth" }) }, + modelRuntime: { + getAuth: async () => ({ + auth: { apiKey: "sk-ant-oat01-test" }, + credential: { type: "oauth", access: "test", expires: Date.now() + 60_000 }, + }), + isUsingOAuth: () => true, getError: () => undefined, }, reload: async () => {}, diff --git a/packages/coding-agent/test/llama-extension.test.ts b/packages/coding-agent/test/llama-extension.test.ts new file mode 100644 index 000000000..4251ad9d0 --- /dev/null +++ b/packages/coding-agent/test/llama-extension.test.ts @@ -0,0 +1,295 @@ +import { once } from "node:events"; +import { createServer, type RequestListener, type Server, type ServerResponse } from "node:http"; +import type { AddressInfo } from "node:net"; +import type { AuthContext, AuthPrompt, ModelsStoreEntry } from "@earendil-works/pi-ai"; +import { afterEach, describe, expect, it } from "vitest"; +import { createEventBus } from "../src/core/event-bus.ts"; +import { createExtensionRuntime, loadExtensionFromFactory } from "../src/core/extensions/loader.ts"; +import { LlamaClient, type LlamaProgress, normalizeLlamaServerUrl } from "../src/extensions/llama/client.ts"; +import { findHuggingFaceToken, HuggingFaceClient } from "../src/extensions/llama/huggingface.ts"; +import llamaExtension from "../src/extensions/llama/index.ts"; +import { createLlamaProvider, LLAMA_PROVIDER_ID } from "../src/extensions/llama/provider.ts"; + +const servers: Server[] = []; + +async function listen(handler: RequestListener): Promise<{ server: Server; url: string }> { + const server = createServer(handler); + servers.push(server); + server.listen(0, "127.0.0.1"); + await once(server, "listening"); + const address = server.address() as AddressInfo; + return { server, url: `http://127.0.0.1:${address.port}` }; +} + +function json(response: ServerResponse, value: unknown): void { + response.writeHead(200, { "Content-Type": "application/json" }); + response.end(JSON.stringify(value)); +} + +afterEach(async () => { + await Promise.all( + servers.splice(0).map( + (server) => + new Promise((resolve) => { + server.close(() => resolve()); + server.closeAllConnections(); + }), + ), + ); +}); + +describe("llama.cpp extension", () => { + it("registers a native provider and /llama command", async () => { + const runtime = createExtensionRuntime(); + const extension = await loadExtensionFromFactory( + llamaExtension, + process.cwd(), + createEventBus(), + runtime, + "", + ); + + expect(extension.commands.get("llama")?.description).toBe("Manage llama.cpp router models"); + // Atomic's loader tracks native and legacy registrations in one pending list. + expect( + runtime.pendingProviderRegistrations.map((entry) => ("provider" in entry ? entry.provider.id : entry.name)), + ).toEqual([LLAMA_PROVIDER_ID]); + }); + + it("normalizes management and inference URLs", () => { + expect(normalizeLlamaServerUrl("http://127.0.0.1:8080/v1/")).toBe("http://127.0.0.1:8080"); + expect(normalizeLlamaServerUrl("https://example.com/prefix/v1")).toBe("https://example.com/prefix"); + expect(() => normalizeLlamaServerUrl("file:///tmp/llama")).toThrow("http or https"); + }); + + it("exposes only loaded models with router metadata", () => { + const controller = createLlamaProvider(); + controller.setCatalog( + [ + { + id: "loaded", + status: { value: "loaded", args: ["llama-server", "--n-gpu-layers", "999"] }, + architecture: { input_modalities: ["text", "image"] }, + meta: { n_ctx: 65536, n_ctx_train: 131072 }, + }, + { id: "unloaded", status: { value: "unloaded" } }, + { id: "loading", status: { value: "loading" } }, + ], + "http://localhost:8080", + ); + + expect(controller.provider.getModels()).toEqual([ + expect.objectContaining({ + id: "loaded", + baseUrl: "http://localhost:8080/v1", + contextWindow: 65536, + maxTokens: 65536, + input: ["text", "image"], + }), + ]); + }); + + it("persists and restores loaded models for cache-only startup refreshes", async () => { + let cachedEntry: ModelsStoreEntry | undefined; + const store = { + read: async () => cachedEntry, + write: async (entry: ModelsStoreEntry) => { + cachedEntry = structuredClone(entry); + }, + delete: async () => { + cachedEntry = undefined; + }, + }; + const { url } = await listen((request, response) => { + if (request.url === "/models") { + json(response, { + data: [ + { id: "loaded", status: { value: "loaded" }, meta: { n_ctx: 32768 } }, + { id: "unloaded", status: { value: "unloaded" } }, + ], + }); + return; + } + response.writeHead(404).end(); + }); + + const first = createLlamaProvider(); + await first.provider.refreshModels?.({ + credential: { type: "api_key", key: "local", env: { LLAMA_BASE_URL: url } }, + store, + allowNetwork: true, + }); + expect(first.provider.getModels().map((model) => model.id)).toEqual(["loaded"]); + expect(cachedEntry?.models.map((model) => model.id)).toEqual(["loaded"]); + + const second = createLlamaProvider(); + await second.provider.refreshModels?.({ + credential: { type: "api_key", key: "local", env: { LLAMA_BASE_URL: url } }, + store, + allowNetwork: false, + }); + expect(second.provider.getModels()).toEqual([ + expect.objectContaining({ id: "loaded", baseUrl: `${url}/v1`, contextWindow: 32768 }), + ]); + }); + + it("stays dormant until configured and stores URL plus optional key", async () => { + const { provider } = createLlamaProvider(); + const auth = provider.auth.apiKey!; + const emptyContext: AuthContext = { + env: async () => undefined, + fileExists: async () => false, + }; + expect(await auth.check?.({ ctx: emptyContext })).toBeUndefined(); + expect(await auth.resolve({ ctx: emptyContext })).toBeUndefined(); + + const { url } = await listen((request, response) => { + expect(request.headers.authorization).toBe("Bearer secret"); + json(response, { data: [] }); + }); + const answers = [url, "secret"]; + const credential = await auth.login!({ + prompt: async (_prompt: AuthPrompt) => answers.shift()!, + notify: () => {}, + }); + expect(credential).toEqual({ + type: "api_key", + key: "secret", + env: { LLAMA_BASE_URL: url }, + }); + expect(await auth.resolve({ ctx: emptyContext, credential })).toEqual({ + auth: { apiKey: "secret", baseUrl: `${url}/v1` }, + env: { LLAMA_BASE_URL: url }, + source: "stored credential", + }); + }); + + it("searches Hugging Face and reads quantizations plus access requirements", async () => { + const { url } = await listen((request, response) => { + expect(request.headers.authorization).toBe("Bearer hf-secret"); + if (request.url?.startsWith("/api/models?")) { + const requestUrl = new URL(request.url, "http://localhost"); + expect(requestUrl.searchParams.get("search")).toBe("qwen coder"); + expect(requestUrl.searchParams.get("filter")).toBe("gguf"); + expect(requestUrl.searchParams.get("sort")).toBe("downloads"); + json(response, [{ id: "owner/model-GGUF", downloads: 1200 }]); + return; + } + if (request.url === "/api/models/owner/model-GGUF?blobs=true") { + json(response, { + id: "owner/model-GGUF", + gated: "manual", + siblings: [ + { rfilename: "model-Q5_K_M.gguf", size: 6000 }, + { rfilename: "model-Q4_K_M-00001-of-00002.gguf", size: 2000 }, + { rfilename: "model-Q4_K_M-00002-of-00002.gguf", size: 3000 }, + { rfilename: "mmproj-F16.gguf", size: 1000 }, + ], + }); + return; + } + response.writeHead(404).end(); + }); + const client = new HuggingFaceClient("hf-secret", url); + + expect(await client.search("qwen coder")).toEqual([{ id: "owner/model-GGUF", downloads: 1200 }]); + expect(await client.details("owner/model-GGUF")).toEqual({ + id: "owner/model-GGUF", + gated: "manual", + quantizations: [ + { name: "Q4_K_M", size: 5000 }, + { name: "Q5_K_M", size: 6000 }, + ], + }); + expect(await findHuggingFaceToken({ HF_TOKEN: " hf-secret " })).toBe("hf-secret"); + }); + + it("loads with SSE progress and waits for the loaded catalog state", async () => { + let status: "unloaded" | "loading" | "loaded" = "unloaded"; + const streams = new Set(); + const send = (event: unknown) => { + for (const response of streams) response.write(`data: ${JSON.stringify(event)}\n\n`); + }; + const { url } = await listen((request, response) => { + if (request.url === "/models/sse") { + response.writeHead(200, { "Content-Type": "text/event-stream" }); + streams.add(response); + request.on("close", () => streams.delete(response)); + return; + } + if (request.url === "/models/load" && request.method === "POST") { + status = "loading"; + json(response, { success: true }); + setTimeout(() => { + send({ + model: "test-model", + event: "status_change", + data: { + status: "loading", + progress: { stages: ["text_model", "mmproj_model"], current: "text_model", value: 0.5 }, + }, + }); + status = "loaded"; + send({ model: "test-model", event: "status_change", data: { status: "loaded" } }); + }, 20); + return; + } + if (request.url === "/models") { + json(response, { data: [{ id: "test-model", status: { value: status } }] }); + return; + } + response.writeHead(404).end(); + }); + + const progress: string[] = []; + const model = await new LlamaClient(url).loadAndWait("test-model", (entry) => progress.push(entry.message)); + expect(model.status.value).toBe("loaded"); + expect(progress).toContain("Loading text model"); + }); + + it("downloads with byte progress and returns the refreshed catalog", async () => { + let status: "missing" | "downloading" | "unloaded" = "missing"; + const streams = new Set(); + const send = (event: unknown) => { + for (const response of streams) response.write(`data: ${JSON.stringify(event)}\n\n`); + }; + const { url } = await listen((request, response) => { + if (request.url === "/models/sse") { + response.writeHead(200, { "Content-Type": "text/event-stream" }); + streams.add(response); + request.on("close", () => streams.delete(response)); + return; + } + if (request.url === "/models" && request.method === "POST") { + status = "downloading"; + json(response, { success: true }); + setTimeout(() => { + send({ + model: "owner/repo:Q4_K_M", + event: "download_progress", + data: { progress: { "https://example/model.gguf": { done: 512, total: 1024 } } }, + }); + status = "unloaded"; + send({ model: "owner/repo:Q4_K_M", event: "download_finished", data: {} }); + }, 20); + return; + } + if (request.url?.startsWith("/models")) { + json(response, { + data: status === "missing" ? [] : [{ id: "owner/repo:Q4_K_M", status: { value: status } }], + }); + return; + } + response.writeHead(404).end(); + }); + + const progress: LlamaProgress[] = []; + const models = await new LlamaClient(url).downloadAndWait("owner/repo:Q4_K_M", (entry) => progress.push(entry)); + expect(models).toEqual([{ id: "owner/repo:Q4_K_M", status: { value: "unloaded" } }]); + expect(progress).toContainEqual({ + message: "Downloading model", + ratio: 0.5, + detail: "512 B / 1.00 KiB", + }); + }); +}); diff --git a/packages/coding-agent/test/main-runtime-api-key.test.ts b/packages/coding-agent/test/main-runtime-api-key.test.ts new file mode 100644 index 000000000..62a557096 --- /dev/null +++ b/packages/coding-agent/test/main-runtime-api-key.test.ts @@ -0,0 +1,25 @@ +import { expect, it, vi } from "vitest"; +import { applyCliRuntimeApiKey } from "../src/main-runtime-api-key.ts"; + +it("applies --api-key without a networked catalog refresh before reading available models", async () => { + const calls: string[] = []; + const modelRuntime = { + setRuntimeApiKey: vi.fn(async (_providerId: string, _apiKey: string, options: { allowNetwork?: boolean }) => { + calls.push(`set:${String(options.allowNetwork)}`); + }), + getAvailable: vi.fn(async () => { + calls.push("available"); + return []; + }), + }; + + await applyCliRuntimeApiKey(modelRuntime, "custom-provider", "runtime-secret"); + + expect(modelRuntime.setRuntimeApiKey).toHaveBeenCalledWith( + "custom-provider", + "runtime-secret", + { allowNetwork: false }, + ); + expect(modelRuntime.getAvailable).toHaveBeenCalledOnce(); + expect(calls).toEqual(["set:false", "available"]); +}); diff --git a/packages/coding-agent/test/model-auth-compatibility.test.ts b/packages/coding-agent/test/model-auth-compatibility.test.ts index 472814274..6fbb76265 100644 --- a/packages/coding-agent/test/model-auth-compatibility.test.ts +++ b/packages/coding-agent/test/model-auth-compatibility.test.ts @@ -1,438 +1,305 @@ -import type { Credential } from "@earendil-works/pi-ai"; -import { mkdtempSync, rmSync, writeFileSync } from "node:fs"; -import { tmpdir } from "node:os"; -import { join } from "node:path"; -import { afterEach, describe, expect, test, vi } from "vitest"; -import { FileAuthStorageBackend } from "../src/core/auth-storage-backends.ts"; +import { type AuthType, type CredentialStore, InMemoryCredentialStore } from "@earendil-works/pi-ai"; +import { describe, expect, it } from "vitest"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { getModelRequestAuth } from "../src/core/model-registry-auth.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; -import { FileModelsStore } from "../src/core/models-store.ts"; -import { - getOAuthApiKey, - getOAuthProvider, - registerOAuthProvider, - resetLegacyOAuthProviders, -} from "../src/core/oauth-provider-bridge.ts"; - -afterEach(() => { - vi.restoreAllMocks(); - resetLegacyOAuthProviders(); -}); - -describe("Pi 0.80.10 model auth compatibility", () => { - test("preserves the synchronous AuthStorage API behind an async CredentialStore adapter", async () => { - const storage = AuthStorage.inMemory({ alpha: { type: "api_key", key: "one" } }); - const credentials = storage.asCredentialStore(); - - expect(storage.list()).toEqual(["alpha"]); - expect(await credentials.list()).toEqual([{ providerId: "alpha", type: "api_key" }]); - await credentials.modify("alpha", async (current) => ({ ...current!, type: "api_key", key: "two" })); - expect(storage.get("alpha")).toEqual({ type: "api_key", key: "two" }); - await credentials.delete("alpha"); - expect(storage.list()).toEqual([]); +import { ModelRuntime } from "../src/core/model-runtime.ts"; + +function authOptions(runtime: ModelRuntime, type?: AuthType) { + return runtime + .getProviders() + .flatMap((provider) => [ + ...(!type || type === "oauth" + ? provider.auth.oauth + ? [{ type: "oauth" as const, provider, method: provider.auth.oauth }] + : [] + : []), + ...(!type || type === "api_key" + ? provider.auth.apiKey + ? [{ type: "api_key" as const, provider, method: provider.auth.apiKey }] + : [] + : []), + ]); +} + +function testModel(id: string) { + return { + id, + name: id, + reasoning: false, + input: ["text"] as ("text" | "image")[], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 10000, + maxTokens: 1000, + }; +} + +describe("ModelRuntime auth options", () => { + it("accepts a pi-ai CredentialStore", async () => { + const credentials = new InMemoryCredentialStore(); + await credentials.modify("anthropic", async () => ({ type: "api_key", key: "stored-key" })); + const runtime = await ModelRuntime.create({ credentials, modelsPath: null }); + + expect((await runtime.getAuth("anthropic"))?.auth.apiKey).toBe("stored-key"); }); - test("exposes runtime-only API keys to provider-owned auth without persisting them", async () => { - const storage = AuthStorage.inMemory({ alpha: { type: "api_key", key: "stored" } }); - storage.setRuntimeApiKey("runtime", "temporary"); - const credentials = storage.asCredentialStore(); + it("scopes provider availability reads and records refresh failures", async () => { + const base = new InMemoryCredentialStore(); + const reads: string[] = []; + let failReads = false; + const credentials: CredentialStore = { + read: async (providerId) => { + reads.push(providerId); + if (failReads) throw new Error(`read failed for ${providerId}`); + return base.read(providerId); + }, + list: () => base.list(), + modify: (providerId, fn) => base.modify(providerId, fn), + delete: (providerId) => base.delete(providerId), + }; + const runtime = await ModelRuntime.create({ credentials, modelsPath: null }); - expect(await credentials.read("runtime")).toEqual({ type: "api_key", key: "temporary" }); - expect(await credentials.list()).toEqual([ - { providerId: "alpha", type: "api_key" }, - { providerId: "runtime", type: "api_key" }, - ]); - expect(storage.get("runtime")).toBeUndefined(); - expect(storage.list()).toEqual(["alpha"]); + reads.length = 0; + await runtime.getAvailable("anthropic"); + expect(new Set(reads)).toEqual(new Set(["anthropic"])); + + failReads = true; + await expect(runtime.getAvailable("anthropic")).rejects.toThrow("Credential store read failed for anthropic"); + expect(runtime.getError()).toContain("Availability refresh: Credential store read failed for anthropic"); + + failReads = false; + await runtime.getAvailable(); + expect(runtime.getError()).toBeUndefined(); }); - test("preserves callback login and provider-owned label metadata", () => { - const providers = AuthStorage.inMemory().getOAuthProviders(); - expect(providers.find((provider) => provider.id === "anthropic")?.usesCallbackServer).toBe(true); - expect(providers.find((provider) => provider.id === "openai-codex")?.usesCallbackServer).toBe(true); - expect(providers.find((provider) => provider.id === "xai")?.loginLabel).toBe( - "Sign in with SuperGrok or X Premium", + it("projects provider-owned methods, names, and status", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + const options = authOptions(runtime); + + expect(options).toEqual( + expect.arrayContaining([ + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "amazon-bedrock", name: "Amazon Bedrock" }), + method: expect.objectContaining({ name: "AWS credentials or bearer token" }), + }), + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "google-vertex", name: "Google Vertex AI" }), + method: expect.objectContaining({ name: "Google Cloud credentials" }), + }), + expect.objectContaining({ + type: "oauth", + provider: expect.objectContaining({ id: "anthropic", name: "Anthropic" }), + }), + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "cloudflare-ai-gateway", name: "Cloudflare AI Gateway" }), + }), + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "cloudflare-workers-ai", name: "Cloudflare Workers AI" }), + }), + ]), ); - const anthropic = providers.find((provider) => provider.id === "anthropic")!; - expect(typeof anthropic.login).toBe("function"); - expect(typeof anthropic.refreshToken).toBe("function"); - expect(anthropic.getApiKey({ refresh: "r", access: "a", expires: 1 })).toBe("a"); - expect(getOAuthProvider("anthropic")?.name).toBe(anthropic.name); + expect(authOptions(runtime, "api_key").every((option) => option.type === "api_key")).toBe(true); + expect(authOptions(runtime, "oauth").every((option) => option.type === "oauth")).toBe(true); + expect(options.some((option) => option.provider.id === "openai-codex" && option.type === "api_key")).toBe(false); }); - test("preserves the legacy credentials-map OAuth API-key helper", async () => { - const original = { refresh: "old-refresh", access: "old-access", expires: 0 }; - const refreshed = { refresh: "new-refresh", access: "new-access", expires: Date.now() + 60_000 }; - const refreshToken = vi.fn(async () => refreshed); - registerOAuthProvider({ - id: "legacy-probe", - name: "Legacy Probe", - login: async () => original, - refreshToken, - getApiKey: (credentials) => `key:${credentials.access}`, + it("attaches the provider's active auth status to every method option", async () => { + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory({ + anthropic: { + type: "oauth", + access: "access", + refresh: "refresh", + expires: Date.now() + 60_000, + }, + }), + modelsPath: null, }); - expect(await getOAuthApiKey("legacy-probe", { - decoy: { refresh: "decoy", access: "decoy", expires: 1 }, - "legacy-probe": original, - })).toEqual({ newCredentials: refreshed, apiKey: "key:new-access" }); - expect(refreshToken).toHaveBeenCalledWith(original); - expect(await getOAuthApiKey("legacy-probe", {})).toBeNull(); - await expect(getOAuthApiKey("missing-provider", {})).rejects.toThrow("Unknown OAuth provider"); - const upstreamError = new Error("sensitive upstream detail"); - refreshToken.mockRejectedValueOnce(upstreamError); - const failure = await getOAuthApiKey("legacy-probe", { "legacy-probe": original }).catch((error) => error); - expect(failure).toMatchObject({ - message: "Failed to refresh OAuth token for legacy-probe", cause: upstreamError, - }); + const options = authOptions(runtime).filter((option) => option.provider.id === "anthropic"); + expect(options).toHaveLength(2); + expect(await runtime.checkAuth("anthropic")).toMatchObject({ type: "oauth" }); }); - test("runtime API-key overrides bypass expired stored OAuth", async () => { - const storage = AuthStorage.inMemory({ - anthropic: { type: "oauth", refresh: "expired", access: "expired", expires: 0 }, + it("constructs an API key method for an extension API-key provider", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + runtime.registerProvider("extension-api-key", { + name: "Extension API Key", + baseUrl: "https://example.test/v1", + apiKey: "$EXTENSION_TEST_API_KEY", + api: "openai-completions", + models: [testModel("extension-model")], }); - storage.setRuntimeApiKey("anthropic", "runtime-wins"); - const registry = ModelRegistry.inMemory(storage); - const model = registry.getAll().find((candidate) => candidate.provider === "anthropic")!; - await expect(registry.getApiKeyAndHeaders(model)).resolves.toMatchObject({ - ok: true, - apiKey: "runtime-wins", + const options = authOptions(runtime).filter((option) => option.provider.id === "extension-api-key"); + expect(options).toHaveLength(1); + expect(options[0]).toMatchObject({ + type: "api_key", + provider: { id: "extension-api-key", name: "Extension API Key" }, + method: { name: "API key" }, }); + expect(options[0]?.method.login).toBeTypeOf("function"); }); - test("runtime API-key overrides stored API-key request auth", async () => { - const storage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "stored-key" } }); - storage.setRuntimeApiKey("anthropic", "runtime-key"); - const registry = ModelRegistry.inMemory(storage); - const model = registry.getAll().find((candidate) => candidate.provider === "anthropic")!; - - expect(await registry.getApiKeyAndHeaders(model)).toMatchObject({ ok: true, apiKey: "runtime-key" }); - }); - - test("legacy OAuth replacement for a built-in provider bypasses built-in auth", async () => { - const storage = AuthStorage.inMemory({ - anthropic: { type: "oauth", refresh: "expired", access: "expired", expires: 0 }, - }); - const registry = ModelRegistry.inMemory(storage); - const refreshToken = vi.fn(async () => ({ - refresh: "custom-refresh", - access: "custom-access", - expires: Date.now() + 60_000, - })); - registry.registerProvider("anthropic", { - oauth: { - name: "Custom Anthropic", - login: async () => ({ refresh: "r", access: "a", expires: 1 }), - refreshToken, - getApiKey: (credential) => `legacy:${credential.access}`, - }, + it("resolves configured auth from request-scoped environment overrides", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + runtime.registerProvider("request-env-provider", { + baseUrl: "https://example.test/v1", + apiKey: "$REQUEST_SCOPED_API_KEY", + headers: { "x-request-value": "$REQUEST_SCOPED_HEADER" }, + api: "openai-completions", + models: [testModel("request-env-model")], }); - const fetchSpy = vi.spyOn(globalThis, "fetch"); - const model = registry.getAll().find((candidate) => candidate.provider === "anthropic")!; - - expect(await registry.getApiKeyAndHeaders(model)).toMatchObject({ ok: true, apiKey: "legacy:custom-access" }); - expect(refreshToken).toHaveBeenCalledOnce(); - expect(fetchSpy).not.toHaveBeenCalled(); - }); - test("global legacy OAuth registration replaces built-in provider auth", async () => { - const storage = AuthStorage.inMemory({ - anthropic: { type: "oauth", refresh: "expired", access: "expired", expires: 0 }, + const auth = await runtime.getAuth("request-env-provider", { + env: { REQUEST_SCOPED_API_KEY: "request-key", REQUEST_SCOPED_HEADER: "request-header" }, }); - const refreshToken = vi.fn(async () => ({ - refresh: "global-refresh", - access: "global-access", - expires: Date.now() + 60_000, - })); - registerOAuthProvider({ - id: "anthropic", - name: "Global Anthropic", - login: async () => ({ refresh: "r", access: "a", expires: 1 }), - refreshToken, - getApiKey: (credential) => `global:${credential.access}`, - }); - const fetchSpy = vi.spyOn(globalThis, "fetch"); - const registry = ModelRegistry.inMemory(storage); - const model = registry.getAll().find((candidate) => candidate.provider === "anthropic")!; - expect(await registry.getApiKeyAndHeaders(model)).toMatchObject({ ok: true, apiKey: "global:global-access" }); - expect(refreshToken).toHaveBeenCalledOnce(); - expect(fetchSpy).not.toHaveBeenCalled(); + expect(auth?.auth).toEqual({ apiKey: "request-key", headers: { "x-request-value": "request-header" } }); }); - test("global legacy OAuth registration authorizes built-in catalog refresh", async () => { - const storage = AuthStorage.inMemory({ - anthropic: { type: "oauth", refresh: "expired", access: "expired", expires: 0 }, - }); - const refreshToken = vi.fn(async () => ({ - refresh: "catalog-refresh", - access: "catalog-access", - expires: Date.now() + 60_000, - })); - registerOAuthProvider({ - id: "anthropic", - name: "Global Anthropic Catalog", - login: async () => ({ refresh: "r", access: "a", expires: 1 }), - refreshToken, - getApiKey: (credential) => `catalog:${credential.access}`, + it("lets an explicit Authorization header override authHeader case-insensitively", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + let capturedHeaders: Record | undefined; + runtime.registerProvider("auth-header-provider", { + baseUrl: "https://example.test/v1", + apiKey: "generated-key", + authHeader: true, + api: "openai-completions", + streamSimple: (_model, _context, options) => { + capturedHeaders = options?.headers; + throw new Error("captured"); + }, + models: [testModel("auth-header-model")], }); - const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response(undefined, { status: 404 })); - const registry = ModelRegistry.inMemory(storage); + const model = runtime.getModel("auth-header-provider", "auth-header-model"); + expect(model).toBeDefined(); - const result = await registry.refresh({ force: true }); + await runtime.completeSimple(model!, { messages: [] }, { headers: { authorization: "Explicit token" } }); - expect(result.errors.has("anthropic")).toBe(false); - expect(refreshToken).toHaveBeenCalledOnce(); - expect(fetchSpy.mock.calls.some(([url]) => new URL(String(url)).hostname === "platform.claude.com")).toBe(false); - expect(storage.get("anthropic")).toMatchObject({ type: "oauth", access: "catalog-access" }); + expect(capturedHeaders).toEqual({ authorization: "Explicit token" }); }); - test("cache-only catalog restore does not refresh global legacy OAuth", async () => { - const storage = AuthStorage.inMemory({ - anthropic: { type: "oauth", refresh: "cache-refresh", access: "cache-access", expires: 0 }, - }); - const refreshToken = vi.fn(async () => ({ - refresh: "unexpected-refresh", - access: "unexpected-access", - expires: Date.now() + 60_000, - })); - registerOAuthProvider({ - id: "anthropic", - name: "Cache-only Anthropic", - login: async () => ({ refresh: "r", access: "a", expires: 1 }), - refreshToken, - getApiKey: (credential) => `cache:${credential.access}`, + it("transforms fully assembled headers once without forwarding the transform", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + let capturedHeaders: Record | undefined; + let transforms = 0; + runtime.registerProvider("header-provider", { + baseUrl: "https://example.test/v1", + apiKey: "generated-key", + authHeader: true, + headers: { "x-provider": "provider" }, + api: "openai-completions", + streamSimple: (_model, _context, options) => { + expect(options).not.toHaveProperty("transformHeaders"); + capturedHeaders = options?.headers; + throw new Error("captured"); + }, + models: [{ ...testModel("header-model"), headers: { "x-model": "model" } }], }); - const fetchSpy = vi.spyOn(globalThis, "fetch"); - const registry = ModelRegistry.inMemory(storage); - - await registry.refresh({ allowNetwork: false }); - - expect(refreshToken).not.toHaveBeenCalled(); - expect(fetchSpy).not.toHaveBeenCalled(); - expect(storage.get("anthropic")).toMatchObject({ type: "oauth", access: "cache-access" }); - }); - - test("runtime-only credentials authorize provider-owned catalog refresh", async () => { - const storage = AuthStorage.inMemory(); - storage.setRuntimeApiKey("anthropic", "runtime-only"); - const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response(undefined, { status: 404 })); - const registry = ModelRegistry.inMemory(storage); + const model = runtime.getModel("header-provider", "header-model"); + expect(model).toBeDefined(); - const result = await registry.refresh({ force: true }); + await runtime.completeSimple( + model!, + { messages: [] }, + { + headers: { "x-explicit": "explicit" }, + transformHeaders: async (headers) => { + transforms++; + expect(headers).toEqual({ + Authorization: "Bearer generated-key", + "x-provider": "provider", + "x-model": "model", + "x-explicit": "explicit", + }); + return { ...headers, "x-transformed": "yes" }; + }, + }, + ); - expect(result.aborted).toBe(false); - expect(fetchSpy.mock.calls.some(([url]) => String(url).includes("/providers/anthropic"))).toBe(true); - expect(storage.get("anthropic")).toBeUndefined(); + expect(transforms).toBe(1); + expect(capturedHeaders).toEqual({ + Authorization: "Bearer generated-key", + "x-provider": "provider", + "x-model": "model", + "x-explicit": "explicit", + "x-transformed": "yes", + }); }); - test("retains credential-specific Copilot apiKey and baseUrl from provider-owned OAuth", async () => { - const storage = AuthStorage.inMemory({ - "github-copilot": { - type: "oauth", - refresh: "github-token", - access: "tid=example;proxy-ep=proxy.enterprise.example.com;", - expires: Date.now() + 60_000, + it("does not fabricate an API key method for an extension OAuth-only provider", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + runtime.registerProvider("extension-oauth", { + name: "Extension OAuth", + baseUrl: "https://example.test/v1", + api: "openai-completions", + oauth: { + name: "Extension subscription", + login: async () => ({ access: "access", refresh: "refresh", expires: Date.now() + 60_000 }), + refreshToken: async (credentials) => credentials, + getApiKey: (credentials) => credentials.access, }, + models: [testModel("extension-model")], }); - const registry = ModelRegistry.inMemory(storage); - const model = registry.getAll().find((candidate) => candidate.provider === "github-copilot"); - expect(model).toBeDefined(); - const auth = await registry.getApiKeyAndHeaders(model!); - - expect(auth).toMatchObject({ - ok: true, - apiKey: "tid=example;proxy-ep=proxy.enterprise.example.com;", - baseUrl: "https://api.enterprise.example.com", + const options = authOptions(runtime).filter((option) => option.provider.id === "extension-oauth"); + expect(options).toHaveLength(1); + expect(options[0]).toMatchObject({ + type: "oauth", + provider: { id: "extension-oauth", name: "Extension OAuth" }, + method: { name: "Extension subscription" }, }); }); - test("filters Copilot models with the effective runtime credential", () => { - const baseline = ModelRegistry.inMemory(AuthStorage.inMemory()).getAll() - .filter((model) => model.provider === "github-copilot"); - expect(baseline.length).toBeGreaterThan(1); - const storage = AuthStorage.inMemory({ - "github-copilot": { - type: "oauth", - refresh: "github-token", - access: "stored-token", - expires: Date.now() + 60_000, - availableModelIds: [baseline[0]!.id], - }, + it("runtime API-key overrides bypass expired stored OAuth", async () => { + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory({ + anthropic: { type: "oauth", access: "expired", refresh: "refresh", expires: 1 }, + }), + modelsPath: null, }); - storage.setRuntimeApiKey("github-copilot", "runtime-key"); - const registry = ModelRegistry.inMemory(storage); + runtime.setRuntimeApiKey("anthropic", "runtime-key"); + + expect((await runtime.getAuth("anthropic"))?.auth.apiKey).toBe("runtime-key"); + expect(await runtime.checkAuth("anthropic")).toMatchObject({ type: "api_key" }); + }); + + it("runtime API-key overrides stored API-key request auth without persistence", async () => { + const credentials = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "stored-key" } }); + const runtime = await ModelRuntime.create({ credentials, modelsPath: null }); + runtime.setRuntimeApiKey("anthropic", "runtime-key"); - expect(registry.getAvailable().filter((model) => model.provider === "github-copilot")).toHaveLength(baseline.length); + expect((await runtime.getAuth("anthropic"))?.auth.apiKey).toBe("runtime-key"); + expect(await credentials.read("anthropic")).toEqual({ type: "api_key", key: "stored-key" }); }); - test("refreshes OAuth for an extension provider with an initially empty catalog", async () => { - const storage = AuthStorage.inMemory({ - "catalog-only": { type: "oauth", refresh: "old-refresh", access: "old-access", expires: 0 }, + it("refreshes provider-owned OAuth extensions with an initially empty catalog", async () => { + const credentials = AuthStorage.inMemory({ + "oauth-catalog": { type: "oauth", access: "access", refresh: "refresh", expires: Date.now() + 60_000 }, }); - const registry = ModelRegistry.inMemory(storage); - const refreshToken = vi.fn(async () => ({ - refresh: "new-refresh", - access: "new-access", - expires: Date.now() + 60_000, - })); - let observedCredential: Credential | undefined; - registry.registerProvider("catalog-only", { + const runtime = await ModelRuntime.create({ credentials, modelsPath: null }); + let observedCredentialType: string | undefined; + runtime.registerProvider("oauth-catalog", { + baseUrl: "https://example.test/v1", + api: "openai-completions", oauth: { - name: "Catalog only", - login: async () => ({ refresh: "r", access: "a", expires: 1 }), - refreshToken, + name: "OAuth Catalog", + login: async () => ({ access: "access", refresh: "refresh", expires: Date.now() + 60_000 }), + refreshToken: async (credential) => credential, getApiKey: (credential) => credential.access, }, refreshModels: async ({ credential }) => { - observedCredential = credential; - return []; + observedCredentialType = credential?.type; + return [testModel("discovered")]; }, }); - const result = await registry.refresh({ force: true }); - - expect(result.errors.size).toBe(0); - expect(refreshToken).toHaveBeenCalledOnce(); - expect(observedCredential).toMatchObject({ type: "oauth", access: "new-access" }); - }); - - test("preserves provider-owned auth headers and null removals", async () => { - const storage = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(storage); - const model = { - ...registry.getAll()[0], - provider: "header-provider", - headers: { "X-Removed": "static", "X-Static": "kept-by-provider-auth" }, - }; - const auth = await getModelRequestAuth( - model, - storage, - new Map([["header-provider", { headers: { "X-Config": "config" } }]]), - new Map(), - { apiKey: "provider-key", headers: { "x-removed": null, "X-Auth": "oauth" }, baseUrl: "https://auth.example/v1" }, - ); - - expect(auth).toEqual({ - ok: true, - apiKey: "provider-key", - headers: { "x-removed": null, "X-Static": "kept-by-provider-auth", "X-Auth": "oauth", "X-Config": "config" }, - baseUrl: "https://auth.example/v1", - }); - }); - - - test("rechecks conditional file catalog writes after acquiring the persistence lock", async () => { - const directory = mkdtempSync(join(tmpdir(), "atomic-conditional-model-store-")); - try { - const path = join(directory, "models-store.json"); - const store = new FileModelsStore(path); - const model = ModelRegistry.inMemory(AuthStorage.inMemory()).getAll()[0]!; - await store.write("conditional", { models: [{ ...model, id: "seed" }] }); - let releaseLock!: () => void; - let markLocked!: () => void; - const lockGate = new Promise((resolve) => { releaseLock = resolve; }); - const locked = new Promise((resolve) => { markLocked = resolve; }); - const blocker = new FileAuthStorageBackend(path).withLockAsync(async () => { - markLocked(); - await lockGate; - return { result: undefined }; - }); - await locked; - let current = true; - const attemptedWrite = store.writeIf( - "conditional", - { models: [{ ...model, id: "late" }] }, - () => current, - ); - - current = false; - releaseLock(); - await Promise.all([blocker, attemptedWrite]); - expect((await store.read("conditional"))?.models.map((entry) => entry.id)).toEqual(["seed"]); - let releaseDeleteLock!: () => void; - let markDeleteLocked!: () => void; - const deleteLockGate = new Promise((resolve) => { releaseDeleteLock = resolve; }); - const deleteLocked = new Promise((resolve) => { markDeleteLocked = resolve; }); - const deleteBlocker = new FileAuthStorageBackend(path).withLockAsync(async () => { - markDeleteLocked(); - await deleteLockGate; - return { result: undefined }; - }); - await deleteLocked; - current = true; - const attemptedDelete = store.deleteIf("conditional", () => current); - current = false; - releaseDeleteLock(); - await Promise.all([deleteBlocker, attemptedDelete]); - expect((await store.read("conditional"))?.models.map((entry) => entry.id)).toEqual(["seed"]); - } finally { - rmSync(directory, { recursive: true, force: true }); - } - }); - test("restores persisted provider catalogs before dependent reads", async () => { - const directory = mkdtempSync(join(tmpdir(), "atomic-model-runtime-")); - try { - const storage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "test-key" } }); - const baseline = ModelRegistry.inMemory(storage).getAll().find((model) => model.provider === "anthropic")!; - await new FileModelsStore(join(directory, "models-store.json")).write("anthropic", { - models: [ - { ...baseline, name: "Refreshed Existing" }, - { ...baseline, id: "persisted-dynamic", name: "Persisted Dynamic" }, - ], - checkedAt: Date.now(), - lastModified: Date.now() + 60_000, - }); - writeFileSync(join(directory, "models.json"), JSON.stringify({ - providers: { anthropic: { baseUrl: "https://proxy.example/v1" } }, - })); - const registry = ModelRegistry.create(storage, join(directory, "models.json")); - - await registry.refresh({ allowNetwork: false }); - - expect(registry.find("anthropic", baseline.id)?.name).toBe("Refreshed Existing"); - expect(registry.find("anthropic", baseline.id)?.baseUrl).toBe("https://proxy.example/v1"); - expect(registry.find("anthropic", "persisted-dynamic")?.baseUrl).toBe("https://proxy.example/v1"); - } finally { - rmSync(directory, { recursive: true, force: true }); - } - }); - - test("retains a first persisted provider overlay when network refresh times out", async () => { - const directory = mkdtempSync(join(tmpdir(), "atomic-model-timeout-")); - try { - const storage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "test-key" } }); - const baseline = ModelRegistry.inMemory(storage).getAll().find((model) => model.provider === "anthropic")!; - await new FileModelsStore(join(directory, "models-store.json")).write("anthropic", { - models: [{ ...baseline, id: "cached-dynamic", name: "Cached Dynamic" }], - checkedAt: 0, - lastModified: Date.now() + 60_000, - }); - vi.spyOn(globalThis, "fetch").mockImplementation(async () => new Promise(() => {})); - const registry = ModelRegistry.create(storage, join(directory, "models.json")); - - const result = await registry.refresh({ timeoutMs: 5 }); - - expect(result.aborted).toBe(true); - expect(registry.find("anthropic", "cached-dynamic")?.name).toBe("Cached Dynamic"); - } finally { - rmSync(directory, { recursive: true, force: true }); - } - }); + await runtime.refresh({ allowNetwork: true }); - test("persists provider-scoped refreshed catalogs across runtime instances", async () => { - const directory = mkdtempSync(join(tmpdir(), "atomic-model-store-")); - const path = join(directory, "models-store.json"); - try { - const first = new FileModelsStore(path); - await first.write("dynamic", { models: [], checkedAt: 123 }); - const second = new FileModelsStore(path); - expect(await second.read("dynamic")).toEqual({ models: [], checkedAt: 123 }); - } finally { - rmSync(directory, { recursive: true, force: true }); - } + expect(observedCredentialType).toBe("oauth"); + expect(runtime.getModel("oauth-catalog", "discovered")).toBeDefined(); }); }); diff --git a/packages/coding-agent/test/model-registry-api-key-resolution.suite.ts b/packages/coding-agent/test/model-registry-api-key-resolution.suite.ts index 6fba8df99..16a1234c9 100644 --- a/packages/coding-agent/test/model-registry-api-key-resolution.suite.ts +++ b/packages/coding-agent/test/model-registry-api-key-resolution.suite.ts @@ -1,7 +1,7 @@ import { readFileSync, writeFileSync } from "node:fs"; import { join } from "node:path"; import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createModelRegistry } from "./model-runtime-test-utils.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; @@ -42,7 +42,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("!echo test-api-key-from-command"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBe("test-api-key-from-command"); @@ -53,7 +53,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("!echo ' spaced-key '"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBe("spaced-key"); @@ -64,7 +64,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("!printf 'line1\\nline2'"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBe("line1\nline2"); @@ -75,7 +75,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("!exit 1"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBeUndefined(); @@ -86,7 +86,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("!nonexistent-command-12345"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBeUndefined(); @@ -97,7 +97,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("!printf ''"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBeUndefined(); @@ -112,7 +112,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("$TEST_API_KEY_12345"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBe("env-api-key-value"); @@ -133,7 +133,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("literal_api_key_value"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBe("literal_api_key_value"); @@ -144,7 +144,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey("!echo 'hello world' | tr ' ' '-'"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const apiKey = await registry.getApiKeyForProvider("custom-provider"); expect(apiKey).toBe("hello-world"); @@ -161,7 +161,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(command), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); await registry.getApiKeyForProvider("custom-provider"); await registry.getApiKeyForProvider("custom-provider"); await registry.getApiKeyForProvider("custom-provider"); @@ -180,10 +180,10 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(command), }); - const registry1 = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry1 = await createModelRegistry(context.authStorage, context.modelsJsonPath); await registry1.getApiKeyForProvider("custom-provider"); - const registry2 = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry2 = await createModelRegistry(context.authStorage, context.modelsJsonPath); await registry2.getApiKeyForProvider("custom-provider"); const count = parseInt(readFileSync(counterFile, "utf-8").trim(), 10); @@ -196,7 +196,7 @@ describeModelRegistry((context) => { "provider-b": providerWithApiKey("!echo key-b"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const keyA = await registry.getApiKeyForProvider("provider-a"); const keyB = await registry.getApiKeyForProvider("provider-b"); @@ -215,7 +215,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(command), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const key1 = await registry.getApiKeyForProvider("custom-provider"); const key2 = await registry.getApiKeyForProvider("custom-provider"); @@ -226,7 +226,7 @@ describeModelRegistry((context) => { expect(count).toBe(2); }); - test("provider auth status reports apiKey environment variables from models.json", () => { + test("provider auth status reports apiKey environment variables from models.json", async () => { const envVarName = "TEST_API_KEY_STATUS_TEST_98765"; const originalEnv = process.env[envVarName]; @@ -237,7 +237,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(`$${envVarName}`), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getProviderAuthStatus("custom-provider")).toEqual({ configured: true, @@ -253,7 +253,7 @@ describeModelRegistry((context) => { } }); - test("provider auth status reports missing explicit env refs as unconfigured", () => { + test("provider auth status reports missing explicit env refs as unconfigured", async () => { const envVarName = "TEST_API_KEY_STATUS_MISSING_98765"; const originalEnv = process.env[envVarName]; delete process.env[envVarName]; @@ -263,7 +263,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(`$${envVarName}`), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getProviderAuthStatus("custom-provider")).toEqual({ configured: false, @@ -277,7 +277,7 @@ describeModelRegistry((context) => { } }); - test("missing explicit env apiKey keeps provider unavailable", () => { + test("missing explicit env apiKey keeps provider unavailable", async () => { const envVarName = "TEST_API_KEY_MISSING_AVAILABILITY_98765"; const originalEnv = process.env[envVarName]; delete process.env[envVarName]; @@ -287,7 +287,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(`$${envVarName}`), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getProviderAuthStatus("custom-provider")).toEqual({ configured: false, @@ -302,12 +302,12 @@ describeModelRegistry((context) => { } }); - test("provider auth status reports non-env apiKey values from models.json as a config key", () => { + test("provider auth status reports non-env apiKey values from models.json as a config key", async () => { writeRawModelsJson({ "custom-provider": providerWithApiKey("literal_api_key_value"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getProviderAuthStatus("custom-provider")).toEqual({ configured: true, @@ -315,7 +315,7 @@ describeModelRegistry((context) => { }); }); - test("provider auth status reports command apiKey values from models.json without executing them", () => { + test("provider auth status reports command apiKey values from models.json without executing them", async () => { const counterFile = join(context.tempDir, "status-counter"); writeFileSync(counterFile, "0"); const counterPath = toShPath(counterFile); @@ -324,7 +324,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(command), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getProviderAuthStatus("custom-provider")).toEqual({ configured: true, @@ -344,7 +344,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(`$${envVarName}`), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const key1 = await registry.getApiKeyForProvider("custom-provider"); expect(key1).toBe("first-value"); @@ -372,7 +372,7 @@ describeModelRegistry((context) => { "custom-provider": providerWithApiKey(command), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const available = registry.getAvailable(); expect(available.some((m) => m.provider === "custom-provider")).toBe(true); @@ -392,7 +392,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("custom-provider", "test-model"); expect(model).toBeDefined(); @@ -421,7 +421,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("custom-provider", "test-model"); expect(model).toBeDefined(); diff --git a/packages/coding-agent/test/model-registry-base-url-overrides.suite.ts b/packages/coding-agent/test/model-registry-base-url-overrides.suite.ts index 353873d3b..af5da428d 100644 --- a/packages/coding-agent/test/model-registry-base-url-overrides.suite.ts +++ b/packages/coding-agent/test/model-registry-base-url-overrides.suite.ts @@ -1,5 +1,5 @@ import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createModelRegistry } from "./model-runtime-test-utils.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; @@ -15,12 +15,12 @@ describeModelRegistry((context) => { emptyContext, } = context; describe("baseUrl override (no custom models)", () => { - test("overriding baseUrl keeps all built-in models", () => { + test("overriding baseUrl keeps all built-in models", async () => { writeRawModelsJson({ anthropic: overrideConfig("https://my-proxy.example.com/v1"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const anthropicModels = getModelsForProvider(registry, "anthropic"); // Should have multiple built-in models, not just one @@ -28,12 +28,12 @@ describeModelRegistry((context) => { expect(anthropicModels.some((m) => m.id.includes("claude"))).toBe(true); }); - test("overriding baseUrl changes URL on all built-in models", () => { + test("overriding baseUrl changes URL on all built-in models", async () => { writeRawModelsJson({ anthropic: overrideConfig("https://my-proxy.example.com/v1"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const anthropicModels = getModelsForProvider(registry, "anthropic"); // All models should have the new baseUrl @@ -49,7 +49,7 @@ describeModelRegistry((context) => { }), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const anthropicModels = getModelsForProvider(registry, "anthropic"); for (const model of anthropicModels) { @@ -70,7 +70,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getError()).toBeUndefined(); const anthropicModels = getModelsForProvider(registry, "anthropic"); @@ -83,14 +83,14 @@ describeModelRegistry((context) => { } }); - test("models.json baseUrl override wins over GitHub Copilot env routing", () => { + test("models.json baseUrl override wins over GitHub Copilot env routing", async () => { const previous = process.env.COPILOT_GITHUB_TOKEN; process.env.COPILOT_GITHUB_TOKEN = "github_pat_enterprise"; try { writeRawModelsJson({ "github-copilot": overrideConfig("https://copilot-proxy.example.com"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("github-copilot", "gpt-5.5"); expect(model?.baseUrl).toBe("https://copilot-proxy.example.com"); @@ -109,7 +109,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("github-copilot", "gpt-5.5"); expect(model).toBeDefined(); @@ -134,7 +134,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("github-copilot", "gpt-5.5"); expect(model).toBeDefined(); @@ -145,12 +145,12 @@ describeModelRegistry((context) => { } }); - test("baseUrl-only override does not affect other providers", () => { + test("baseUrl-only override does not affect other providers", async () => { writeRawModelsJson({ anthropic: overrideConfig("https://my-proxy.example.com/v1"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const googleModels = getModelsForProvider(registry, "google"); // Google models should still have their original baseUrl @@ -158,7 +158,7 @@ describeModelRegistry((context) => { expect(googleModels[0].baseUrl).not.toBe("https://my-proxy.example.com/v1"); }); - test("can mix baseUrl override and models merge", () => { + test("can mix baseUrl override and models merge", async () => { writeRawModelsJson({ // baseUrl-only for anthropic anthropic: overrideConfig("https://anthropic-proxy.example.com/v1"), @@ -170,7 +170,7 @@ describeModelRegistry((context) => { ), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); // Anthropic: multiple built-in models with new baseUrl const anthropicModels = getModelsForProvider(registry, "anthropic"); @@ -187,7 +187,7 @@ describeModelRegistry((context) => { writeRawModelsJson({ anthropic: overrideConfig("https://first-proxy.example.com/v1"), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(getModelsForProvider(registry, "anthropic")[0].baseUrl).toBe("https://first-proxy.example.com/v1"); diff --git a/packages/coding-agent/test/model-registry-cost-tiers.suite.ts b/packages/coding-agent/test/model-registry-cost-tiers.suite.ts index 7ce754c08..38c15bd94 100644 --- a/packages/coding-agent/test/model-registry-cost-tiers.suite.ts +++ b/packages/coding-agent/test/model-registry-cost-tiers.suite.ts @@ -1,6 +1,6 @@ import { calculateCost, type Usage } from "@earendil-works/pi-ai"; import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createInMemoryModelRegistry, createModelRegistry } from "./model-runtime-test-utils.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; const replacementTier = { @@ -24,7 +24,7 @@ function usage(input: number, output: number, cacheRead = 0, cacheWrite = 0): Us describeModelRegistry((context) => { describe("request-wide model cost tiers", () => { - test("custom models retain complete tiers and price only strictly above aggregate input threshold", () => { + test("custom models retain complete tiers and price only strictly above aggregate input threshold", async () => { context.writeRawModelsJson({ demo: { baseUrl: "https://example.com/v1", @@ -49,7 +49,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("demo", "tiered-model"); expect(registry.getError()).toBeUndefined(); expect(model?.cost.tiers).toEqual([replacementTier]); @@ -66,7 +66,7 @@ describeModelRegistry((context) => { expect(aboveThreshold.cacheRead).toBeCloseTo(0.4); }); - test("custom models reject incomplete cost tiers", () => { + test("custom models reject incomplete cost tiers", async () => { context.writeRawModelsJson({ demo: { baseUrl: "https://example.com/v1", @@ -87,12 +87,12 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.find("demo", "invalid-tier")).toBeUndefined(); expect(registry.getError()).toContain("cacheWrite"); }); - test("model overrides reject incomplete cost tiers", () => { + test("model overrides reject incomplete cost tiers", async () => { context.writeRawModelsJson({ openai: { modelOverrides: { @@ -103,45 +103,45 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getError()).toContain("cacheWrite"); }); - test("scalar built-in override preserves inherited GPT-5.6 tiers", () => { + test("scalar built-in override preserves inherited GPT-5.6 tiers", async () => { context.writeRawModelsJson({ openai: { modelOverrides: { "gpt-5.6-sol": { cost: { input: 99 } } } }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("openai", "gpt-5.6-sol"); expect(model?.cost.input).toBe(99); expect(model?.cost.output).toBeGreaterThan(0); expect(model?.cost.tiers?.length).toBeGreaterThan(0); }); - test("explicit tier override replaces inherited tiers and preserves unspecified scalar rates", () => { - const baseline = ModelRegistry.create(context.authStorage).find("openai", "gpt-5.6-sol"); + test("explicit tier override replaces inherited tiers and preserves unspecified scalar rates", async () => { + const baseline = (await createInMemoryModelRegistry(context.authStorage)).find("openai", "gpt-5.6-sol"); context.writeRawModelsJson({ openai: { modelOverrides: { "gpt-5.6-sol": { cost: { tiers: [replacementTier] } } } }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("openai", "gpt-5.6-sol"); expect(model?.cost.input).toBe(baseline?.cost.input); expect(model?.cost.output).toBe(baseline?.cost.output); expect(model?.cost.tiers).toEqual([replacementTier]); }); - test("empty tier override clears inherited tiers", () => { + test("empty tier override clears inherited tiers", async () => { context.writeRawModelsJson({ openai: { modelOverrides: { "gpt-5.6-sol": { cost: { tiers: [] } } } }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.find("openai", "gpt-5.6-sol")?.cost.tiers).toEqual([]); }); - test("dynamic provider model override replaces tiers while preserving unspecified scalar rates", () => { + test("dynamic provider model override replaces tiers while preserving unspecified scalar rates", async () => { context.writeRawModelsJson({ "extension-provider": { modelOverrides: { @@ -149,7 +149,7 @@ describeModelRegistry((context) => { }, }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider("extension-provider", { baseUrl: "https://provider.test/v1", apiKey: "TEST_KEY", diff --git a/packages/coding-agent/test/model-registry-custom-models.suite.ts b/packages/coding-agent/test/model-registry-custom-models.suite.ts index 62d9ea385..7b043213c 100644 --- a/packages/coding-agent/test/model-registry-custom-models.suite.ts +++ b/packages/coding-agent/test/model-registry-custom-models.suite.ts @@ -1,6 +1,6 @@ import type { AnthropicMessagesCompat, OpenAICompletionsCompat, OpenAIResponsesCompat } from "@earendil-works/pi-ai/compat"; import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; describeModelRegistry((context) => { @@ -15,7 +15,7 @@ describeModelRegistry((context) => { emptyContext, } = context; describe("custom models merge behavior", () => { - test("built-in provider custom models inherit api and baseUrl without explicit fields", () => { + test("built-in provider custom models inherit api and baseUrl without explicit fields", async () => { // Built-in providers already have api/baseUrl on every model, and auth // comes from env vars / auth storage. No need to specify them. writeRawModelsJson({ @@ -31,7 +31,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getError()).toBeUndefined(); const model = registry.find("openrouter", "fake-provider/fake-model"); @@ -40,7 +40,7 @@ describeModelRegistry((context) => { expect(model?.baseUrl).toBe("https://openrouter.ai/api/v1"); }); - test("non-built-in provider custom models still require baseUrl and apiKey", () => { + test("non-built-in provider custom models still require baseUrl and apiKey", async () => { writeRawModelsJson({ "my-custom-provider": { models: [ @@ -54,16 +54,16 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getError()).toContain("baseUrl"); }); - test("custom provider with same name as built-in merges with built-in models", () => { + test("custom provider with same name as built-in merges with built-in models", async () => { writeModelsJson({ anthropic: providerConfig("https://my-proxy.example.com/v1", [{ id: "claude-custom" }]), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const anthropicModels = getModelsForProvider(registry, "anthropic"); expect(anthropicModels.length).toBeGreaterThan(1); @@ -71,7 +71,7 @@ describeModelRegistry((context) => { expect(anthropicModels.some((m) => m.id.includes("claude"))).toBe(true); }); - test("custom model with same id replaces built-in model by id", () => { + test("custom model with same id replaces built-in model by id", async () => { writeModelsJson({ openrouter: providerConfig( "https://my-proxy.example.com/v1", @@ -80,7 +80,7 @@ describeModelRegistry((context) => { ), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnetModels = models.filter((m) => m.id === "anthropic/claude-sonnet-4"); @@ -88,23 +88,23 @@ describeModelRegistry((context) => { expect(sonnetModels[0].baseUrl).toBe("https://my-proxy.example.com/v1"); }); - test("custom provider with same name as built-in does not affect other built-in providers", () => { + test("custom provider with same name as built-in does not affect other built-in providers", async () => { writeModelsJson({ anthropic: providerConfig("https://my-proxy.example.com/v1", [{ id: "claude-custom" }]), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(getModelsForProvider(registry, "google").length).toBeGreaterThan(0); expect(getModelsForProvider(registry, "openai").length).toBeGreaterThan(0); }); - test("provider-level baseUrl applies to both built-in and custom models", () => { + test("provider-level baseUrl applies to both built-in and custom models", async () => { writeModelsJson({ anthropic: providerConfig("https://merged-proxy.example.com/v1", [{ id: "claude-custom" }]), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const anthropicModels = getModelsForProvider(registry, "anthropic"); for (const model of anthropicModels) { @@ -112,7 +112,7 @@ describeModelRegistry((context) => { } }); - test("provider-level compat applies to custom models", () => { + test("provider-level compat applies to custom models", async () => { writeRawModelsJson({ demo: { baseUrl: "https://example.com/v1", @@ -135,14 +135,14 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const compat = registry.find("demo", "demo-model")?.compat as OpenAICompletionsCompat | undefined; expect(compat?.supportsUsageInStreaming).toBe(false); expect(compat?.maxTokensField).toBe("max_tokens"); }); - test("model-level compat overrides provider-level compat for custom models", () => { + test("model-level compat overrides provider-level compat for custom models", async () => { writeRawModelsJson({ demo: { baseUrl: "https://example.com/v1", @@ -169,14 +169,14 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const compat = registry.find("demo", "demo-model")?.compat as OpenAICompletionsCompat | undefined; expect(compat?.supportsUsageInStreaming).toBe(true); expect(compat?.maxTokensField).toBe("max_completion_tokens"); }); - test("provider-level compat applies to built-in models", () => { + test("provider-level compat applies to built-in models", async () => { writeRawModelsJson({ openrouter: { compat: { @@ -186,7 +186,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); expect(models.length).toBeGreaterThan(0); @@ -197,7 +197,7 @@ describeModelRegistry((context) => { } }); - test("model schema accepts thinkingLevelMap and compat schema accepts supportsStrictMode and cacheControlFormat", () => { + test("model schema accepts thinkingLevelMap and compat schema accepts supportsStrictMode and cacheControlFormat", async () => { writeRawModelsJson({ demo: { baseUrl: "https://example.com/v1", @@ -224,7 +224,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const model = registry.find("demo", "demo-model"); const compat = model?.compat as OpenAICompletionsCompat | undefined; @@ -234,7 +234,7 @@ describeModelRegistry((context) => { expect(compat?.cacheControlFormat).toBe("anthropic"); }); - test("compat schema accepts chat template thinking configuration", () => { + test("compat schema accepts chat template thinking configuration", async () => { writeRawModelsJson({ demo: { baseUrl: "https://example.com/v1", @@ -260,7 +260,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const compat = registry.find("demo", "demo-model")?.compat as OpenAICompletionsCompat | undefined; expect(registry.getError()).toBeUndefined(); @@ -271,7 +271,7 @@ describeModelRegistry((context) => { }); }); - test("compat schema accepts Anthropic eager tool input streaming flag", () => { + test("compat schema accepts Anthropic eager tool input streaming flag", async () => { writeRawModelsJson({ demo: { baseUrl: "https://example.com", @@ -293,14 +293,14 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const compat = registry.find("demo", "demo-model")?.compat as AnthropicMessagesCompat | undefined; expect(registry.getError()).toBeUndefined(); expect(compat?.supportsEagerToolInputStreaming).toBe(false); }); - test("compat schema accepts long cache retention flag", () => { + test("compat schema accepts long cache retention flag", async () => { writeRawModelsJson({ demo: { baseUrl: "https://example.com", @@ -322,14 +322,14 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const compat = registry.find("demo", "demo-model")?.compat as AnthropicMessagesCompat | undefined; expect(registry.getError()).toBeUndefined(); expect(compat?.supportsLongCacheRetention).toBe(false); }); - test("compat schema accepts Pi 0.80.7 Responses session affinity settings", () => { + test("compat schema accepts Pi 0.80.7 Responses session affinity settings", async () => { writeRawModelsJson({ demo: { baseUrl: "https://example.com/v1", @@ -352,7 +352,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const compat = registry.find("demo", "demo-model")?.compat as OpenAIResponsesCompat | undefined; expect(registry.getError()).toBeUndefined(); @@ -360,7 +360,7 @@ describeModelRegistry((context) => { expect(compat?.supportsToolSearch).toBe(true); }); - test("model-level baseUrl overrides provider-level baseUrl for custom models", () => { + test("model-level baseUrl overrides provider-level baseUrl for custom models", async () => { writeRawModelsJson({ "opencode-go": { baseUrl: "https://opencode.ai/zen/go/v1", @@ -389,7 +389,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const m25 = registry.find("opencode-go", "minimax-m2.5"); const glm5 = registry.find("opencode-go", "glm-5"); @@ -397,7 +397,7 @@ describeModelRegistry((context) => { expect(glm5?.baseUrl).toBe("https://opencode.ai/zen/go/v1"); }); - test("modelOverrides still apply when provider also defines models", () => { + test("modelOverrides still apply when provider also defines models", async () => { writeRawModelsJson({ openrouter: { baseUrl: "https://my-proxy.example.com/v1", @@ -422,7 +422,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); expect(models.some((m) => m.id === "custom/openrouter-model")).toBe(true); @@ -435,14 +435,14 @@ describeModelRegistry((context) => { writeModelsJson({ anthropic: providerConfig("https://first-proxy.example.com/v1", [{ id: "claude-custom" }]), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(getModelsForProvider(registry, "anthropic").some((m) => m.id === "claude-custom")).toBe(true); // Update and refresh writeModelsJson({ anthropic: providerConfig("https://second-proxy.example.com/v1", [{ id: "claude-custom-2" }]), }); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); const anthropicModels = getModelsForProvider(registry, "anthropic"); expect(anthropicModels.some((m) => m.id === "claude-custom")).toBe(false); @@ -454,12 +454,12 @@ describeModelRegistry((context) => { writeModelsJson({ anthropic: providerConfig("https://proxy.example.com/v1", [{ id: "claude-custom" }]), }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(getModelsForProvider(registry, "anthropic").some((m) => m.id === "claude-custom")).toBe(true); // Remove custom models and refresh writeModelsJson({}); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); const anthropicModels = getModelsForProvider(registry, "anthropic"); expect(anthropicModels.some((m) => m.id === "claude-custom")).toBe(false); diff --git a/packages/coding-agent/test/model-registry-dynamic-providers.suite.ts b/packages/coding-agent/test/model-registry-dynamic-providers.suite.ts index 31924cec4..06696d7ce 100644 --- a/packages/coding-agent/test/model-registry-dynamic-providers.suite.ts +++ b/packages/coding-agent/test/model-registry-dynamic-providers.suite.ts @@ -1,15 +1,13 @@ -import { getApiProvider } from "@earendil-works/pi-ai/compat"; -import { getOAuthProvider } from "../src/core/oauth-provider-bridge.ts"; -import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { describe, expect, test, vi } from "vitest"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; describeModelRegistry((context) => { - const { providerConfig, getModelsForProvider, writeRawModelsJson, openAiModel, emptyContext } = context; + const { providerConfig, getModelsForProvider, writeRawModelsJson } = context; describe("dynamic provider lifecycle", () => { - test("getProviderDisplayName resolves registered, OAuth, built-in, and fallback names", () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + test("getProviderDisplayName resolves registered, OAuth, built-in, and fallback names", async () => { + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(registry.getProviderDisplayName("openai")).toBe("OpenAI"); expect(registry.getProviderDisplayName("github-copilot")).toBe("GitHub Copilot"); @@ -70,7 +68,7 @@ describeModelRegistry((context) => { }, }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider("extension-provider", { baseUrl: "https://provider.test/v1", apiKey: "TEST_KEY", @@ -97,12 +95,12 @@ describeModelRegistry((context) => { if (!model) throw new Error("missing extension model"); expect(await registry.getApiKeyAndHeaders(model)).toMatchObject({ ok: true, - headers: { "X-Base": "base", "X-Override": "override", "X-Shared": "override" }, + headers: { "X-Base": "base", "X-Override": "override", "X-Shared": "base" }, }); }); test("failed registerProvider does not persist invalid streamSimple config", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect(() => registry.registerProvider("broken-provider", { @@ -112,11 +110,11 @@ describeModelRegistry((context) => { }), ).toThrow('Provider broken-provider: "api" is required when registering streamSimple.'); - await expect(registry.refresh()).resolves.toBeDefined(); + await expect(registry.refresh()).resolves.toBeUndefined(); }); test("failed registerProvider does not remove existing provider models", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider("demo-provider", { baseUrl: "https://provider.test/v1", @@ -156,12 +154,12 @@ describeModelRegistry((context) => { ).toThrow('Provider demo-provider, model broken-model: no "api" specified.'); expect(registry.find("demo-provider", "demo-model")).toBeDefined(); - await expect(registry.refresh()).resolves.toBeDefined(); + await expect(registry.refresh()).resolves.toBeUndefined(); expect(registry.find("demo-provider", "demo-model")).toBeDefined(); }); - test("unregisterProvider removes custom OAuth provider and restores built-in OAuth provider", () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + test("unregisterProvider removes custom OAuth provider and restores built-in OAuth provider", async () => { + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider("anthropic", { oauth: { @@ -176,49 +174,20 @@ describeModelRegistry((context) => { }, }); - expect(getOAuthProvider("anthropic")?.name).toBe("Custom Anthropic OAuth"); + expect(registry.getProvider("anthropic")?.auth.oauth?.name).toBe("Custom Anthropic OAuth"); registry.unregisterProvider("anthropic"); - expect(getOAuthProvider("anthropic")?.name).not.toBe("Custom Anthropic OAuth"); + expect(registry.getProvider("anthropic")?.auth.oauth?.name).not.toBe("Custom Anthropic OAuth"); }); - test("unregisterProvider removes custom streamSimple override and restores built-in API stream handler", () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); - - registry.registerProvider("stream-override-provider", { - api: "openai-completions", - streamSimple: () => { - throw new Error("custom streamSimple override"); - }, - }); - - let threwCustomOverride = false; - try { - getApiProvider("openai-completions")?.streamSimple(openAiModel, emptyContext); - } catch (error) { - threwCustomOverride = error instanceof Error && error.message === "custom streamSimple override"; - } - expect(threwCustomOverride).toBe(true); - - registry.unregisterProvider("stream-override-provider"); - - let threwCustomOverrideAfterUnregister = false; - try { - getApiProvider("openai-completions")?.streamSimple(openAiModel, emptyContext); - } catch (error) { - threwCustomOverrideAfterUnregister = - error instanceof Error && error.message === "custom streamSimple override"; - } - expect(threwCustomOverrideAfterUnregister).toBe(false); - }); describe("dynamic provider override persistence", () => { test("baseUrl-only override keeps built-in provider models after refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider("anthropic", { baseUrl: "https://proxy.test/anthropic" }); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); const anthropicModels = getModelsForProvider(registry, "anthropic"); expect(anthropicModels.length).toBeGreaterThan(1); @@ -226,40 +195,40 @@ describeModelRegistry((context) => { }); test("models-only override replaces built-in provider models after refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider("anthropic", { ...providerConfig("https://custom.test/anthropic", [{ id: "custom-claude" }], "anthropic-messages"), baseUrl: "https://custom.test/anthropic", }); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); expect(getModelsForProvider(registry, "anthropic").map((m) => m.id)).toEqual(["custom-claude"]); expect(registry.find("anthropic", "custom-claude")?.baseUrl).toBe("https://custom.test/anthropic"); }); test("models plus baseUrl override replaces built-in provider models after refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider("anthropic", { ...providerConfig("https://custom.test/anthropic", [{ id: "custom-claude" }], "anthropic-messages"), baseUrl: "https://custom.test/anthropic", }); registry.registerProvider("anthropic", { baseUrl: "https://proxy.test/anthropic" }); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); expect(getModelsForProvider(registry, "anthropic").map((m) => m.id)).toEqual(["custom-claude"]); expect(registry.find("anthropic", "custom-claude")?.baseUrl).toBe("https://proxy.test/anthropic"); }); test("models-only custom provider registration survives refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider( "custom-provider", providerConfig("https://custom.test/v1", [{ id: "custom-a" }, { id: "custom-b" }], "openai-completions"), ); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); expect(getModelsForProvider(registry, "custom-provider").map((m) => m.id)).toEqual([ "custom-a", @@ -268,14 +237,14 @@ describeModelRegistry((context) => { }); test("baseUrl-only override keeps custom provider models after refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider( "custom-provider", providerConfig("https://custom.test/v1", [{ id: "custom-a" }, { id: "custom-b" }], "openai-completions"), ); registry.registerProvider("custom-provider", { baseUrl: "https://proxy.test/custom" }); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); expect(getModelsForProvider(registry, "custom-provider").map((m) => m.id)).toEqual([ "custom-a", @@ -289,14 +258,14 @@ describeModelRegistry((context) => { }); test("headers-only override keeps custom provider models after refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider( "custom-provider", providerConfig("https://custom.test/v1", [{ id: "custom-a" }, { id: "custom-b" }], "openai-completions"), ); registry.registerProvider("custom-provider", { headers: { "x-proxy": "enabled" } }); - await registry.refresh(); + await getModelRuntime(registry).refresh({ allowNetwork: false }); const models = getModelsForProvider(registry, "custom-provider"); expect(models.map((m) => m.id)).toEqual(["custom-a", "custom-b"]); @@ -308,13 +277,13 @@ describeModelRegistry((context) => { }); test("async catalog refresh publishes successful results only after completion", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const initial = providerConfig("https://dynamic.test/v1", [{ id: "old" }]); let resolveRefresh!: (models: NonNullable) => void; const pending = new Promise>((resolve) => (resolveRefresh = resolve)); registry.registerProvider("dynamic", { ...initial, refreshModels: () => pending }); - const refresh = registry.refresh(); + const refresh = getModelRuntime(registry).refresh(); expect(getModelsForProvider(registry, "dynamic").map((model) => model.id)).toEqual(["old"]); resolveRefresh(providerConfig("https://dynamic.test/v1", [{ id: "new" }]).models!); const result = await refresh; @@ -325,12 +294,12 @@ describeModelRegistry((context) => { }); test("unrelated registration does not discard an in-flight provider refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const initial = providerConfig("https://dynamic.test/v1", [{ id: "old" }]); let resolveRefresh!: (models: NonNullable) => void; const pending = new Promise>((resolve) => (resolveRefresh = resolve)); registry.registerProvider("dynamic", { ...initial, refreshModels: () => pending }); - const refresh = registry.refresh(); + const refresh = getModelRuntime(registry).refresh(); registry.registerProvider("anthropic", { headers: { "x-unrelated": "yes" } }); resolveRefresh(providerConfig("https://dynamic.test/v1", [{ id: "new" }]).models!); @@ -342,7 +311,7 @@ describeModelRegistry((context) => { test("async catalog refresh returns partial provider errors without discarding successes", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const good = providerConfig("https://good.test/v1", [{ id: "old-good" }]); const bad = providerConfig("https://bad.test/v1", [{ id: "old-bad" }]); registry.registerProvider("good", { @@ -351,7 +320,7 @@ describeModelRegistry((context) => { }); registry.registerProvider("bad", { ...bad, refreshModels: async ({ allowNetwork }) => { throw new Error(allowNetwork ? "catalog failed" : "cache fallback failed"); } }); - const result = await registry.refresh(); + const result = await getModelRuntime(registry).refresh(); expect(result.errors.get("bad")?.message).toBe("catalog failed"); expect(getModelsForProvider(registry, "good").map((model) => model.id)).toEqual(["new-good"]); @@ -359,12 +328,12 @@ describeModelRegistry((context) => { }); test("stale refresh completion cannot resurrect an unregistered provider", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const initial = providerConfig("https://dynamic.test/v1", [{ id: "old" }]); let resolveRefresh!: (models: NonNullable) => void; const pending = new Promise>((resolve) => (resolveRefresh = resolve)); registry.registerProvider("dynamic", { ...initial, refreshModels: () => pending }); - const staleRefresh = registry.refresh(); + const staleRefresh = getModelRuntime(registry).refresh(); registry.unregisterProvider("dynamic"); resolveRefresh(providerConfig("https://dynamic.test/v1", [{ id: "resurrected" }]).models!); @@ -373,27 +342,34 @@ describeModelRegistry((context) => { expect(registry.find("dynamic", "resurrected")).toBeUndefined(); }); - test("stale refresh completion cannot overwrite a re-registered provider or its cache", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + test("stale refresh completion cannot overwrite a re-registered provider but can update its shared cache", async () => { + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const initial = providerConfig("https://dynamic.test/v1", [{ id: "old" }]); let releaseStale!: () => void; - let markStaleStarted!: () => void; + let markAllStaleStarted!: () => void; + let markAllStaleFinished!: () => void; const staleGate = new Promise((resolve) => { releaseStale = resolve; }); - const staleStarted = new Promise((resolve) => { markStaleStarted = resolve; }); + const allStaleStarted = new Promise((resolve) => { markAllStaleStarted = resolve; }); + const allStaleFinished = new Promise((resolve) => { markAllStaleFinished = resolve; }); + let staleStartCount = 0; + let staleFinishCount = 0; let persistedAfterStale: string[] | undefined; registry.registerProvider("dynamic", { ...initial, refreshModels: async ({ store }) => { - markStaleStarted(); + staleStartCount += 1; + if (staleStartCount === 2) markAllStaleStarted(); await staleGate; const staleModels = providerConfig("https://dynamic.test/v1", [{ id: "stale-store" }]).models!; await store.write({ models: staleModels, checkedAt: Date.now() }); persistedAfterStale = (await store.read())?.models.map((model) => model.id); + staleFinishCount += 1; + if (staleFinishCount === 2) markAllStaleFinished(); return staleModels; }, }); - const staleRefresh = registry.refresh(); - await staleStarted; + const staleRefresh = getModelRuntime(registry).refresh(); + await allStaleStarted; const freshModels = providerConfig("https://manual.test/v1", [{ id: "manual" }]).models!; registry.registerProvider("dynamic", { @@ -403,92 +379,55 @@ describeModelRegistry((context) => { return freshModels; }, }); - await registry.refresh(); + await getModelRuntime(registry).refresh(); releaseStale(); - await staleRefresh; + await Promise.all([staleRefresh, allStaleFinished]); expect(getModelsForProvider(registry, "dynamic").map((model) => model.id)).toEqual(["manual"]); - expect(persistedAfterStale).toEqual(["manual"]); - }); - - test("async catalog refresh times out and retains the stale snapshot", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); - const initial = providerConfig("https://slow.test/v1", [{ id: "cached" }]); - registry.registerProvider("slow", { - ...initial, - refreshModels: async () => new Promise>(() => {}), - }); - - const result = await registry.refresh({ timeoutMs: 5 }); - - expect(result.aborted).toBe(true); - expect(getModelsForProvider(registry, "slow").map((model) => model.id)).toEqual(["cached"]); - }); - - test("aborted extension refresh cannot mutate its persisted cache", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); - const initial = providerConfig("https://slow.test/v1", [{ id: "cached" }]); - let releaseCallback!: () => void; - let callbackFinished!: () => void; - const callbackGate = new Promise((resolve) => { releaseCallback = resolve; }); - const finished = new Promise((resolve) => { callbackFinished = resolve; }); - let persistedAfterAbort: string[] | undefined; - registry.registerProvider("slow", { - ...initial, - refreshModels: async ({ store }) => { - await callbackGate; - await store.write({ models: initial.models!, checkedAt: Date.now() }); - persistedAfterAbort = (await store.read())?.models.map((model) => model.id); - callbackFinished(); - return initial.models!; - }, - }); - - const result = await registry.refresh({ timeoutMs: 5 }); - releaseCallback(); - await finished; - - expect(result.aborted).toBe(true); - expect(persistedAfterAbort).toBeUndefined(); + expect(persistedAfterStale).toEqual(["stale-store"]); }); test("pre-aborted refresh returns without invoking providers", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const initial = providerConfig("https://slow.test/v1", [{ id: "cached" }]); - let called = false; + let calls = 0; registry.registerProvider("slow", { ...initial, refreshModels: async () => { - called = true; + calls += 1; return initial.models!; }, }); + await vi.waitFor(() => expect(calls).toBeGreaterThan(0)); + calls = 0; const controller = new AbortController(); controller.abort(); - const result = await registry.refresh({ signal: controller.signal, timeoutMs: 100 }); + const result = await getModelRuntime(registry).refresh({ signal: controller.signal }); expect(result.aborted).toBe(true); - expect(called).toBe(false); + expect(calls).toBe(0); expect(registry.find("slow", "cached")).toBeDefined(); }); - test("additive provider overrides retain built-in credential filtering", () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + test("additive provider overrides retain built-in credential filtering", async () => { + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const allowed = registry.getAll().find((model) => model.provider === "github-copilot")!; - context.authStorage.set("github-copilot", { + await context.authStorage.modify("github-copilot", async () => ({ type: "oauth", refresh: "r", access: "a", expires: Date.now() + 60_000, availableModelIds: [allowed.id], - }); + })); + await getModelRuntime(registry).refresh({ allowNetwork: false }); const availableIds = () => registry.getAvailable() .filter((model) => model.provider === "github-copilot") .map((model) => model.id); expect(availableIds()).toEqual([allowed.id]); registry.registerProvider("github-copilot", { headers: { "x-test": "1" } }); + await getModelRuntime(registry).refresh({ allowNetwork: false }); expect(availableIds()).toEqual([allowed.id]); }); diff --git a/packages/coding-agent/test/model-registry-hot-reload.test.ts b/packages/coding-agent/test/model-registry-hot-reload.test.ts index 423c92fc5..8fa2cb883 100644 --- a/packages/coding-agent/test/model-registry-hot-reload.test.ts +++ b/packages/coding-agent/test/model-registry-hot-reload.test.ts @@ -1,9 +1,11 @@ import { mkdtempSync, rmSync, writeFileSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; +import type { TUI } from "@earendil-works/pi-tui"; import { afterEach, describe, expect, test, vi } from "vitest"; 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 type { SettingsManager } from "../src/core/settings-manager.ts"; import { ModelSelectorComponent } from "../src/modes/interactive/components/model-selector.ts"; import { initTheme } from "../src/modes/interactive/theme/theme.ts"; @@ -15,7 +17,7 @@ afterEach(() => { function writeModels(path: string, ids: string[]): void { writeFileSync(path, JSON.stringify({ providers: { - layered: { + configured: { api: "openai-completions", baseUrl: "https://example.test/v1", apiKey: "local", @@ -25,38 +27,39 @@ function writeModels(path: string, ids: string[]): void { })); } -describe("model registry layered hot reload", () => { - test("reloads legacy and primary models.json layers on every refresh", async () => { +describe("model config hot reload", () => { + test("reloads the configured models.json file on every refresh", async () => { const directory = mkdtempSync(join(tmpdir(), "atomic-model-reload-")); tempDirs.push(directory); - const legacy = join(directory, "legacy-models.json"); - const primary = join(directory, "primary-models.json"); - writeModels(legacy, ["legacy"]); - writeModels(primary, ["primary-before"]); - const registry = ModelRegistry.create(AuthStorage.inMemory(), [legacy, primary]); - expect(registry.find("layered", "legacy")).toBeDefined(); - expect(registry.find("layered", "primary-before")).toBeDefined(); + const modelsPath = join(directory, "models.json"); + writeModels(modelsPath, ["before"]); + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory(), + modelsPath, + allowModelNetwork: false, + }); + expect(runtime.getModel("configured", "before")).toBeDefined(); - writeModels(primary, ["primary-after"]); - await registry.refresh({ allowNetwork: false }); + writeModels(modelsPath, ["after"]); + await runtime.refresh({ allowNetwork: false }); - expect(registry.find("layered", "legacy")).toBeDefined(); - expect(registry.find("layered", "primary-before")).toBeUndefined(); - expect(registry.find("layered", "primary-after")).toBeDefined(); + expect(runtime.getModel("configured", "before")).toBeUndefined(); + expect(runtime.getModel("configured", "after")).toBeDefined(); }); - test("reloads model layers from every newly opened /model picker", async () => { + test("refreshes the model config from every newly opened /model picker", async () => { initTheme("dark"); const refresh = vi.fn(async () => ({ aborted: false, errors: new Map() })); - const registry = { + const runtime = { getError: () => undefined, - getAvailable: () => [], - find: () => undefined, + getAvailableSnapshot: () => [], + getModel: () => undefined, refresh, - }; - const tui = { requestRender: vi.fn() }; + } as unknown as ModelRuntime; + const tui = { requestRender: vi.fn() } as unknown as TUI; + const settings = { setDefaultModelAndProvider: vi.fn() } as unknown as SettingsManager; const openPicker = () => new ModelSelectorComponent( - tui as never, undefined, {} as never, registry as never, [], () => {}, () => {}, + tui, undefined, settings, runtime, [], () => {}, () => {}, ); openPicker(); diff --git a/packages/coding-agent/test/model-registry-layered-overrides.suite.ts b/packages/coding-agent/test/model-registry-layered-overrides.suite.ts deleted file mode 100644 index 3f697bb82..000000000 --- a/packages/coding-agent/test/model-registry-layered-overrides.suite.ts +++ /dev/null @@ -1,162 +0,0 @@ -import { writeFileSync } from "node:fs"; -import { join } from "node:path"; -import { expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; -import { describeModelRegistry } from "./model-registry-fixtures.ts"; - -const sonnetId = "anthropic/claude-sonnet-4"; -const opusId = "anthropic/claude-opus-4"; - -function modelWithHeaders(id: string, headers?: Record): Record { - return { id, headers }; -} - -describeModelRegistry((context) => { - function createLayeredRegistry(...providerLayers: Array>) { - const paths = providerLayers.map((providers, index) => { - const path = join(context.tempDir, `layer-${index}-models.json`); - writeFileSync(path, JSON.stringify({ providers })); - return path; - }); - return ModelRegistry.create(context.authStorage, paths); - } - - async function resolveHeaders(registry: ModelRegistry): Promise> { - const model = registry.find("openrouter", sonnetId); - if (!model) throw new Error("missing layered model"); - const auth = await registry.getApiKeyAndHeaders(model); - if (!auth.ok) throw new Error(auth.error); - return auth.headers ?? {}; - } - - test("layered modelOverrides retain disjoint model IDs under the same provider", () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [opusId]: { name: "Primary Opus" } } } }, - { openrouter: { modelOverrides: { [sonnetId]: { name: "Legacy Sonnet" } } } }, - ); - - expect(registry.find("openrouter", sonnetId)?.name).toBe("Legacy Sonnet"); - expect(registry.find("openrouter", opusId)?.name).toBe("Primary Opus"); - }); - - test("primary exact modelOverride replaces the legacy entry wholesale", () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [sonnetId]: { maxTokens: 12_345 } } } }, - { openrouter: { modelOverrides: { [sonnetId]: { name: "Legacy Name", cost: { input: 99 } } } } }, - ); - const model = registry.find("openrouter", sonnetId); - - expect(model?.maxTokens).toBe(12_345); - expect(model?.name).not.toBe("Legacy Name"); - expect(model?.cost.input).not.toBe(99); - }); - - test("primary empty modelOverride suppresses legacy-only fields", () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [sonnetId]: {} } } }, - { openrouter: { modelOverrides: { [sonnetId]: { name: "Legacy Name" } } } }, - ); - - expect(registry.find("openrouter", sonnetId)?.name).not.toBe("Legacy Name"); - }); - - test("primary empty provider override map retains legacy modelOverrides", () => { - const registry = createLayeredRegistry( - { openrouter: { baseUrl: "https://primary.example/v1", modelOverrides: {} } }, - { openrouter: { modelOverrides: { [sonnetId]: { name: "Legacy Sonnet" } } } }, - ); - - expect(registry.find("openrouter", sonnetId)?.name).toBe("Legacy Sonnet"); - }); - - test("more than two modelOverride layers retain ordered overlays", () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [sonnetId]: { name: "Primary Sonnet" } } } }, - { openrouter: { modelOverrides: { [opusId]: { name: "Middle Opus" } } } }, - { openrouter: { modelOverrides: { [sonnetId]: { name: "Legacy Sonnet" } } } }, - ); - - expect(registry.find("openrouter", sonnetId)?.name).toBe("Primary Sonnet"); - expect(registry.find("openrouter", opusId)?.name).toBe("Middle Opus"); - }); - - test("primary exact replacement without headers clears a legacy override header", async () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [sonnetId]: { name: "Primary Name" } } } }, - { openrouter: { modelOverrides: { [sonnetId]: { headers: { "X-Legacy": "legacy" } } } } }, - ); - - expect((await resolveHeaders(registry))["X-Legacy"]).toBeUndefined(); - }); - - test("a legacy custom-model header survives a primary override without headers", async () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [sonnetId]: { name: "Primary Name" } } } }, - { openrouter: { models: [modelWithHeaders(sonnetId, { "X-Model": "legacy-model" })] } }, - ); - - expect((await resolveHeaders(registry))["X-Model"]).toBe("legacy-model"); - }); - - test("a primary override header beats a retained legacy custom-model header", async () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [sonnetId]: { headers: { "X-Layered": "primary-override" } } } } }, - { openrouter: { models: [modelWithHeaders(sonnetId, { "X-Layered": "legacy-model" })] } }, - ); - - expect((await resolveHeaders(registry))["X-Layered"]).toBe("primary-override"); - }); - - test("a primary custom-model header beats a retained legacy override header", async () => { - const registry = createLayeredRegistry( - { openrouter: { models: [modelWithHeaders(sonnetId, { "X-Layered": "primary-model" })] } }, - { openrouter: { modelOverrides: { [sonnetId]: { headers: { "X-Layered": "legacy-override" } } } } }, - ); - - expect((await resolveHeaders(registry))["X-Layered"]).toBe("primary-model"); - }); - - test("a custom-model header wins over an override header in the same file", async () => { - const registry = createLayeredRegistry({ - openrouter: { - models: [modelWithHeaders(sonnetId, { "X-Layered": "same-file-model" })], - modelOverrides: { [sonnetId]: { headers: { "X-Layered": "same-file-override" } } }, - }, - }); - - expect((await resolveHeaders(registry))["X-Layered"]).toBe("same-file-model"); - }); - - test("primary override headers win an exact layered override conflict", async () => { - const registry = createLayeredRegistry( - { openrouter: { modelOverrides: { [sonnetId]: { headers: { "X-Layered": "primary" } } } } }, - { openrouter: { modelOverrides: { [sonnetId]: { headers: { "X-Layered": "legacy" } } } } }, - ); - - expect((await resolveHeaders(registry))["X-Layered"]).toBe("primary"); - }); - - test("extension providers receive disjoint layered modelOverrides", () => { - const registry = createLayeredRegistry( - { "layered-extension": { modelOverrides: { "model-b": { name: "Primary B" } } } }, - { "layered-extension": { modelOverrides: { "model-a": { name: "Legacy A" } } } }, - ); - registry.registerProvider("layered-extension", { - baseUrl: "https://provider.test/v1", - apiKey: "TEST_KEY", - api: "openai-completions", - models: ["model-a", "model-b"].map((id) => ({ - id, - name: id, - reasoning: false, - input: ["text"], - cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, - contextWindow: 128_000, - maxTokens: 4096, - })), - }); - - expect(registry.find("layered-extension", "model-a")?.name).toBe("Legacy A"); - expect(registry.find("layered-extension", "model-b")?.name).toBe("Primary B"); - }); -}); diff --git a/packages/coding-agent/test/model-registry-model-overrides.suite.ts b/packages/coding-agent/test/model-registry-model-overrides.suite.ts index 16b1a07af..a7bfd151d 100644 --- a/packages/coding-agent/test/model-registry-model-overrides.suite.ts +++ b/packages/coding-agent/test/model-registry-model-overrides.suite.ts @@ -1,6 +1,6 @@ import type { OpenAICompletionsCompat } from "@earendil-works/pi-ai/compat"; import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createModelRegistry } from "./model-runtime-test-utils.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; describeModelRegistry((context) => { @@ -15,7 +15,7 @@ describeModelRegistry((context) => { emptyContext, } = context; describe("modelOverrides (per-model customization)", () => { - test("model override applies to a single built-in model", () => { + test("model override applies to a single built-in model", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -26,7 +26,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnet = models.find((m) => m.id === "anthropic/claude-sonnet-4"); @@ -37,7 +37,7 @@ describeModelRegistry((context) => { expect(opus?.name).not.toBe("Custom Sonnet Name"); }); - test("model override with compat.openRouterRouting", () => { + test("model override with compat.openRouterRouting", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -50,7 +50,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnet = models.find((m) => m.id === "anthropic/claude-sonnet-4"); @@ -58,7 +58,7 @@ describeModelRegistry((context) => { expect(compat?.openRouterRouting).toEqual({ only: ["amazon-bedrock"] }); }); - test("model override deep merges compat settings", () => { + test("model override deep merges compat settings", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -71,7 +71,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnet = models.find((m) => m.id === "anthropic/claude-sonnet-4"); @@ -80,7 +80,7 @@ describeModelRegistry((context) => { expect(compat?.openRouterRouting).toEqual({ order: ["anthropic", "together"] }); }); - test("model override deep merges chatTemplateKwargs", () => { + test("model override deep merges chatTemplateKwargs", async () => { writeRawModelsJson({ openrouter: { compat: { @@ -100,7 +100,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const sonnet = getModelsForProvider(registry, "openrouter").find((m) => m.id === "anthropic/claude-sonnet-4"); const compat = sonnet?.compat as OpenAICompletionsCompat | undefined; @@ -111,7 +111,7 @@ describeModelRegistry((context) => { }); }); - test("multiple model overrides on same provider", () => { + test("multiple model overrides on same provider", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -125,7 +125,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnet = models.find((m) => m.id === "anthropic/claude-sonnet-4"); @@ -137,7 +137,7 @@ describeModelRegistry((context) => { expect(opusCompat?.openRouterRouting).toEqual({ only: ["anthropic"] }); }); - test("model override combined with baseUrl override", () => { + test("model override combined with baseUrl override", async () => { writeRawModelsJson({ openrouter: { baseUrl: "https://my-proxy.example.com/v1", @@ -149,7 +149,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnet = models.find((m) => m.id === "anthropic/claude-sonnet-4"); @@ -163,7 +163,7 @@ describeModelRegistry((context) => { expect(opus?.name).not.toBe("Proxied Sonnet"); }); - test("model override for non-existent model ID is ignored", () => { + test("model override for non-existent model ID is ignored", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -174,7 +174,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); // Should not create a new model @@ -183,7 +183,7 @@ describeModelRegistry((context) => { expect(registry.getError()).toBeUndefined(); }); - test("model override can change cost fields partially", () => { + test("model override can change cost fields partially", async () => { writeRawModelsJson({ openrouter: { modelOverrides: { @@ -194,7 +194,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnet = models.find((m) => m.id === "anthropic/claude-sonnet-4"); @@ -215,7 +215,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const models = getModelsForProvider(registry, "openrouter"); const sonnet = models.find((m) => m.id === "anthropic/claude-sonnet-4"); expect(sonnet).toBeDefined(); @@ -238,7 +238,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); expect( getModelsForProvider(registry, "openrouter").find((m) => m.id === "anthropic/claude-sonnet-4")?.name, ).toBe("First Name"); @@ -271,7 +271,7 @@ describeModelRegistry((context) => { }, }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const customName = getModelsForProvider(registry, "openrouter").find( (m) => m.id === "anthropic/claude-sonnet-4", )?.name; diff --git a/packages/coding-agent/test/model-registry-provider-membership.suite.ts b/packages/coding-agent/test/model-registry-provider-membership.suite.ts index 165665071..a487e024c 100644 --- a/packages/coding-agent/test/model-registry-provider-membership.suite.ts +++ b/packages/coding-agent/test/model-registry-provider-membership.suite.ts @@ -1,10 +1,10 @@ import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; describeModelRegistry((context) => { describe("provider membership", () => { - test("tracks exact provider membership independently from authentication and model materialization", () => { + test("tracks exact provider membership independently from authentication and model materialization", async () => { context.writeRawModelsJson({ "configured-provider": context.providerConfig( "https://configured.test/v1", @@ -12,13 +12,14 @@ describeModelRegistry((context) => { "openai-completions", ), }); - context.authStorage.setRuntimeApiKey("stale-auth-only", "stale-token"); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + await context.authStorage.modify("stale-auth-only", async () => ({ type: "api_key", key: "stale-token" })); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); + const runtime = getModelRuntime(registry); - expect(registry.hasProvider("openai")).toBe(true); - expect(registry.hasProvider("OpenAI")).toBe(false); - expect(registry.hasProvider("configured-provider")).toBe(true); - expect(registry.hasProvider("stale-auth-only")).toBe(false); + expect(runtime.getProvider("openai") !== undefined).toBe(true); + expect(runtime.getProvider("OpenAI") !== undefined).toBe(false); + expect(runtime.getProvider("configured-provider") !== undefined).toBe(true); + expect(runtime.getProvider("stale-auth-only") !== undefined).toBe(false); registry.registerProvider("extension-without-models", { api: "openai-completions", @@ -26,17 +27,17 @@ describeModelRegistry((context) => { throw new Error("not called"); }, }); - expect(registry.hasProvider("extension-without-models")).toBe(true); + expect(runtime.getProvider("extension-without-models") !== undefined).toBe(true); registry.unregisterProvider("extension-without-models"); - expect(registry.hasProvider("extension-without-models")).toBe(false); + expect(runtime.getProvider("extension-without-models") !== undefined).toBe(false); const anthropic = registry.getProvider("anthropic"); if (!anthropic) throw new Error("missing built-in provider fixture"); registry.registerProvider({ ...anthropic, id: "native-member" }); - expect(registry.hasProvider("native-member")).toBe(true); + expect(runtime.getProvider("native-member") !== undefined).toBe(true); registry.unregisterProvider("native-member"); - expect(registry.hasProvider("native-member")).toBe(false); - expect(registry.hasProvider("anthropic")).toBe(true); + expect(runtime.getProvider("native-member") !== undefined).toBe(false); + expect(runtime.getProvider("anthropic") !== undefined).toBe(true); }); }); }); diff --git a/packages/coding-agent/test/model-registry-provider-runtime-ownership.suite.ts b/packages/coding-agent/test/model-registry-provider-runtime-ownership.suite.ts index eb7f64e84..ffc08a048 100644 --- a/packages/coding-agent/test/model-registry-provider-runtime-ownership.suite.ts +++ b/packages/coding-agent/test/model-registry-provider-runtime-ownership.suite.ts @@ -7,10 +7,9 @@ import { streamSimple, unregisterApiProviders, } from "@earendil-works/pi-ai/compat"; -import { getOAuthProvider } from "../src/core/oauth-provider-bridge.ts"; import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; describeModelRegistry((context) => { const { providerConfig, getModelsForProvider, openAiModel, emptyContext } = context; @@ -18,8 +17,8 @@ describeModelRegistry((context) => { describe("dynamic provider lifecycle", () => { describe("dynamic provider override persistence", () => { test("one registry cannot erase another registry's API or OAuth registrations", async () => { - const first = ModelRegistry.create(context.authStorage, context.modelsJsonPath); - const second = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const first = await createModelRegistry(context.authStorage, context.modelsJsonPath); + const second = await createModelRegistry(context.authStorage, context.modelsJsonPath); const api = "registry-isolation-api" as Api; const oauth = (name: string) => ({ name, @@ -38,21 +37,22 @@ describeModelRegistry((context) => { streamSimple: () => { throw new Error("second"); }, }); const secondApi = getApiProvider(api); - await first.refresh({ allowNetwork: false }); + await getModelRuntime(first).refresh({ allowNetwork: false }); expect(getApiProvider(api)).toBe(secondApi); - expect(getOAuthProvider("registry-isolation")?.name).toBe("second"); + expect(second.getProvider("registry-isolation")?.auth.oauth?.name).toBe("second"); first.unregisterProvider("registry-isolation"); expect(getApiProvider(api)).toBe(secondApi); - expect(getOAuthProvider("registry-isolation")?.name).toBe("second"); + expect(second.getProvider("registry-isolation")?.auth.oauth?.name).toBe("second"); second.unregisterProvider("registry-isolation"); expect(getApiProvider(api)).toBeUndefined(); }); - test("unregistering the latest registry restores the previous API owner", () => { - const first = ModelRegistry.create(context.authStorage, context.modelsJsonPath); - const second = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + + test("pi scopes extension streams to each runtime instead of stacking global API owners", async () => { + const first = await createModelRegistry(context.authStorage, context.modelsJsonPath); + const second = await createModelRegistry(context.authStorage, context.modelsJsonPath); const api = "registry-fallback-api" as Api; first.registerProvider("registry-first", { api, @@ -65,12 +65,12 @@ describeModelRegistry((context) => { second.unregisterProvider("registry-second"); - expect(() => getApiProvider(api)?.streamSimple({ ...openAiModel, api }, emptyContext)).toThrow("first-owner"); + expect(getApiProvider(api)).toBeUndefined(); + expect(first.getProvider("registry-first")).toBeDefined(); first.unregisterProvider("registry-first"); }); - - test("unregistering an Atomic override restores an external API owner", () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + test("unregistering an Atomic override restores an external API owner", async () => { + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const api = "external-fallback-api" as Api; registerApiProvider({ api, @@ -89,8 +89,8 @@ describeModelRegistry((context) => { unregisterApiProviders("external-owner"); }); - test("unregistering an unrelated API does not reclassify an active override", () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + test("unregistering an unrelated API does not reclassify an active override", async () => { + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const externalSource = "active-openai-override"; let customDispatches = 0; registerApiProvider({ @@ -115,9 +115,10 @@ describeModelRegistry((context) => { }); test("passes runtime-only credentials to extension catalog refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); + const runtime = getModelRuntime(registry); let observedCredential: Credential | undefined; - context.authStorage.setRuntimeApiKey("dynamic-probe", "runtime-secret"); + await runtime.setRuntimeApiKey("dynamic-probe", "runtime-secret"); registry.registerProvider("dynamic-probe", { refreshModels: async ({ credential }) => { observedCredential = credential; @@ -125,15 +126,16 @@ describeModelRegistry((context) => { }, }); - const result = await registry.refresh({ allowNetwork: false }); + const result = await runtime.refresh({ allowNetwork: false }); expect(result.errors.size).toBe(0); expect(observedCredential).toEqual({ type: "api_key", key: "runtime-secret" }); - expect(context.authStorage.get("dynamic-probe")).toBeUndefined(); + expect(await context.authStorage.read("dynamic-probe")).toBeUndefined(); }); test("passes configured API keys to extension catalog refresh", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); + const runtime = getModelRuntime(registry); let observedCredential: Credential | undefined; registry.registerProvider("configured-catalog", { apiKey: "literal-secret", @@ -143,14 +145,14 @@ describeModelRegistry((context) => { }, }); - const result = await registry.refresh({ allowNetwork: true }); + const result = await runtime.refresh({ allowNetwork: true }); expect(result.errors.size).toBe(0); expect(observedCredential).toEqual({ type: "api_key", key: "literal-secret" }); }); test("ignores undefined fields in partial provider updates", async () => { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); registry.registerProvider( "partial-provider", providerConfig("https://partial.test/v1", [{ id: "kept-model" }], "openai-completions"), @@ -174,7 +176,7 @@ describeModelRegistry((context) => { }; await expectPreservedProvider(); - await registry.refresh({ allowNetwork: false }); + await getModelRuntime(registry).refresh({ allowNetwork: false }); await expectPreservedProvider(); }); }); diff --git a/packages/coding-agent/test/model-registry-refresh-credential-resolution.test.ts b/packages/coding-agent/test/model-registry-refresh-credential-resolution.test.ts index 3b619483c..60895ec3a 100644 --- a/packages/coding-agent/test/model-registry-refresh-credential-resolution.test.ts +++ b/packages/coding-agent/test/model-registry-refresh-credential-resolution.test.ts @@ -1,7 +1,7 @@ import { writeFileSync } from "node:fs"; import { join } from "node:path"; import { describe, expect, test } from "vitest"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; import { describeModelRegistry } from "./model-registry-fixtures.ts"; describeModelRegistry((context) => { @@ -11,7 +11,7 @@ describeModelRegistry((context) => { const original = process.env[envVarName]; process.env[envVarName] = "environment-catalog-key"; try { - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); let observedKey: string | undefined; registry.registerProvider("environment-catalog", { apiKey: `$${envVarName}`, @@ -32,7 +32,7 @@ describeModelRegistry((context) => { test("resolves configured command-backed API keys", async () => { const tokenFile = join(context.tempDir, "catalog-token"); writeFileSync(tokenFile, "command-catalog-key"); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); let observedKey: string | undefined; registry.registerProvider("command-catalog", { apiKey: `!sh -c 'cat "${context.toShPath(tokenFile)}"'`, @@ -53,12 +53,15 @@ describeModelRegistry((context) => { const tokenFile = join(context.tempDir, "stored-catalog-token"); writeFileSync(tokenFile, "stored-command-key"); try { - context.authStorage.set("stored-environment", { type: "api_key", key: `$${envVarName}` }); - context.authStorage.set("stored-command", { + await context.authStorage.modify("stored-environment", async () => ({ + type: "api_key", + key: `$${envVarName}`, + })); + await context.authStorage.modify("stored-command", async () => ({ type: "api_key", key: `!sh -c 'cat "${context.toShPath(tokenFile)}"'`, - }); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + })); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); const observed = new Map(); for (const providerId of ["stored-environment", "stored-command"]) { registry.registerProvider(providerId, { @@ -81,9 +84,12 @@ describeModelRegistry((context) => { }); test("keeps runtime over stored over configured key precedence", async () => { - context.authStorage.set("credential-precedence", { type: "api_key", key: "stored-key" }); - context.authStorage.setRuntimeApiKey("credential-precedence", "runtime-key"); - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + await context.authStorage.modify("credential-precedence", async () => ({ + type: "api_key", + key: "stored-key", + })); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); + await getModelRuntime(registry).setRuntimeApiKey("credential-precedence", "runtime-key"); let observedKey: string | undefined; registry.registerProvider("credential-precedence", { apiKey: "configured-key", @@ -95,13 +101,16 @@ describeModelRegistry((context) => { await registry.refresh(); expect(observedKey).toBe("runtime-key"); - expect(context.authStorage.get("credential-precedence")).toEqual({ type: "api_key", key: "stored-key" }); + expect(await context.authStorage.read("credential-precedence")).toEqual({ type: "api_key", key: "stored-key" }); }); test("does not pass unresolved stored API-key expressions literally", async () => { - context.authStorage.set("missing-expression", { type: "api_key", key: "$ATOMIC_MISSING_CATALOG_KEY" }); + await context.authStorage.modify("missing-expression", async () => ({ + type: "api_key", + key: "$ATOMIC_MISSING_CATALOG_KEY", + })); delete process.env.ATOMIC_MISSING_CATALOG_KEY; - const registry = ModelRegistry.create(context.authStorage, context.modelsJsonPath); + const registry = await createModelRegistry(context.authStorage, context.modelsJsonPath); let observedKey: string | undefined; registry.registerProvider("missing-expression", { refreshModels: async ({ credential }) => { diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index 6ed4fd52c..79b2c2d72 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -3,7 +3,6 @@ import "./model-registry-base-url-overrides.suite.ts"; import "./model-registry-custom-models.suite.ts"; import "./model-registry-cost-tiers.suite.ts"; import "./model-registry-model-overrides.suite.ts"; -import "./model-registry-layered-overrides.suite.ts"; import "./model-registry-dynamic-providers.suite.ts"; import "./model-registry-provider-runtime-ownership.suite.ts"; import "./model-registry-provider-membership.suite.ts"; diff --git a/packages/coding-agent/test/model-resolver-initial.test.ts b/packages/coding-agent/test/model-resolver-initial.test.ts index a16ccd4d7..746a5efae 100644 --- a/packages/coding-agent/test/model-resolver-initial.test.ts +++ b/packages/coding-agent/test/model-resolver-initial.test.ts @@ -7,6 +7,7 @@ import { restoreModelFromSession, } from "../src/core/model-resolver.ts"; import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createInMemoryModelRegistry, createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; const allModels: Model<"anthropic-messages">[] = [ { @@ -90,14 +91,14 @@ describe("default model selection", () => { }); test("findInitialModel accepts explicit provider custom model ids", async () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = await findInitialModel({ cliProvider: "openrouter", cliModel: "openrouter/openai/ghost-model", scopedModels: [], isContinuing: false, - modelRegistry: registry, + modelRuntime: registry instanceof ModelRegistry ? getModelRuntime(registry) : registry, }); expect(result.model?.provider).toBe("openrouter"); expect(result.model?.id).toBe("openai/ghost-model"); @@ -105,17 +106,17 @@ describe("default model selection", () => { test("findInitialModel reports an unusable complete saved default without switching providers", async () => { const availableModel = allModels[1]!; const registry = { - find: () => undefined, - hasProvider: () => false, - getAvailable: async () => [availableModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModel: () => undefined, + getProvider: () => undefined, + getAvailableSnapshot: () => [availableModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = await findInitialModel({ scopedModels: [], isContinuing: false, defaultProvider: ["cur", "sor"].join(""), defaultModelId: ["composer", "-2"].join(""), defaultThinkingLevel: "medium", - modelRegistry: registry, + modelRuntime: registry instanceof ModelRegistry ? getModelRuntime(registry) : registry, }); expect(result.model).toBeUndefined(); expect(result.fallbackMessage).toBe( @@ -126,16 +127,16 @@ describe("default model selection", () => { test("findInitialModel keeps normal fallback for an unknown model on a supported provider", async () => { const availableModel = allModels[1]!; const registry = { - find: () => undefined, - hasProvider: (provider: string) => provider === "openai", - getAvailable: async () => [availableModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModel: () => undefined, + getProvider: (provider: string) => provider === "openai" ? ({ id: "openai" } as never) : undefined, + getAvailableSnapshot: () => [availableModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = await findInitialModel({ scopedModels: [], isContinuing: false, defaultProvider: "openai", defaultModelId: "unknown-saved-model", - modelRegistry: registry, + modelRuntime: registry instanceof ModelRegistry ? getModelRuntime(registry) : registry, }); expect(result.model).toBe(availableModel); expect(result.fallbackMessage).toBeUndefined(); @@ -144,8 +145,8 @@ describe("default model selection", () => { test("findInitialModel keeps automatic selection permissive when a saved-default field is omitted", async () => { const availableModel = allModels[1]!; const registry = { - getAvailable: async () => [availableModel], - } as unknown as Parameters[0]["modelRegistry"]; + getAvailableSnapshot: () => [availableModel], + } as unknown as Parameters[0]["modelRuntime"]; for (const partialDefault of [ { defaultProvider: "openai" }, @@ -155,14 +156,14 @@ describe("default model selection", () => { scopedModels: [], isContinuing: false, ...partialDefault, - modelRegistry: registry, + modelRuntime: registry instanceof ModelRegistry ? getModelRuntime(registry) : registry, }); expect(result.model).toBe(availableModel); expect(result.fallbackMessage).toBeUndefined(); } }); test("findInitialModel resolves a valid authenticated custom-provider default", async () => { - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const registry = await createInMemoryModelRegistry(AuthStorage.inMemory()); registry.registerProvider("custom-openai", { baseUrl: "https://custom.example/v1", apiKey: "test-key", @@ -183,7 +184,7 @@ describe("default model selection", () => { isContinuing: false, defaultProvider: "custom-openai", defaultModelId: "custom-default", - modelRegistry: registry, + modelRuntime: registry instanceof ModelRegistry ? getModelRuntime(registry) : registry, }); expect(result.model?.provider).toBe("custom-openai"); expect(result.model?.id).toBe("custom-default"); @@ -192,8 +193,9 @@ describe("default model selection", () => { test("restoreModelFromSession does not synthesize removed catalog-backed OpenAI ids", async () => { const openaiBaseModel = allModels[1]!; const registry = { - find: () => undefined, - getAvailable: async () => [openaiBaseModel], + getModel: () => undefined, + getProvider: () => ({ id: "openai" } as never), + getAvailableSnapshot: () => [openaiBaseModel], canRestoreUnknownModel: () => false, } as unknown as Parameters[4]; const result = await restoreModelFromSession("openai", "gpt-5.6", undefined, false, registry); @@ -203,7 +205,7 @@ describe("default model selection", () => { expect(result.fallbackMessage).toContain("model no longer exists"); }); test("restoreModelFromSession restores missing ids for registered OpenAI-compatible providers", async () => { - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const registry = await createInMemoryModelRegistry(AuthStorage.inMemory()); registry.registerProvider("custom-openai", { baseUrl: "https://custom.example/v1", apiKey: "test-key", @@ -226,7 +228,7 @@ describe("default model selection", () => { "newly-discovered-model", undefined, false, - registry, + getModelRuntime(registry), ); expect(result.fallbackMessage).toBeUndefined(); @@ -237,20 +239,23 @@ describe("default model selection", () => { const unauthenticatedExact = { ...allModels[1]!, id: "saved-exact" }; const fallbackModel = allModels[0]!; const registry = { - find: () => unauthenticatedExact, + getModel: () => unauthenticatedExact, + getProvider: () => ({ id: unauthenticatedExact.provider } as never), hasConfiguredAuth: () => false, - getAvailable: async () => [fallbackModel], + getAvailableSnapshot: () => [fallbackModel], } as unknown as Parameters[4]; const result = await restoreModelFromSession("openai", "saved-exact", undefined, false, registry); expect(result.model).toBe(fallbackModel); expect(result.model?.id).not.toBe("saved-exact"); expect(result.fallbackMessage).toContain("no auth configured"); }); - test("restoreModelFromSession scrubs inherited context-window options from fallback models", async () => { + test("restoreModelFromSession restores missing ids from registered provider models", async () => { const registry = { - find: () => undefined, - getAvailable: async () => [copilotSelectableBaseModel], + getModel: () => undefined, + getProvider: () => ({ id: "github-copilot" } as never), canRestoreUnknownModel: () => true, + getAvailableSnapshot: () => [copilotSelectableBaseModel], + hasConfiguredAuth: () => true, } as unknown as Parameters[4]; const result = await restoreModelFromSession( "github-copilot", @@ -278,12 +283,12 @@ describe("default model selection", () => { maxTokens: 8192, }; const registry = { - getAvailable: async () => [aiGatewayModel], - } as unknown as Parameters[0]["modelRegistry"]; + getAvailableSnapshot: () => [aiGatewayModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = await findInitialModel({ scopedModels: [], isContinuing: false, - modelRegistry: registry, + modelRuntime: registry instanceof ModelRegistry ? getModelRuntime(registry) : registry, }); expect(result.model?.provider).toBe("vercel-ai-gateway"); expect(result.model?.id).toBe("anthropic/claude-opus-4-6"); @@ -292,18 +297,18 @@ describe("default model selection", () => { const savedModel = allModels[0]!; const availableModel = allModels[1]!; const registry = { - find: () => savedModel, - hasProvider: (provider: string) => provider === savedModel.provider, + getModel: () => savedModel, + getProvider: (provider: string) => provider === savedModel.provider, hasConfiguredAuth: (model: Model<"anthropic-messages">) => model === availableModel, - getAvailable: async () => [availableModel], - } as unknown as Parameters[0]["modelRegistry"]; + getAvailableSnapshot: () => [availableModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = await findInitialModel({ scopedModels: [], isContinuing: false, defaultProvider: savedModel.provider, defaultModelId: savedModel.id, - modelRegistry: registry, + modelRuntime: registry instanceof ModelRegistry ? getModelRuntime(registry) : registry, }); expect(result.model).toBe(availableModel); diff --git a/packages/coding-agent/test/model-resolver-scope-upstream.test.ts b/packages/coding-agent/test/model-resolver-scope-upstream.test.ts index 3026d2662..8ce2c1cab 100644 --- a/packages/coding-agent/test/model-resolver-scope-upstream.test.ts +++ b/packages/coding-agent/test/model-resolver-scope-upstream.test.ts @@ -1,7 +1,7 @@ import type { Model } from "@earendil-works/pi-ai/compat"; import { describe, expect, test } from "vitest"; import { resolveModelScopeWithDiagnostics } from "../src/core/model-resolver.ts"; -import type { ModelRegistry } from "../src/core/model-registry.ts"; +import type { ModelRuntime } from "../src/core/model-runtime.ts"; function model(id: string, provider = "openrouter"): Model<"openai-completions"> { return { @@ -18,12 +18,12 @@ function model(id: string, provider = "openrouter"): Model<"openai-completions"> }; } -function registry(models: Model<"openai-completions">[]): ModelRegistry { - return { getAvailable: () => models } as Pick as ModelRegistry; +function runtime(models: Model<"openai-completions">[]): ModelRuntime { + return { getAvailableSnapshot: () => models } as Pick as ModelRuntime; } async function resolve(patterns: string[], models: Model<"openai-completions">[]) { - return resolveModelScopeWithDiagnostics(patterns, registry(models)); + return resolveModelScopeWithDiagnostics(patterns, runtime(models)); } describe("upstream model scope resolution", () => { diff --git a/packages/coding-agent/test/model-resolver.test.ts b/packages/coding-agent/test/model-resolver.test.ts index 62bfa6eea..e422c0493 100644 --- a/packages/coding-agent/test/model-resolver.test.ts +++ b/packages/coding-agent/test/model-resolver.test.ts @@ -227,11 +227,11 @@ describe("parseModelPattern", () => { describe("resolveCliModel", () => { test("separates thinking from an unknown provider-prefixed custom model id", () => { const registry = { - getAll: () => [...allModels, openaiCodexBaseModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [...allModels, openaiCodexBaseModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "openai-codex/gpt-5.6-sol:xhigh", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("openai-codex"); @@ -241,12 +241,12 @@ describe("resolveCliModel", () => { }); test("separates thinking from an explicit provider custom model id", () => { const registry = { - getAll: () => [...allModels, openaiCodexBaseModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [...allModels, openaiCodexBaseModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliProvider: "openai-codex", cliModel: "gpt-5.6-sol:high", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.id).toBe("gpt-5.6-sol"); @@ -255,11 +255,11 @@ describe("resolveCliModel", () => { }); test("separates off thinking from a custom model id", () => { const registry = { - getAll: () => [...allModels, openaiCodexBaseModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [...allModels, openaiCodexBaseModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "openai-codex/gpt-5.6-sol:off", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.id).toBe("gpt-5.6-sol"); @@ -268,12 +268,12 @@ describe("resolveCliModel", () => { }); test("preserves an unrecognized colon suffix on a custom model id", () => { const registry = { - getAll: () => [...allModels, openaiCodexBaseModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [...allModels, openaiCodexBaseModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliProvider: "openai-codex", cliModel: "gpt-5.6-sol:preview", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.id).toBe("gpt-5.6-sol:preview"); @@ -287,11 +287,11 @@ describe("resolveCliModel", () => { name: "GPT-5.6 Sol XHigh", }; const registry = { - getAll: () => [...allModels, openaiCodexBaseModel, registeredModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [...allModels, openaiCodexBaseModel, registeredModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "openai-codex/gpt-5.6-sol:xhigh", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model).toBe(registeredModel); @@ -310,12 +310,12 @@ describe("resolveCliModel", () => { name: "OpenAI Foo High", }; const registry = { - getAll: () => [...allModels, inferredProviderModel, registeredGatewayModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [...allModels, inferredProviderModel, registeredGatewayModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "openai/foo:high", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); @@ -325,11 +325,11 @@ describe("resolveCliModel", () => { }); test("resolves --model provider/id without --provider", () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "openai/gpt-4o", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("openai"); @@ -337,12 +337,12 @@ describe("resolveCliModel", () => { }); test("resolves fuzzy patterns within an explicit provider", () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliProvider: "openai", cliModel: "4o", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("openai"); @@ -350,11 +350,11 @@ describe("resolveCliModel", () => { }); test("supports --model : (without explicit --thinking)", () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "sonnet:high", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.id).toBe("claude-sonnet-4-5"); @@ -362,11 +362,11 @@ describe("resolveCliModel", () => { }); test("prefers exact model id match over provider inference (OpenRouter-style ids)", () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "openai/gpt-4o:extended", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("openrouter"); @@ -374,12 +374,12 @@ describe("resolveCliModel", () => { }); test("does not strip invalid :suffix as thinking level in --model (treat as raw id)", () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliProvider: "openai", cliModel: "gpt-4o:extended", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("openai"); @@ -387,12 +387,12 @@ describe("resolveCliModel", () => { }); test("allows custom model ids for explicit providers without double prefixing", () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliProvider: "openrouter", cliModel: "openrouter/openai/ghost-model", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("openrouter"); @@ -400,12 +400,12 @@ describe("resolveCliModel", () => { }); test("scrubs inherited context-window options from explicit provider fallback models", () => { const registry = { - getAll: () => [copilotSelectableBaseModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [copilotSelectableBaseModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliProvider: "github-copilot", cliModel: "future-copilot-model", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("github-copilot"); @@ -414,12 +414,12 @@ describe("resolveCliModel", () => { }); test("returns a clear error when there are no models", () => { const registry = { - getAll: () => [], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliProvider: "openai", cliModel: "gpt-4o", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.model).toBeUndefined(); expect(result.error).toContain("No models available"); @@ -452,11 +452,11 @@ describe("resolveCliModel", () => { maxTokens: 8192, }; const registry = { - getAll: () => [...allModels, zaiModel, gatewayModel], - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => [...allModels, zaiModel, gatewayModel], + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "zai/glm-5", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("zai"); @@ -464,11 +464,11 @@ describe("resolveCliModel", () => { }); test("resolves provider-prefixed fuzzy patterns (openrouter/qwen -> openrouter model)", () => { const registry = { - getAll: () => allModels, - } as unknown as Parameters[0]["modelRegistry"]; + getModels: () => allModels, + } as unknown as Parameters[0]["modelRuntime"]; const result = resolveCliModel({ cliModel: "openrouter/qwen", - modelRegistry: registry, + modelRuntime: registry, }); expect(result.error).toBeUndefined(); expect(result.model?.provider).toBe("openrouter"); diff --git a/packages/coding-agent/test/model-runtime-auth-options.test.ts b/packages/coding-agent/test/model-runtime-auth-options.test.ts new file mode 100644 index 000000000..74f65be45 --- /dev/null +++ b/packages/coding-agent/test/model-runtime-auth-options.test.ts @@ -0,0 +1,256 @@ +import { type AuthType, type CredentialStore, InMemoryCredentialStore } from "@earendil-works/pi-ai"; +import { describe, expect, it } from "vitest"; +import { AuthStorage } from "../src/core/auth-storage.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; + +function authOptions(runtime: ModelRuntime, type?: AuthType) { + return runtime + .getProviders() + .flatMap((provider) => [ + ...(!type || type === "oauth" + ? provider.auth.oauth + ? [{ type: "oauth" as const, provider, method: provider.auth.oauth }] + : [] + : []), + ...(!type || type === "api_key" + ? provider.auth.apiKey + ? [{ type: "api_key" as const, provider, method: provider.auth.apiKey }] + : [] + : []), + ]); +} + +function testModel(id: string) { + return { + id, + name: id, + reasoning: false, + input: ["text"] as ("text" | "image")[], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 10000, + maxTokens: 1000, + }; +} + +describe("ModelRuntime auth options", () => { + it("accepts a pi-ai CredentialStore", async () => { + const credentials = new InMemoryCredentialStore(); + await credentials.modify("anthropic", async () => ({ type: "api_key", key: "stored-key" })); + const runtime = await ModelRuntime.create({ credentials, modelsPath: null }); + + expect((await runtime.getAuth("anthropic"))?.auth.apiKey).toBe("stored-key"); + }); + + it("scopes provider availability reads and records refresh failures", async () => { + const base = new InMemoryCredentialStore(); + const reads: string[] = []; + let failReads = false; + const credentials: CredentialStore = { + read: async (providerId) => { + reads.push(providerId); + if (failReads) throw new Error(`read failed for ${providerId}`); + return base.read(providerId); + }, + list: () => base.list(), + modify: (providerId, fn) => base.modify(providerId, fn), + delete: (providerId) => base.delete(providerId), + }; + const runtime = await ModelRuntime.create({ credentials, modelsPath: null }); + + reads.length = 0; + await runtime.getAvailable("anthropic"); + expect(new Set(reads)).toEqual(new Set(["anthropic"])); + + failReads = true; + await expect(runtime.getAvailable("anthropic")).rejects.toThrow("Credential store read failed for anthropic"); + expect(runtime.getError()).toContain("Availability refresh: Credential store read failed for anthropic"); + + failReads = false; + await runtime.getAvailable(); + expect(runtime.getError()).toBeUndefined(); + }); + + it("projects provider-owned methods, names, and status", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + const options = authOptions(runtime); + + expect(options).toEqual( + expect.arrayContaining([ + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "amazon-bedrock", name: "Amazon Bedrock" }), + method: expect.objectContaining({ name: "AWS credentials or bearer token" }), + }), + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "google-vertex", name: "Google Vertex AI" }), + method: expect.objectContaining({ name: "Google Cloud credentials" }), + }), + expect.objectContaining({ + type: "oauth", + provider: expect.objectContaining({ id: "anthropic", name: "Anthropic" }), + }), + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "cloudflare-ai-gateway", name: "Cloudflare AI Gateway" }), + }), + expect.objectContaining({ + type: "api_key", + provider: expect.objectContaining({ id: "cloudflare-workers-ai", name: "Cloudflare Workers AI" }), + }), + ]), + ); + expect(authOptions(runtime, "api_key").every((option) => option.type === "api_key")).toBe(true); + expect(authOptions(runtime, "oauth").every((option) => option.type === "oauth")).toBe(true); + expect(options.some((option) => option.provider.id === "openai-codex" && option.type === "api_key")).toBe(false); + }); + + it("attaches the provider's active auth status to every method option", async () => { + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory({ + anthropic: { + type: "oauth", + access: "access", + refresh: "refresh", + expires: Date.now() + 60_000, + }, + }), + modelsPath: null, + }); + + const options = authOptions(runtime).filter((option) => option.provider.id === "anthropic"); + expect(options).toHaveLength(2); + expect(await runtime.checkAuth("anthropic")).toMatchObject({ type: "oauth" }); + }); + + it("constructs an API key method for an extension API-key provider", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + runtime.registerProvider("extension-api-key", { + name: "Extension API Key", + baseUrl: "https://example.test/v1", + apiKey: "$EXTENSION_TEST_API_KEY", + api: "openai-completions", + models: [testModel("extension-model")], + }); + + const options = authOptions(runtime).filter((option) => option.provider.id === "extension-api-key"); + expect(options).toHaveLength(1); + expect(options[0]).toMatchObject({ + type: "api_key", + provider: { id: "extension-api-key", name: "Extension API Key" }, + method: { name: "API key" }, + }); + expect(options[0]?.method.login).toBeTypeOf("function"); + }); + + it("resolves configured auth from request-scoped environment overrides", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + runtime.registerProvider("request-env-provider", { + baseUrl: "https://example.test/v1", + apiKey: "$REQUEST_SCOPED_API_KEY", + headers: { "x-request-value": "$REQUEST_SCOPED_HEADER" }, + api: "openai-completions", + models: [testModel("request-env-model")], + }); + + const auth = await runtime.getAuth("request-env-provider", { + env: { REQUEST_SCOPED_API_KEY: "request-key", REQUEST_SCOPED_HEADER: "request-header" }, + }); + + expect(auth?.auth).toEqual({ apiKey: "request-key", headers: { "x-request-value": "request-header" } }); + }); + + it("lets an explicit Authorization header override authHeader case-insensitively", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + let capturedHeaders: Record | undefined; + runtime.registerProvider("auth-header-provider", { + baseUrl: "https://example.test/v1", + apiKey: "generated-key", + authHeader: true, + api: "openai-completions", + streamSimple: (_model, _context, options) => { + capturedHeaders = options?.headers; + throw new Error("captured"); + }, + models: [testModel("auth-header-model")], + }); + const model = runtime.getModel("auth-header-provider", "auth-header-model"); + expect(model).toBeDefined(); + + await runtime.completeSimple(model!, { messages: [] }, { headers: { authorization: "Explicit token" } }); + + expect(capturedHeaders).toEqual({ authorization: "Explicit token" }); + }); + + it("transforms fully assembled headers once without forwarding the transform", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + let capturedHeaders: Record | undefined; + let transforms = 0; + runtime.registerProvider("header-provider", { + baseUrl: "https://example.test/v1", + apiKey: "generated-key", + authHeader: true, + headers: { "x-provider": "provider" }, + api: "openai-completions", + streamSimple: (_model, _context, options) => { + expect(options).not.toHaveProperty("transformHeaders"); + capturedHeaders = options?.headers; + throw new Error("captured"); + }, + models: [{ ...testModel("header-model"), headers: { "x-model": "model" } }], + }); + const model = runtime.getModel("header-provider", "header-model"); + expect(model).toBeDefined(); + + await runtime.completeSimple( + model!, + { messages: [] }, + { + headers: { "x-explicit": "explicit" }, + transformHeaders: async (headers) => { + transforms++; + expect(headers).toEqual({ + Authorization: "Bearer generated-key", + "x-provider": "provider", + "x-model": "model", + "x-explicit": "explicit", + }); + return { ...headers, "x-transformed": "yes" }; + }, + }, + ); + + expect(transforms).toBe(1); + expect(capturedHeaders).toEqual({ + Authorization: "Bearer generated-key", + "x-provider": "provider", + "x-model": "model", + "x-explicit": "explicit", + "x-transformed": "yes", + }); + }); + + it("does not fabricate an API key method for an extension OAuth-only provider", async () => { + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory(), modelsPath: null }); + runtime.registerProvider("extension-oauth", { + name: "Extension OAuth", + baseUrl: "https://example.test/v1", + api: "openai-completions", + oauth: { + name: "Extension subscription", + login: async () => ({ access: "access", refresh: "refresh", expires: Date.now() + 60_000 }), + refreshToken: async (credentials) => credentials, + getApiKey: (credentials) => credentials.access, + }, + models: [testModel("extension-model")], + }); + + const options = authOptions(runtime).filter((option) => option.provider.id === "extension-oauth"); + expect(options).toHaveLength(1); + expect(options[0]).toMatchObject({ + type: "oauth", + provider: { id: "extension-oauth", name: "Extension OAuth" }, + method: { name: "Extension subscription" }, + }); + }); +}); diff --git a/packages/coding-agent/test/model-runtime-cloudflare-compat.test.ts b/packages/coding-agent/test/model-runtime-cloudflare-compat.test.ts new file mode 100644 index 000000000..f7348d813 --- /dev/null +++ b/packages/coding-agent/test/model-runtime-cloudflare-compat.test.ts @@ -0,0 +1,52 @@ +import { describe, expect, it } from "vitest"; +import { AuthStorage } from "../src/core/auth-storage.ts"; +import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; + +async function createCloudflareRuntime(): Promise<{ modelRuntime: ModelRuntime; modelRegistry: ModelRegistry }> { + const authStorage = AuthStorage.inMemory(); + await authStorage.modify("cloudflare-ai-gateway", async () => ({ + type: "api_key", + key: "test-token", + env: { + CLOUDFLARE_ACCOUNT_ID: "test-account", + CLOUDFLARE_GATEWAY_ID: "test-gateway", + }, + })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); + return { modelRuntime, modelRegistry: new ModelRegistry(modelRuntime) }; +} + +describe("Cloudflare compatibility auth", () => { + it("resolves Cloudflare headers and caller-controlled endpoint environment through ModelRuntime", async () => { + const { modelRuntime } = await createCloudflareRuntime(); + const model = modelRuntime.getModel("cloudflare-ai-gateway", "workers-ai/@cf/moonshotai/kimi-k2.5"); + expect(model).toBeDefined(); + + const resolution = await modelRuntime.getAuth(model!); + + expect(resolution?.auth.headers?.["cf-aig-authorization"]).toBe("Bearer test-token"); + expect(resolution?.env).toEqual({ + CLOUDFLARE_ACCOUNT_ID: "test-account", + CLOUDFLARE_GATEWAY_ID: "test-gateway", + }); + }); + + it("keeps the extension facade request-auth projection intentionally base-URL-free", async () => { + const { modelRegistry } = await createCloudflareRuntime(); + const model = modelRegistry.find("cloudflare-ai-gateway", "workers-ai/@cf/moonshotai/kimi-k2.5"); + expect(model).toBeDefined(); + + const auth = await modelRegistry.getApiKeyAndHeaders(model!); + expect(auth.ok).toBe(true); + if (!auth.ok) throw new Error(auth.error); + + expect(auth.apiKey).toBeUndefined(); + expect(auth.headers?.["cf-aig-authorization"]).toBe("Bearer test-token"); + expect(auth.env).toEqual({ + CLOUDFLARE_ACCOUNT_ID: "test-account", + CLOUDFLARE_GATEWAY_ID: "test-gateway", + }); + expect(auth).not.toHaveProperty("baseUrl"); + }); +}); diff --git a/packages/coding-agent/test/model-runtime-modify-models-compat.test.ts b/packages/coding-agent/test/model-runtime-modify-models-compat.test.ts new file mode 100644 index 000000000..c5f211463 --- /dev/null +++ b/packages/coding-agent/test/model-runtime-modify-models-compat.test.ts @@ -0,0 +1,201 @@ +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 { AuthStorage } from "../src/core/auth-storage.ts"; +import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; + +function model(id: string): Model<"openai-completions"> { + return { + id, + name: id, + api: "openai-completions", + provider: "extension-oauth", + baseUrl: "https://example.test/v1", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 1000, + maxTokens: 100, + }; +} + +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 () => { + const modelsStore = new InMemoryModelsStore(); + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory(), + modelsStore, + modelsPath: null, + allowModelNetwork: false, + }); + runtime.registerProvider("extension-dynamic", { + baseUrl: "http://localhost:8080/v1", + apiKey: "local", + api: "openai-completions", + refreshModels: async () => [ + { + ...model("live"), + provider: "extension-dynamic", + baseUrl: "http://localhost:8080/v1", + }, + ], + }); + + await runtime.refresh({ allowNetwork: false }); + expect(runtime.getModel("extension-dynamic", "live")).toBeDefined(); + expect(await modelsStore.read("extension-dynamic")).toBeUndefined(); + }); + + it("applies legacy OAuth modifyModels after async credential initialization", async () => { + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory({ + "extension-oauth": { + type: "oauth", + access: "access", + refresh: "refresh", + expires: Date.now() + 60_000, + }, + }), + modelsStore: new InMemoryModelsStore(), + modelsPath: null, + allowModelNetwork: false, + }); + runtime.registerProvider("extension-oauth", { + baseUrl: "https://example.test/v1", + api: "openai-completions", + models: [model("base")], + oauth: { + name: "Extension OAuth", + login: async () => { + throw new Error("not used"); + }, + refreshToken: async (credential) => credential, + getApiKey: (credential) => credential.access, + modifyModels: (models, credential) => + credential.access === "access" ? [...models, model("credential-model")] : models, + }, + }); + + await runtime.refresh({ allowNetwork: false }); + expect(runtime.getModel("extension-oauth", "base")).toBeDefined(); + expect(runtime.getModel("extension-oauth", "credential-model")).toBeDefined(); + + await runtime.logout("extension-oauth"); + expect(runtime.getModel("extension-oauth", "credential-model")).toBeUndefined(); + }); +}); diff --git a/packages/coding-agent/test/model-runtime-refresh-bounds.test.ts b/packages/coding-agent/test/model-runtime-refresh-bounds.test.ts new file mode 100644 index 000000000..dc4df33a7 --- /dev/null +++ b/packages/coding-agent/test/model-runtime-refresh-bounds.test.ts @@ -0,0 +1,50 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; +import { AuthStorage } from "../src/core/auth-storage.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; + +afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllEnvs(); +}); + +describe("model refresh timeout boundaries", () => { + it("applies modelRefreshTimeoutMs only to the create-time refresh", async () => { + let createSignal: AbortSignal | undefined; + const refresh = vi.spyOn(ModelRuntime.prototype, "refresh").mockImplementation(async (options = {}) => { + createSignal = options.signal; + if (options.signal) { + await new Promise((resolve) => options.signal?.addEventListener("abort", () => resolve(), { once: true })); + } + return { aborted: options.signal?.aborted ?? false, errors: new Map() }; + }); + + await ModelRuntime.create({ + credentials: AuthStorage.inMemory(), + modelsPath: null, + allowModelNetwork: true, + modelRefreshTimeoutMs: 5, + }); + + expect(createSignal).toBeInstanceOf(AbortSignal); + expect(createSignal?.aborted).toBe(true); + expect(refresh).toHaveBeenCalledWith(expect.objectContaining({ signal: createSignal })); + }); + + it("treats a false offline flag as online", async () => { + vi.stubEnv("ATOMIC_OFFLINE", "0"); + vi.stubEnv("PI_OFFLINE", ""); + let createSignal: AbortSignal | undefined; + vi.spyOn(ModelRuntime.prototype, "refresh").mockImplementation(async (options = {}) => { + createSignal = options.signal; + return { aborted: false, errors: new Map() }; + }); + + await ModelRuntime.create({ + credentials: AuthStorage.inMemory(), + modelsPath: null, + allowModelNetwork: true, + }); + + expect(createSignal).toBeInstanceOf(AbortSignal); + }); +}); diff --git a/packages/coding-agent/test/model-runtime-test-utils.ts b/packages/coding-agent/test/model-runtime-test-utils.ts new file mode 100644 index 000000000..b19ab33db --- /dev/null +++ b/packages/coding-agent/test/model-runtime-test-utils.ts @@ -0,0 +1,42 @@ +import type { CredentialStore } from "@earendil-works/pi-ai"; +import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; + +const runtimes = new WeakMap(); + +function wrap(runtime: ModelRuntime): ModelRegistry { + const registry = new ModelRegistry(runtime); + runtimes.set(registry, runtime); + return registry; +} + +export async function createModelRegistry(credentials: CredentialStore, modelsPath?: string): Promise { + return wrap(await ModelRuntime.create({ credentials, modelsPath, allowModelNetwork: false })); +} + +export async function createInMemoryModelRegistry(credentials: CredentialStore): Promise { + return wrap(await ModelRuntime.create({ credentials, modelsPath: null, allowModelNetwork: false })); +} + +export function getModelRuntime(modelRegistry: ModelRegistry): ModelRuntime { + const runtime = runtimes.get(modelRegistry); + if (!runtime) throw new Error("ModelRegistry was not created by the test helper"); + return runtime; +} + +/** Minimal runtime-shaped test double for pure model-selection/UI tests. */ +export function fakeModelRuntime(overrides: Partial = {}): ModelRuntime { + return { + getModels: () => [], + getAvailableSnapshot: () => [], + getModel: () => undefined, + getProvider: () => undefined, + getProviders: () => [], + hasConfiguredAuth: () => false, + isUsingOAuth: () => false, + getAuth: async () => undefined, + getProviderAuthStatus: () => ({ configured: false }), + refresh: async () => ({ ok: true, providers: [] }), + ...overrides, + } as ModelRuntime; +} diff --git a/packages/coding-agent/test/model-selector-refresh-status.test.ts b/packages/coding-agent/test/model-selector-refresh-status.test.ts index 42e1c5972..5e4d9544b 100644 --- a/packages/coding-agent/test/model-selector-refresh-status.test.ts +++ b/packages/coding-agent/test/model-selector-refresh-status.test.ts @@ -1,8 +1,7 @@ import type { Api, Model } from "@earendil-works/pi-ai/compat"; import type { TUI } from "@earendil-works/pi-tui"; -import { describe, expect, it } from "vitest"; -import { ENV_OFFLINE } from "../src/config.ts"; -import type { ModelRegistry } from "../src/core/model-registry.ts"; +import { afterEach, describe, expect, it, vi } from "vitest"; +import type { ModelRuntime } from "../src/core/model-runtime.ts"; import type { SettingsManager } from "../src/core/settings-manager.ts"; import { ModelSelectorComponent } from "../src/modes/interactive/components/model-selector.ts"; import { initTheme } from "../src/modes/interactive/theme/theme.ts"; @@ -20,26 +19,26 @@ const model = { maxTokens: 1024, } as Model; -type RefreshResult = Awaited>; +type RefreshResult = Awaited>; function deferred(): { promise: Promise; resolve: (value: T) => void } { let resolve!: (value: T) => void; return { promise: new Promise((done) => (resolve = done)), resolve }; } -function createSelector(refresh: ModelRegistry["refresh"]): ModelSelectorComponent { +function createSelector(refresh: ModelRuntime["refresh"]): ModelSelectorComponent { initTheme("dark"); - const registry = { + const runtime = { refresh, getError: () => undefined, - getAvailable: async () => [model], - find: () => model, - } as unknown as ModelRegistry; + getAvailableSnapshot: () => [model], + getModel: () => model, + } as unknown as ModelRuntime; return new ModelSelectorComponent( { requestRender: () => {} } as unknown as TUI, model, { setDefaultModelAndProvider: () => {} } as unknown as SettingsManager, - registry, + runtime, [], () => {}, () => {}, @@ -51,6 +50,10 @@ async function renderedAfterWork(selector: ModelSelectorComponent): Promise { + vi.useRealTimers(); +}); + describe("model selector catalog refresh status", () => { it("renders the stale snapshot immediately and then concise success", async () => { const refresh = deferred(); @@ -64,14 +67,14 @@ describe("model selector catalog refresh status", () => { expect(refreshed).toContain("Model catalogs refreshed."); }); - it("keeps available models and identifies a partial provider failure", async () => { + it("keeps cached models and identifies a partial provider failure", async () => { const selector = createSelector(async () => ({ aborted: false, errors: new Map([["configured", new Error("offline")]]), })); const rendered = await renderedAfterWork(selector); expect(rendered).toContain("cached-model"); - expect(rendered).toContain("Could not refresh configured; showing available models."); + expect(rendered).toContain("Could not refresh configured; showing cached models."); }); it("summarizes multiple provider errors", async () => { @@ -83,30 +86,37 @@ describe("model selector catalog refresh status", () => { ]), })); const rendered = await renderedAfterWork(selector); - expect(rendered).toContain("Could not refresh 2 model catalogs; showing available models."); + expect(rendered).toContain("Could not refresh 2 model catalogs; showing cached models."); }); - it("reports timeout while retaining cached models", async () => { - const selector = createSelector(async () => ({ aborted: true, errors: new Map() })); + it("aborts a slow refresh after the selector timeout and keeps cached models", async () => { + vi.useFakeTimers(); + const selector = createSelector( + (options) => + new Promise((resolve) => { + options?.signal?.addEventListener( + "abort", + () => resolve({ aborted: true, errors: new Map() }), + { once: true }, + ); + }), + ); + await vi.advanceTimersByTimeAsync(15_000); + vi.useRealTimers(); const rendered = await renderedAfterWork(selector); expect(rendered).toContain("cached-model"); expect(rendered).toContain("Model refresh timed out; showing cached models."); }); - it("keeps selector refreshes cache-only in offline mode", async () => { - const previous = process.env[ENV_OFFLINE]; - process.env[ENV_OFFLINE] = "1"; - let observed: Parameters[0]; - try { - const selector = createSelector(async (options) => { - observed = options; - return { aborted: false, errors: new Map() }; - }); - await renderedAfterWork(selector); - expect(observed).toMatchObject({ allowNetwork: false, timeoutMs: 15_000 }); - } finally { - if (previous === undefined) delete process.env[ENV_OFFLINE]; - else process.env[ENV_OFFLINE] = previous; - } + it("delegates network gating to the runtime instead of passing allowNetwork", async () => { + let observed: Parameters[0]; + const selector = createSelector(async (options) => { + observed = options; + return { aborted: false, errors: new Map() }; + }); + await renderedAfterWork(selector); + expect(observed).toMatchObject({ signal: expect.any(AbortSignal) }); + expect(observed).not.toHaveProperty("allowNetwork"); + expect(observed).not.toHaveProperty("timeoutMs"); }); }); diff --git a/packages/coding-agent/test/models-store-snapshot.test.ts b/packages/coding-agent/test/models-store-snapshot.test.ts index 42f596d98..71a72fc25 100644 --- a/packages/coding-agent/test/models-store-snapshot.test.ts +++ b/packages/coding-agent/test/models-store-snapshot.test.ts @@ -15,7 +15,7 @@ const model = { maxTokens: 1024, } as Model; -test("in-memory model store reads return isolated snapshots", async () => { +test("in-memory model store preserves the exact stored entry", async () => { const store = new InMemoryCodingAgentModelsStore(); await store.write("snapshot-provider", { models: [model], checkedAt: 1 }); const snapshot = await store.read("snapshot-provider"); @@ -23,6 +23,6 @@ test("in-memory model store reads return isolated snapshots", async () => { (snapshot!.models[0] as { id: string }).id = "mutated-without-write"; const persisted = await store.read("snapshot-provider"); - expect(persisted?.checkedAt).toBe(1); - expect(persisted?.models[0]?.id).toBe("seed"); + expect(persisted?.checkedAt).toBe(2); + expect(persisted?.models[0]?.id).toBe("mutated-without-write"); }); diff --git a/packages/coding-agent/test/models-store.test.ts b/packages/coding-agent/test/models-store.test.ts new file mode 100644 index 000000000..483153974 --- /dev/null +++ b/packages/coding-agent/test/models-store.test.ts @@ -0,0 +1,51 @@ +import { existsSync, mkdirSync, rmSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import type { Model } from "@earendil-works/pi-ai"; +import { afterEach, describe, expect, it } from "vitest"; +import { FileModelsStore } from "../src/core/models-store.ts"; + +const tempDirs: string[] = []; + +afterEach(() => { + for (const path of tempDirs.splice(0)) { + if (existsSync(path)) rmSync(path, { recursive: true }); + } +}); + +function model(provider: string, id: string): Model<"openai-completions"> { + return { + id, + name: id, + api: "openai-completions", + provider, + baseUrl: "https://example.test/v1", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 1000, + maxTokens: 100, + }; +} + +describe("FileModelsStore", () => { + it("persists provider catalogs without replacing unrelated providers", async () => { + const dir = join(tmpdir(), `pi-models-store-${Date.now()}-${Math.random().toString(36).slice(2)}`); + tempDirs.push(dir); + mkdirSync(dir, { recursive: true }); + const path = join(dir, "models-store.json"); + const store = new FileModelsStore(path); + + await store.write("one", { models: [model("one", "m1")], checkedAt: 100 }); + await store.write("two", { models: [model("two", "m2")], checkedAt: 200 }); + + const reloaded = new FileModelsStore(path); + expect((await reloaded.read("one"))?.models.map((entry) => entry.id)).toEqual(["m1"]); + expect((await reloaded.read("one"))?.checkedAt).toBe(100); + expect((await reloaded.read("two"))?.models.map((entry) => entry.id)).toEqual(["m2"]); + + await reloaded.delete("one"); + expect(await reloaded.read("one")).toBeUndefined(); + expect((await reloaded.read("two"))?.models.map((entry) => entry.id)).toEqual(["m2"]); + }); +}); diff --git a/packages/coding-agent/test/oauth-cancellation.test.ts b/packages/coding-agent/test/oauth-cancellation.test.ts index 23b87a751..3cac837de 100644 --- a/packages/coding-agent/test/oauth-cancellation.test.ts +++ b/packages/coding-agent/test/oauth-cancellation.test.ts @@ -2,15 +2,12 @@ import { ModelsError } from "@earendil-works/pi-ai"; import { afterEach, describe, expect, it, vi } from "vitest"; import { isOAuthLoginCancelled, - loginOAuthProvider, normalizeOAuthLoginError, OAuthLoginTransactionError, - registerLegacyOAuthProvider, - resetLegacyOAuthProviders, -} from "../src/core/oauth-provider-bridge.ts"; +} from "../src/core/oauth-login.ts"; import { loginRuntimeOAuthProvider } from "../src/core/agent-session-runtime-auth.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"; const callbacks = (signal?: AbortSignal) => ({ signal, @@ -21,42 +18,53 @@ const callbacks = (signal?: AbortSignal) => ({ onSelect: async () => undefined, }); -afterEach(() => { - resetLegacyOAuthProviders(); - vi.restoreAllMocks(); -}); +afterEach(() => vi.restoreAllMocks()); describe("OAuth cancellation normalization", () => { - it("normalizes a pre-aborted builtin Kimi login and preserves the native cause", async () => { + it("recognizes exact cancellation, AbortError causes, and active signals", () => { const controller = new AbortController(); - controller.abort(); - let caught: unknown; - try { - await loginOAuthProvider("kimi-coding", callbacks(controller.signal)); - } catch (error) { - caught = error; - } - expect(caught).toBeInstanceOf(Error); - expect((caught as Error).message).toBe("Login cancelled"); - expect((caught as Error).cause).toBeInstanceOf(Error); - expect(((caught as Error).cause as Error).name).toBe("AbortError"); + controller.abort(new Error("cancelled by caller")); + expect(isOAuthLoginCancelled(controller.signal.reason, controller.signal)).toBe(true); + expect(isOAuthLoginCancelled(new Error("Login cancelled"))).toBe(true); + expect(isOAuthLoginCancelled(new Error("wrapped", { cause: new DOMException("aborted", "AbortError") }))).toBe(true); + }); + + it("does not classify ordinary failures or completed-transaction failures as cancellation", () => { + expect(isOAuthLoginCancelled(new Error("provider failed"))).toBe(false); + expect(isOAuthLoginCancelled(new OAuthLoginTransactionError(new DOMException("aborted", "AbortError")))).toBe(false); + }); + + it("normalizes cancellation while preserving the native cause", () => { + const cause = new DOMException("aborted", "AbortError"); + const normalized = normalizeOAuthLoginError(cause); + expect(normalized).toBeInstanceOf(Error); + expect((normalized as Error).message).toBe("Login cancelled"); + expect((normalized as Error).cause).toBe(cause); }); - it("does not write through the builtin registry route for pre-aborted Kimi", async () => { + it("preserves genuine provider and persistence failures", () => { + const providerFailure = new Error("provider failed"); + expect(normalizeOAuthLoginError(providerFailure)).toBe(providerFailure); + const persistenceFailure = new OAuthLoginTransactionError(new DOMException("aborted", "AbortError")); + expect(normalizeOAuthLoginError(persistenceFailure)).toBe(persistenceFailure); + }); + + it("does not write provider-owned credentials for a pre-aborted builtin Kimi login", async () => { const previous = { type: "api_key" as const, key: "previous" }; const authStorage = AuthStorage.inMemory({ "kimi-coding": previous }); - const registry = ModelRegistry.inMemory(authStorage); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const controller = new AbortController(); controller.abort(); await expect(loginRuntimeOAuthProvider( - { modelRegistry: registry } as never, + { modelRuntime } as never, "kimi-coding", callbacks(controller.signal), )).rejects.toMatchObject({ message: "Login cancelled" }); - expect(authStorage.get("kimi-coding")).toEqual(previous); + expect(await authStorage.read("kimi-coding")).toEqual(previous); }); + it("normalizes Kimi cancellation immediately after the device code is shown", async () => { vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response(JSON.stringify({ device_code: "device", @@ -66,106 +74,27 @@ describe("OAuth cancellation normalization", () => { interval: 1, expires_in: 60, }), { status: 200, headers: { "content-type": "application/json" } })); + const authStorage = AuthStorage.inMemory(); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const controller = new AbortController(); - let caught: unknown; - try { - await loginOAuthProvider("kimi-coding", { + + await expect(loginRuntimeOAuthProvider( + { modelRuntime } as never, + "kimi-coding", + { ...callbacks(controller.signal), onDeviceCode: () => controller.abort(new Error("user stopped")), - }); - } catch (error) { - caught = error; - } - expect(caught).toMatchObject({ message: "Login cancelled" }); - expect((caught as Error).cause).toBeInstanceOf(Error); + }, + )).rejects.toMatchObject({ message: "Login cancelled" }); expect(fetch).toHaveBeenCalledOnce(); + expect(await authStorage.read("kimi-coding")).toBeUndefined(); }); - - it("recognizes only exact signal, AbortError cause, and compatibility cancellation", async () => { - const reason = new Error("stop"); - const controller = new AbortController(); - controller.abort(reason); - const nested = new Error("wrapper", { cause: new DOMException("aborted", "AbortError") }); - const cycle = new Error("cycle"); - (cycle as Error & { cause?: unknown }).cause = cycle; - - expect(isOAuthLoginCancelled(reason, controller.signal)).toBe(true); - expect(isOAuthLoginCancelled(nested)).toBe(true); - expect(isOAuthLoginCancelled(new Error("Login cancelled"))).toBe(true); - expect(isOAuthLoginCancelled("Login cancelled")).toBe(true); - expect(isOAuthLoginCancelled(new Error("wrapper", { cause: "Login cancelled" }))).toBe(true); - expect(isOAuthLoginCancelled(cycle)).toBe(false); - for (const message of ["login cancelled by provider", "Abort request failed", "request timeout", "access denied"]) { - expect(isOAuthLoginCancelled(new Error(message))).toBe(false); - } - - const preAborted = new AbortController(); - preAborted.abort(); - await expect(loginOAuthProvider("unknown-provider", callbacks(preAborted.signal))) - .rejects.toThrow("Unknown OAuth provider: unknown-provider"); - }); - - it("normalizes a legacy provider abort but preserves genuine provider failures", async () => { - const abort = new DOMException("operation aborted", "AbortError"); - registerLegacyOAuthProvider("legacy", { - name: "Legacy", - login: async () => { throw abort; }, - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }); - await expect(loginOAuthProvider("legacy", callbacks())).rejects.toMatchObject({ - message: "Login cancelled", - cause: abort, - }); - let runtimeError: unknown; - try { - await loginRuntimeOAuthProvider( - { modelRegistry: { authStorage: AuthStorage.inMemory() } } as never, - "legacy", - callbacks(), - ); - } catch (error) { - runtimeError = error; - } - expect(runtimeError).toMatchObject({ message: "Login cancelled", cause: abort }); - - const failure = new Error("HTTP 403 access denied"); - expect(normalizeOAuthLoginError(failure)).toBe(failure); - }); - - it("normalizes provider-owned runtime aborts that bypass the compatibility bridge", async () => { - const controller = new AbortController(); + it("normalizes provider-owned runtime aborts", async () => { const abort = new DOMException("aborted", "AbortError"); - const session = { - modelRegistry: { - login: async () => { throw abort; }, - authStorage: { login: async () => { throw new Error("unexpected legacy route"); } }, - }, - }; - await expect(loginRuntimeOAuthProvider(session as never, "kimi-coding", callbacks(controller.signal))) - .rejects.toMatchObject({ message: "Login cancelled", cause: abort }); - }); + const session = { modelRuntime: { login: async () => { throw abort; } } }; - it("does not persist when a legacy provider ignores an abort and returns credentials", async () => { - const controller = new AbortController(); - const previous = { type: "api_key" as const, key: "previous" }; - const authStorage = AuthStorage.inMemory({ legacy: previous }); - registerLegacyOAuthProvider("legacy", { - name: "Legacy", - login: async () => { - controller.abort(new Error("user stopped")); - return { access: "must-not-save", refresh: "refresh", expires: Date.now() + 60_000 }; - }, - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }); - - await expect(loginRuntimeOAuthProvider( - { modelRegistry: { authStorage } } as never, - "legacy", - callbacks(controller.signal), - )).rejects.toMatchObject({ message: "Login cancelled", cause: controller.signal.reason }); - expect(authStorage.get("legacy")).toEqual(previous); + await expect(loginRuntimeOAuthProvider(session as never, "kimi-coding", callbacks())) + .rejects.toMatchObject({ message: "Login cancelled", cause: abort }); }); it("does not relabel a persistence failure when the signal aborts during persistence", async () => { @@ -175,41 +104,15 @@ describe("OAuth cancellation normalization", () => { cause: persistenceCause, }); const session = { - modelRegistry: { + modelRuntime: { login: async () => { controller.abort(new Error("user stopped")); throw persistenceFailure; }, - authStorage: { login: async () => { throw new Error("unexpected legacy route"); } }, }, }; - await expect(loginRuntimeOAuthProvider(session as never, "openrouter", callbacks(controller.signal))) - .rejects.toEqual(new OAuthLoginTransactionError(persistenceFailure)); - }); - it("does not relabel a legacy custom credential-store AbortError", async () => { - registerLegacyOAuthProvider("legacy", { - name: "Legacy", - login: async () => ({ access: "access", refresh: "refresh", expires: Date.now() + 60_000 }), - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }); - const controller = new AbortController(); - const persistenceFailure = new DOMException("credential write aborted", "AbortError"); - const session = { - modelRegistry: { - authStorage: { - asCredentialStore: () => ({ - modify: async () => { - controller.abort(new Error("user stopped")); - throw persistenceFailure; - }, - }), - }, - }, - }; - - await expect(loginRuntimeOAuthProvider(session as never, "legacy", callbacks(controller.signal))) + await expect(loginRuntimeOAuthProvider(session as never, "openrouter", callbacks(controller.signal))) .rejects.toEqual(new OAuthLoginTransactionError(persistenceFailure)); }); }); diff --git a/packages/coding-agent/test/oauth-provider-metadata.test.ts b/packages/coding-agent/test/oauth-provider-metadata.test.ts new file mode 100644 index 000000000..1f99819a6 --- /dev/null +++ b/packages/coding-agent/test/oauth-provider-metadata.test.ts @@ -0,0 +1,43 @@ +import type { Provider } from "@earendil-works/pi-ai"; +import { builtinProviders } from "@earendil-works/pi-ai/providers/all"; +import { describe, expect, it } from "vitest"; +import { collectOAuthProviderMetadata } from "../src/core/oauth-provider-metadata.ts"; +import type { ProviderConfigInput } from "../src/core/provider-composer.ts"; + +function oauthProvider(id: string, loginLabel?: string): Provider { + return { + id, + name: id === "anthropic" ? "Anthropic" : "OpenAI Codex", + auth: { + oauth: { + name: `${id} OAuth`, + ...(loginLabel ? { loginLabel } : {}), + login: async () => ({ type: "oauth", access: "token" }), + refresh: async (credential) => credential, + toAuth: async () => ({ key: "token" }), + }, + }, + } as Provider; +} + +describe("collectOAuthProviderMetadata", () => { + it("preserves builtin callback-server and login-label metadata", () => { + const metadata = collectOAuthProviderMetadata(builtinProviders(), new Map()); + + expect(metadata.find(({ id }) => id === "anthropic")).toMatchObject({ usesCallbackServer: true }); + expect(metadata.find(({ id }) => id === "openai-codex")).toMatchObject({ usesCallbackServer: true }); + expect(metadata.find(({ id }) => id === "xai")).toMatchObject({ + loginLabel: "Sign in with SuperGrok or X Premium", + }); + }); + + it("prefers explicit extension metadata over builtin defaults", () => { + const extensions = new Map([ + ["anthropic", { oauth: { loginLabel: "Corporate Claude", usesCallbackServer: false } }], + ]); + + expect(collectOAuthProviderMetadata([oauthProvider("anthropic", "Sign in to Claude")], extensions)).toEqual([ + { id: "anthropic", name: "Anthropic", loginLabel: "Corporate Claude", usesCallbackServer: false }, + ]); + }); +}); diff --git a/packages/coding-agent/test/oauth-selector.test.ts b/packages/coding-agent/test/oauth-selector.test.ts index 4e031a127..8229d9ef1 100644 --- a/packages/coding-agent/test/oauth-selector.test.ts +++ b/packages/coding-agent/test/oauth-selector.test.ts @@ -1,14 +1,11 @@ import { setKeybindings } from "@earendil-works/pi-tui"; -import { afterEach, beforeAll, beforeEach, describe, expect, it } from "vitest"; -import { AuthStorage } from "../src/core/auth-storage.ts"; +import { beforeAll, beforeEach, describe, expect, it } from "vitest"; import { KeybindingsManager } from "../src/core/keybindings.ts"; -import { BUILT_IN_PROVIDER_DISPLAY_NAMES } from "../src/core/provider-display-names.ts"; import { OAuthSelectorComponent } from "../src/modes/interactive/components/oauth-selector.ts"; -import { isApiKeyLoginProvider } from "../src/modes/interactive/interactive-mode.ts"; +import { InteractiveMode } from "../src/modes/interactive/interactive-mode.ts"; import { initTheme } from "../src/modes/interactive/theme/theme.ts"; import { stripAnsi } from "../src/utils/ansi.ts"; - -const originalOpenAiApiKey = process.env.OPENAI_API_KEY; +import { fakeModelRuntime } from "./model-runtime-test-utils.ts"; describe("OAuthSelectorComponent", () => { beforeAll(() => { @@ -19,120 +16,181 @@ describe("OAuthSelectorComponent", () => { setKeybindings(new KeybindingsManager()); }); - afterEach(() => { - if (originalOpenAiApiKey === undefined) { - delete process.env.OPENAI_API_KEY; - } else { - process.env.OPENAI_API_KEY = originalOpenAiApiKey; - } + it("projects provider-owned auth options without provider-specific filtering", () => { + const getLoginProviderOptions = ( + InteractiveMode as unknown as { + prototype: { + getLoginProviderOptions( + this: object, + authType?: "oauth" | "api_key", + ): Array<{ id: string; name: string; authType: string; method?: { name: string; login?: unknown } }>; + }; + } + ).prototype.getLoginProviderOptions; + const providers = [ + { + id: "anthropic", + name: "Anthropic", + auth: { + oauth: { name: "Anthropic (Claude Pro/Max)", login: async () => ({}) }, + apiKey: { name: "Anthropic API key", login: async () => ({}) }, + }, + }, + { + id: "google-vertex", + name: "Google Vertex AI", + auth: { apiKey: { name: "Google Cloud credentials" } }, + }, + ]; + const fakeThis = { + session: { + modelRuntime: { + getProviders: () => providers, + getOAuthProviderMetadata: () => [{ id: "anthropic", name: "Anthropic" }], + getProviderAuthStatus: () => ({ configured: false }), + isUsingOAuth: () => false, + }, + }, + }; + + const apiKeyOptions = getLoginProviderOptions.call(fakeThis, "api_key"); + expect(apiKeyOptions).toMatchObject([ + { id: "anthropic", name: "Anthropic", authType: "api_key" }, + { id: "google-vertex", name: "Google Vertex AI", authType: "api_key" }, + ]); + expect(getLoginProviderOptions.call(fakeThis, "oauth")).toMatchObject([ + { id: "anthropic", name: "Anthropic", authType: "oauth" }, + ]); }); - it("uses built-in provider metadata to distinguish API-key-capable providers", () => { - const oauthProviderIds = new Set(["anthropic", "github-copilot", "openai-codex", "custom-oauth"]); - const builtInProviderIds = new Set(["anthropic", "github-copilot", "openai-codex", "amazon-bedrock", "openai"]); - - expect(isApiKeyLoginProvider("anthropic", oauthProviderIds, builtInProviderIds)).toBe(true); - expect(BUILT_IN_PROVIDER_DISPLAY_NAMES.anthropic).toBe("Anthropic"); - expect(isApiKeyLoginProvider("openai", oauthProviderIds, builtInProviderIds)).toBe(true); - expect(isApiKeyLoginProvider("github-copilot", oauthProviderIds, builtInProviderIds)).toBe(true); - expect(isApiKeyLoginProvider("openai-codex", oauthProviderIds, builtInProviderIds)).toBe(false); - expect(isApiKeyLoginProvider("amazon-bedrock", oauthProviderIds, builtInProviderIds)).toBe(true); - expect(isApiKeyLoginProvider("custom-oauth", oauthProviderIds, builtInProviderIds)).toBe(false); - expect(isApiKeyLoginProvider("custom-api", oauthProviderIds, builtInProviderIds)).toBe(true); + it("offers a stored API key for logout even when the provider advertises OAuth only", () => { + const getLogoutProviderOptions = ( + InteractiveMode as unknown as { + prototype: { getLogoutProviderOptions(this: object): Array<{ id: string; name: string; authType: string }> }; + } + ).prototype.getLogoutProviderOptions; + const fakeThis = { + session: { + modelRuntime: { + getProviders: () => [{ id: "oauth-only", name: "OAuth Only", auth: { oauth: {} } }], + getOAuthProviderMetadata: () => [{ id: "oauth-only", name: "OAuth Only" }], + getProviderAuthStatus: () => ({ configured: true, source: "stored" }), + getStoredCredentialType: () => "api_key", + isUsingOAuth: () => false, + }, + }, + }; + + expect(getLogoutProviderOptions.call(fakeThis)).toEqual([ + { id: "oauth-only", name: "OAuth Only", authType: "api_key" }, + ]); }); - it("shows stored OAuth auth distinctly in the API key selector", () => { - const authStorage = AuthStorage.inMemory({ - anthropic: { - type: "oauth", - access: "access-token", - refresh: "refresh-token", - expires: Date.now() + 60_000, + it("labels an engine-published provider by its stored API-key credential", () => { + const getLogoutProviderOptions = ( + InteractiveMode as unknown as { + prototype: { getLogoutProviderOptions(this: object): Array<{ id: string; name: string; authType: string }> }; + } + ).prototype.getLogoutProviderOptions; + const fakeThis = { + session: { + modelRuntime: { + getProviders: () => [], + getOAuthProviderMetadata: () => [{ id: "engine-oauth", name: "Engine OAuth" }], + getProviderAuthStatus: () => ({ configured: true, source: "stored" }), + getStoredCredentialType: () => "api_key", + isUsingOAuth: () => false, + }, }, - }); + }; + + expect(getLogoutProviderOptions.call(fakeThis)).toEqual([ + { id: "engine-oauth", name: "Engine OAuth", authType: "api_key" }, + ]); + }); + + it("renders an option without compiled auth status as unconfigured", () => { const selector = new OAuthSelectorComponent( "login", - authStorage, - [{ id: "anthropic", name: "Anthropic", authType: "api_key" }], + fakeModelRuntime(), + [{ id: "google", name: "Google", authType: "api_key" }], () => {}, () => {}, ); const output = stripAnsi(selector.render(120).join("\n")); + expect(output).toContain("unconfigured"); + expect(output).not.toContain("✓ configured"); + }); + + it("shows stored OAuth auth distinctly in the API key selector", () => { + const selector = new OAuthSelectorComponent( + "login", + fakeModelRuntime({ getStoredCredentialType: () => "oauth" }), + [{ id: "anthropic", name: "Anthropic", authType: "api_key" }], + () => {}, + () => {}, + () => ({ configured: true, source: "stored", label: "OAuth" }), + ); - expect(output).toContain("Anthropic"); + const output = stripAnsi(selector.render(120).join("\n")); expect(output).toContain("subscription configured"); + expect(output).not.toContain("API key configured"); }); - it("shows environment API key auth as configured", () => { - process.env.OPENAI_API_KEY = "test-openai-key"; - const authStorage = AuthStorage.inMemory(); + it("shows a stored API key distinctly in the subscription selector", () => { const selector = new OAuthSelectorComponent( "login", - authStorage, - [{ id: "openai", name: "OpenAI", authType: "api_key" }], + fakeModelRuntime({ getStoredCredentialType: () => "api_key" }), + [{ id: "anthropic", name: "Anthropic", authType: "oauth" }], () => {}, () => {}, + () => ({ configured: true, source: "stored", label: "API key" }), ); const output = stripAnsi(selector.render(120).join("\n")); - - expect(output).toContain("OpenAI"); - expect(output).toContain("✓ env: OPENAI_API_KEY"); - expect(output).not.toContain("unconfigured"); + expect(output).toContain("API key configured"); + expect(output).not.toContain("subscription configured"); }); - it("shows custom provider environment API key auth from status resolver", () => { - const authStorage = AuthStorage.inMemory(); + it("shows environment API key auth as configured", () => { const selector = new OAuthSelectorComponent( "login", - authStorage, - [{ id: "ollama", name: "ollama", authType: "api_key" }], + fakeModelRuntime(), + [{ id: "openai", name: "OpenAI", authType: "api_key" }], () => {}, () => {}, - () => ({ configured: true, source: "environment", label: "OLLAMA_API_KEY" }), + () => ({ configured: true, source: "environment", label: "OPENAI_API_KEY" }), ); const output = stripAnsi(selector.render(120).join("\n")); - - expect(output).toContain("ollama"); - expect(output).toContain("✓ env: OLLAMA_API_KEY"); + expect(output).toContain("✓ env: OPENAI_API_KEY"); expect(output).not.toContain("unconfigured"); }); it("shows models.json API key auth as configured", () => { - const authStorage = AuthStorage.inMemory(); const selector = new OAuthSelectorComponent( "login", - authStorage, + fakeModelRuntime(), [{ id: "local-proxy", name: "local-proxy", authType: "api_key" }], () => {}, () => {}, () => ({ configured: true, source: "models_json_key" }), ); - const output = stripAnsi(selector.render(120).join("\n")); - - expect(output).toContain("local-proxy"); - expect(output).toContain("✓ key in models.json"); - expect(output).not.toContain("unconfigured"); + expect(stripAnsi(selector.render(120).join("\n"))).toContain("✓ key in models.json"); }); it("shows models.json command auth as configured", () => { - const authStorage = AuthStorage.inMemory(); const selector = new OAuthSelectorComponent( "login", - authStorage, + fakeModelRuntime(), [{ id: "op-proxy", name: "op-proxy", authType: "api_key" }], () => {}, () => {}, () => ({ configured: true, source: "models_json_command" }), ); - const output = stripAnsi(selector.render(120).join("\n")); - - expect(output).toContain("op-proxy"); - expect(output).toContain("✓ command in models.json"); - expect(output).not.toContain("unconfigured"); + expect(stripAnsi(selector.render(120).join("\n"))).toContain("✓ command in models.json"); }); }); diff --git a/packages/coding-agent/test/package-command-model-refresh.test.ts b/packages/coding-agent/test/package-command-model-refresh.test.ts index a0ec2db89..7728ca6a6 100644 --- a/packages/coding-agent/test/package-command-model-refresh.test.ts +++ b/packages/coding-agent/test/package-command-model-refresh.test.ts @@ -60,7 +60,7 @@ describe("atomic update --models", () => { it("loads and force-refreshes Atomic extension providers", async () => { let factoryCalls = 0; let refreshCalls = 0; - let observedOptions: { allowNetwork: boolean; force?: boolean } | undefined; + const observedOptions: Array<{ allowNetwork: boolean; force?: boolean }> = []; const log = vi.spyOn(console, "log").mockImplementation(() => {}); await expect(handlePackageCommand(["update", "--models"], { @@ -68,9 +68,10 @@ describe("atomic update --models", () => { (pi) => { factoryCalls += 1; pi.registerProvider("extension-catalog", { + apiKey: "test-key", refreshModels: async ({ allowNetwork, force }) => { refreshCalls += 1; - observedOptions = { allowNetwork, force }; + observedOptions.push({ allowNetwork, force }); return []; }, }); @@ -79,8 +80,8 @@ describe("atomic update --models", () => { })).resolves.toBe(true); expect(factoryCalls).toBeGreaterThanOrEqual(1); - expect(refreshCalls).toBe(1); - expect(observedOptions).toEqual({ allowNetwork: true, force: true }); + expect(refreshCalls).toBeGreaterThanOrEqual(1); + expect(observedOptions).toContainEqual({ allowNetwork: true, force: true }); expect(log.mock.calls.flat().join("\n")).toContain("Model catalogs refreshed"); }); diff --git a/packages/coding-agent/test/pi-0.82.1-auth.test.ts b/packages/coding-agent/test/pi-0.82.1-auth.test.ts index ed6dc6d1d..74e9db1f2 100644 --- a/packages/coding-agent/test/pi-0.82.1-auth.test.ts +++ b/packages/coding-agent/test/pi-0.82.1-auth.test.ts @@ -9,7 +9,7 @@ import { SettingsManager } from "../src/core/settings-manager.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { getBuiltinApiKeyLoginOptions } from "../src/modes/interactive/interactive-auth-routing.ts"; import { resolveLoginProviderReference } from "../src/modes/interactive/login-provider-options.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { createInMemoryModelRegistry, createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; function jsonResponse(body: object, status = 200): Response { return new Response(JSON.stringify(body), { @@ -33,25 +33,23 @@ afterEach(() => { }); describe("Pi 0.82.1 authentication integration", () => { - it("exposes OpenRouter and Kimi Code through provider-owned OAuth metadata", () => { - const providers = AuthStorage.inMemory().getOAuthProviders(); - expect(providers.find(({ id }) => id === "openrouter")).toMatchObject({ + it("exposes OpenRouter and Kimi Code through provider-owned OAuth metadata", async () => { + const registry = await createInMemoryModelRegistry(AuthStorage.inMemory()); + const runtime = getModelRuntime(registry); + expect(runtime.getProvider("openrouter")?.auth.oauth).toMatchObject({ name: "OpenRouter OAuth", - loginLabel: "Sign in with OpenRouter", }); - expect(providers.find(({ id }) => id === "kimi-coding")).toMatchObject({ + expect(runtime.getProvider("kimi-coding")?.auth.oauth).toMatchObject({ name: "Kimi Code (subscription)", - loginLabel: "Sign in with Kimi Code", }); }); - it("routes direct /login references for both subscription providers", () => { - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); - const oauthOptions = registry.authStorage.getOAuthProviders().map(({ id }) => ({ - id, - name: registry.getProviderDisplayName(id), - authType: "oauth" as const, - })); + it("routes direct /login references for both subscription providers", async () => { + const registry = await createInMemoryModelRegistry(AuthStorage.inMemory()); + const runtime = getModelRuntime(registry); + const oauthOptions = runtime.getProviders() + .filter((provider) => provider.auth.oauth) + .map((provider) => ({ id: provider.id, name: provider.name ?? provider.id, authType: "oauth" as const })); const options = [ ...oauthOptions, ...getBuiltinApiKeyLoginOptions((id) => registry.getProviderDisplayName(id)), @@ -75,7 +73,7 @@ describe("Pi 0.82.1 authentication integration", () => { process.env.ANTHROPIC_AUTH_TOKEN = "gateway-token"; delete process.env.ANTHROPIC_API_KEY; delete process.env.ANTHROPIC_OAUTH_TOKEN; - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const registry = await createInMemoryModelRegistry(AuthStorage.inMemory()); const model = registry.getAll().find(({ provider }) => provider === "anthropic"); expect(model).toBeDefined(); @@ -95,31 +93,33 @@ describe("Pi 0.82.1 authentication integration", () => { delete process.env.ANTHROPIC_API_KEY; delete process.env.ANTHROPIC_OAUTH_TOKEN; const storage = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(storage); + const registry = await createInMemoryModelRegistry(storage); + const runtime = getModelRuntime(registry); - expect(storage.hasAuth("anthropic")).toBe(true); - expect(storage.getAuthStatus("anthropic")).toEqual({ + expect(runtime.hasConfiguredAuth("anthropic")).toBe(true); + expect(runtime.getProviderAuthStatus("anthropic")).toEqual({ configured: true, source: "environment", label: "ANTHROPIC_AUTH_TOKEN", }); - expect(registry.getAvailable().some(({ provider }) => provider === "anthropic")).toBe(true); - await expect(registry.checkAuth("anthropic")).resolves.toEqual({ + expect(runtime.getAvailableSnapshot().some(({ provider }) => provider === "anthropic")).toBe(true); + await expect(runtime.checkAuth("anthropic")).resolves.toEqual({ source: "ANTHROPIC_AUTH_TOKEN", type: "api_key", }); }); - it("does not treat known but empty Anthropic environment variables as configured", () => { + it("does not treat known but empty Anthropic environment variables as configured", async () => { process.env.ANTHROPIC_AUTH_TOKEN = ""; process.env.ANTHROPIC_OAUTH_TOKEN = ""; process.env.ANTHROPIC_API_KEY = ""; const storage = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(storage); + const registry = await createInMemoryModelRegistry(storage); + const runtime = getModelRuntime(registry); - expect(storage.hasAuth("anthropic")).toBe(false); - expect(storage.getAuthStatus("anthropic")).toEqual({ configured: false }); - expect(registry.getAvailable().some(({ provider }) => provider === "anthropic")).toBe(false); + expect(runtime.hasConfiguredAuth("anthropic")).toBe(false); + expect(runtime.getProviderAuthStatus("anthropic")).toEqual({ configured: false }); + expect(runtime.getAvailableSnapshot().some(({ provider }) => provider === "anthropic")).toBe(false); }); it("selects, restores, and cycles Anthropic models with bearer-only auth", async () => { @@ -128,7 +128,8 @@ describe("Pi 0.82.1 authentication integration", () => { delete process.env.ANTHROPIC_OAUTH_TOKEN; const directory = mkdtempSync(join(tmpdir(), "atomic-anthropic-bearer-models-")); const storage = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(storage); + const registry = await createInMemoryModelRegistry(storage); + const runtime = getModelRuntime(registry); const models = registry.getAll().filter(({ provider }) => provider === "anthropic").slice(0, 2); expect(models).toHaveLength(2); const [first, second] = models; @@ -139,7 +140,7 @@ describe("Pi 0.82.1 authentication integration", () => { cliModel: first!.id, scopedModels: [], isContinuing: false, - modelRegistry: registry, + modelRuntime: runtime, }); expect(direct.model).toBe(first); @@ -148,10 +149,10 @@ describe("Pi 0.82.1 authentication integration", () => { isContinuing: false, defaultProvider: first!.provider, defaultModelId: first!.id, - modelRegistry: registry, + modelRuntime: runtime, }); expect(configuredDefault.model).toBe(first); - await expect(restoreModelFromSession(first!.provider, first!.id, undefined, false, registry)).resolves.toEqual({ + await expect(restoreModelFromSession(first!.provider, first!.id, undefined, false, runtime)).resolves.toEqual({ model: first, fallbackMessage: undefined, }); @@ -160,7 +161,7 @@ describe("Pi 0.82.1 authentication integration", () => { cwd: directory, agentDir: directory, authStorage: storage, - modelRegistry: registry, + modelRuntime: runtime, model: first, scopedModels: models.map((model) => ({ model })), settingsManager: SettingsManager.inMemory(), @@ -181,17 +182,18 @@ describe("Pi 0.82.1 authentication integration", () => { const storage = AuthStorage.inMemory({ "kimi-coding": { type: "oauth", access: "old", refresh: "refresh-old", expires: 0 }, }); - const registry = ModelRegistry.inMemory(storage); + const registry = await createInMemoryModelRegistry(storage); + const runtime = getModelRuntime(registry); vi.spyOn(globalThis, "fetch").mockResolvedValue(jsonResponse({ access_token: "access-new", refresh_token: "refresh-new", expires_in: 3600, })); - await expect(registry.getAuth("kimi-coding")).resolves.toMatchObject({ + await expect(runtime.getAuth("kimi-coding")).resolves.toMatchObject({ auth: { headers: { Authorization: "Bearer access-new" } }, }); - expect(storage.get("kimi-coding")).toMatchObject({ + expect(await storage.read("kimi-coding")).toMatchObject({ type: "oauth", access: "access-new", refresh: "refresh-new", @@ -202,7 +204,7 @@ describe("Pi 0.82.1 authentication integration", () => { const storage = AuthStorage.inMemory({ "kimi-coding": { type: "oauth", access: "old", refresh: "revoked", expires: 0 }, }); - const registry = ModelRegistry.inMemory(storage); + const registry = await createInMemoryModelRegistry(storage); vi.spyOn(globalThis, "fetch").mockResolvedValue(jsonResponse({ error: "invalid_grant", error_description: "credential revoked by gateway", @@ -231,14 +233,15 @@ describe("Pi 0.82.1 authentication integration", () => { const storage = AuthStorage.inMemory({ "corp-radius": { type: "oauth", access: "old", refresh: "refresh", expires: 0 }, }); - const registry = ModelRegistry.create(storage, join(directory, "models.json")); + const registry = await createModelRegistry(storage, join(directory, "models.json")); + const runtime = getModelRuntime(registry); const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValue(jsonResponse({ access_token: "new-access", refresh_token: "new-refresh", expires_in: 3600, })); - await expect(registry.getAuth("corp-radius")).resolves.toMatchObject({ + await expect(runtime.getAuth("corp-radius")).resolves.toMatchObject({ auth: { apiKey: "new-access" }, }); expect(fetchSpy).toHaveBeenCalledWith( diff --git a/packages/coding-agent/test/provider-context-usage.test.ts b/packages/coding-agent/test/provider-context-usage.test.ts index 9f31d8a7d..3541e6b00 100644 --- a/packages/coding-agent/test/provider-context-usage.test.ts +++ b/packages/coding-agent/test/provider-context-usage.test.ts @@ -12,6 +12,7 @@ import { getLatestCompactionBoundaryEntry, SessionManager } from "../src/core/se import { SettingsManager } from "../src/core/settings-manager.ts"; import { createTestResourceLoader } from "./utilities.ts"; import { appendTestCompaction } from "./verbatim-compaction-test-helpers.ts"; +import { createInMemoryModelRegistry, createModelRegistry } from "./model-runtime-test-utils.ts"; const model = getModel("anthropic", "claude-sonnet-4-5")!; @@ -84,12 +85,12 @@ describe("provider-bound context usage", () => { sessionManager.appendMessage(createUserMessage("continue", Date.now() + 1)); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(model.provider, "test-key"); + await authStorage.modify(model.provider, async () => ({ type: "api_key", key: "test-key" })); const { session } = await createAgentSession({ cwd, model, authStorage, - modelRegistry: ModelRegistry.inMemory(authStorage), + modelRegistry: await createInMemoryModelRegistry(authStorage), settingsManager: SettingsManager.inMemory(), sessionManager, resourceLoader: createTestResourceLoader(), diff --git a/packages/coding-agent/test/remote-catalog-provider.test.ts b/packages/coding-agent/test/remote-catalog-provider.test.ts index e832b5163..db00f188c 100644 --- a/packages/coding-agent/test/remote-catalog-provider.test.ts +++ b/packages/coding-agent/test/remote-catalog-provider.test.ts @@ -1,4 +1,10 @@ -import { createProvider, InMemoryModelsStore, type Model } from "@earendil-works/pi-ai"; +import { + createProvider, + InMemoryModelsStore, + type Model, + type ModelsStoreEntry, + type ProviderModelsStore, +} from "@earendil-works/pi-ai"; import { afterEach, describe, expect, it, vi } from "vitest"; import { VERSION } from "../src/config.ts"; import { withRemoteCatalog } from "../src/core/remote-catalog-provider.ts"; @@ -18,10 +24,30 @@ function model(id: string): Model<"openai-completions"> { }; } -function providerStore(store: InMemoryModelsStore) { +function testProvider(localGeneratedAt?: number) { + return withRemoteCatalog( + createProvider({ + id: "test-provider", + auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, + models: [model("static")], + api: { + stream: () => { + throw new Error("not used"); + }, + streamSimple: () => { + throw new Error("not used"); + }, + }, + }), + "https://pi.dev", + localGeneratedAt, + ); +} + +function scopedStore(store: InMemoryModelsStore): ProviderModelsStore { return { read: () => store.read("test-provider"), - write: (entry: Parameters[1]) => store.write("test-provider", entry), + write: (entry: ModelsStoreEntry) => store.write("test-provider", entry), delete: () => store.delete("test-provider"), }; } @@ -29,31 +55,20 @@ function providerStore(store: InMemoryModelsStore) { afterEach(() => vi.restoreAllMocks()); describe("remote catalog provider", () => { - it("persists keyed catalogs, sends version headers, observes TTL, and supports forced refresh", async () => { + it("parses keyed catalogs, sends version headers, observes the refresh TTL, and supports forced refreshes", async () => { const fetchSpy = vi.spyOn(globalThis, "fetch").mockImplementation( - async () => new Response(JSON.stringify({ dynamic: model("dynamic") }), { - status: 200, - headers: { "content-type": "application/json" }, - }), - ); - const provider = withRemoteCatalog( - createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - }), - "https://catalog.example.test", + async () => + new Response(JSON.stringify({ dynamic: model("dynamic") }), { + status: 200, + headers: { "content-type": "application/json" }, + }), ); + const provider = testProvider(); const store = new InMemoryModelsStore(); - const context = { credential: { type: "api_key" as const }, store: providerStore(store), allowNetwork: true }; - - await provider.refreshModels?.(context); - await provider.refreshModels?.(context); - await provider.refreshModels?.({ ...context, force: true }); + const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true }; + await provider.refreshModels?.(refresh); + await provider.refreshModels?.(refresh); + await provider.refreshModels?.({ ...refresh, force: true }); expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]); expect((await store.read(provider.id))?.models.map((entry) => entry.id)).toEqual(["dynamic"]); @@ -63,355 +78,135 @@ describe("remote catalog provider", () => { }); }); - it("ignores persisted catalogs that are not newer than the bundled generation timestamp", async () => { - const createWrappedProvider = () => withRemoteCatalog( - createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, + it("prefers the newer of the generated and remote catalogs", async () => { + const localGeneratedAt = Date.parse("2026-07-23T10:00:00.000Z"); + const newerHeader = new Date(localGeneratedAt + 60_000).toUTCString(); + const responses = [ + new Response(JSON.stringify({ old: model("old") }), { + headers: { "last-modified": new Date(localGeneratedAt - 60_000).toUTCString() }, }), - "https://catalog.example.test", - 200_000, - ); - for (const entry of [ - { models: [model("legacy")], checkedAt: Date.now() }, - { models: [model("stale")], checkedAt: Date.now(), lastModified: 100_000 }, - { models: [model("fresh")], checkedAt: Date.now(), lastModified: 300_000 }, - ]) { - const provider = createWrappedProvider(); - const store = new InMemoryModelsStore(); - await store.write(provider.id, entry); - await provider.refreshModels?.({ credential: { type: "api_key" }, store: providerStore(store), allowNetwork: false }); - const expected = entry.lastModified === 300_000 ? ["static", "fresh"] : ["static"]; - expect(provider.getModels().map((candidate) => candidate.id)).toEqual(expected); - } - }); - - it("retains cached models on errors and treats 501 routes as unavailable overlays", async () => { - const fetchSpy = vi.spyOn(globalThis, "fetch") - .mockResolvedValueOnce(new Response(JSON.stringify([model("cached")]), { status: 200 })) - .mockResolvedValueOnce(new Response("failure", { status: 503 })) - .mockResolvedValueOnce(new Response("not implemented", { status: 501 })); - const provider = withRemoteCatalog( - createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, + new Response(JSON.stringify({ newer: model("newer") }), { + headers: { "last-modified": newerHeader }, }), - "https://catalog.example.test", - ); + ]; + vi.spyOn(globalThis, "fetch").mockImplementation(async () => responses.shift() as Response); + const provider = testProvider(localGeneratedAt); const store = new InMemoryModelsStore(); - const context = { credential: { type: "api_key" as const }, store: providerStore(store), allowNetwork: true }; - - await provider.refreshModels?.(context); - await store.write(provider.id, { models: [model("cached")], checkedAt: 0 }); - await expect(provider.refreshModels?.(context)).rejects.toThrow("503"); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "cached"]); - await expect(provider.refreshModels?.(context)).resolves.toBeUndefined(); - expect(fetchSpy).toHaveBeenCalledTimes(3); - }); + const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true }; - it("does not publish refreshed models when persistence fails", async () => { - vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response(JSON.stringify([model("new")]), { status: 200 })); - const provider = withRemoteCatalog( - createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - }), - "https://catalog.example.test", - ); - const store = { - read: async () => ({ models: [model("stale")], checkedAt: 0 }), - write: async () => { throw new Error("disk full"); }, - delete: async () => {}, - }; + await provider.refreshModels?.(refresh); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static"]); - await expect(provider.refreshModels?.({ - credential: { type: "api_key" }, - store, - allowNetwork: true, - })).rejects.toThrow("disk full"); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "stale"]); + await provider.refreshModels?.({ ...refresh, force: true }); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "newer"]); + expect(await store.read(provider.id)).toMatchObject({ lastModified: Date.parse(newerHeader) }); }); - it("allows retry after an aborted request ignores its signal", async () => { - const fetchSpy = vi.spyOn(globalThis, "fetch") - .mockImplementationOnce(async () => new Promise(() => {})) - .mockResolvedValueOnce(new Response("not found", { status: 404 })); - const provider = withRemoteCatalog( - createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, + it("revalidates a stored catalog with its etag and keeps the overlay on 304", async () => { + const responses = [ + new Response(JSON.stringify({ dynamic: model("dynamic") }), { + headers: { "content-type": "application/json", etag: '"catalog-1"' }, }), - "https://catalog.example.test", - ); + new Response(null, { status: 304, headers: { etag: '"catalog-1"' } }), + ]; + const fetchSpy = vi.spyOn(globalThis, "fetch").mockImplementation(async () => responses.shift() as Response); + const provider = testProvider(); const store = new InMemoryModelsStore(); - const controller = new AbortController(); - const context = { credential: { type: "api_key" as const }, store: providerStore(store), allowNetwork: true }; - const first = provider.refreshModels?.({ ...context, signal: controller.signal }); - await vi.waitFor(() => expect(fetchSpy).toHaveBeenCalledTimes(1)); + const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true }; - controller.abort(); - await expect(first).resolves.toBeUndefined(); - await expect(provider.refreshModels?.({ ...context, force: true })).resolves.toBeUndefined(); - expect(fetchSpy).toHaveBeenCalledTimes(2); - }); + await provider.refreshModels?.(refresh); + expect(fetchSpy.mock.calls[0]?.[1]?.headers).not.toHaveProperty("if-none-match"); + expect(await store.read(provider.id)).toMatchObject({ etag: '"catalog-1"' }); - it("retries a surviving same-strength caller when the refresh owner later aborts", async () => { - const fetchSpy = vi.spyOn(globalThis, "fetch") - .mockImplementationOnce(async (_url, init) => new Promise((_resolve, reject) => { - init?.signal?.addEventListener("abort", () => reject(new DOMException("Aborted", "AbortError")), { once: true }); - })) - .mockResolvedValueOnce(new Response(JSON.stringify([model("fresh")]), { status: 200 })); - const provider = withRemoteCatalog(createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - }), "https://catalog.example.test"); - const store = providerStore(new InMemoryModelsStore()); - const ownerController = new AbortController(); - const context = { credential: { type: "api_key" as const }, store, allowNetwork: true, force: true }; - const owner = provider.refreshModels?.({ ...context, signal: ownerController.signal }); - await vi.waitFor(() => expect(fetchSpy).toHaveBeenCalledTimes(1)); - const survivor = provider.refreshModels?.(context); - - ownerController.abort(); - await expect(owner).resolves.toBeUndefined(); - await expect(survivor).resolves.toBeUndefined(); - expect(fetchSpy).toHaveBeenCalledTimes(2); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "fresh"]); - }); - - it("retries surviving callers when an aborted owner's delayed write rejects", async () => { - let releaseFirstWrite!: () => void; - const firstWriteGate = new Promise((resolve) => { releaseFirstWrite = resolve; }); - let writeCount = 0; - let persisted: string[] = []; - const store = { - read: async () => undefined, - write: async (entry: { models: readonly Model<"openai-completions">[] }) => { - writeCount += 1; - if (writeCount === 1) { - await firstWriteGate; - throw new Error("transient write"); - } - persisted = entry.models.map((candidate) => candidate.id); - }, - delete: async () => {}, - }; - const fetchSpy = vi.spyOn(globalThis, "fetch") - .mockResolvedValueOnce(new Response(JSON.stringify([model("stale")]), { status: 200 })) - .mockResolvedValueOnce(new Response(JSON.stringify([model("fresh")]), { status: 200 })); - const provider = withRemoteCatalog(createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - })); - const ownerController = new AbortController(); - const context = { credential: { type: "api_key" as const }, store, allowNetwork: true, force: true }; - const owner = provider.refreshModels?.({ ...context, signal: ownerController.signal }); - const ownerFailure = expect(owner).rejects.toThrow("transient write"); - await vi.waitFor(() => expect(writeCount).toBe(1)); - const survivors = [provider.refreshModels?.(context), provider.refreshModels?.(context)]; + const checkedAt = (await store.read(provider.id))?.checkedAt; + await provider.refreshModels?.({ ...refresh, force: true }); - ownerController.abort(); - releaseFirstWrite(); - await ownerFailure; - await Promise.all(survivors); - expect(fetchSpy).toHaveBeenCalledTimes(2); - expect(writeCount).toBe(2); - expect(persisted).toEqual(["fresh"]); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "fresh"]); + expect(fetchSpy.mock.calls[1]?.[1]?.headers).toMatchObject({ "if-none-match": '"catalog-1"' }); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]); + const stored = await store.read(provider.id); + expect(stored?.models.map((entry) => entry.id)).toEqual(["dynamic"]); + expect(stored?.etag).toBe('"catalog-1"'); + expect(stored?.checkedAt).toBeGreaterThanOrEqual(checkedAt ?? 0); }); - it("prevents an aborted stale read from overwriting a newer catalog", async () => { - type StoredEntry = Awaited>; - let resolveOldRead!: (entry: StoredEntry) => void; - const oldRead = new Promise((resolve) => { resolveOldRead = resolve; }); - let readCount = 0; - const store = { - read: async () => { - readCount += 1; - return readCount === 1 - ? oldRead - : { models: [model("fresh")], checkedAt: Date.now() }; - }, - write: async () => {}, - delete: async () => {}, - }; - const provider = withRemoteCatalog(createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - })); - const controller = new AbortController(); - const context = { credential: { type: "api_key" as const }, store, allowNetwork: false }; - const staleRefresh = provider.refreshModels?.({ ...context, signal: controller.signal }); - await vi.waitFor(() => expect(readCount).toBe(1)); + it("drops a stale etag when the overlay becomes unavailable", async () => { + const responses = [ + new Response(JSON.stringify({ dynamic: model("dynamic") }), { + headers: { "content-type": "application/json", etag: '"catalog-1"' }, + }), + new Response("not implemented", { status: 501 }), + ]; + vi.spyOn(globalThis, "fetch").mockImplementation(async () => responses.shift() as Response); + const provider = testProvider(); + const store = new InMemoryModelsStore(); + const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true }; - controller.abort(); - await expect(staleRefresh).resolves.toBeUndefined(); - await provider.refreshModels?.(context); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "fresh"]); + await provider.refreshModels?.(refresh); + await provider.refreshModels?.({ ...refresh, force: true }); - resolveOldRead({ models: [model("stale")], checkedAt: 0 }); - await new Promise((resolve) => setTimeout(resolve, 0)); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "fresh"]); + expect((await store.read(provider.id))?.etag).toBeUndefined(); }); - it("keeps an aborted pending write fenced before a newer refresh persists", async () => { - type StoredEntry = NonNullable>>; - let persisted: StoredEntry | undefined; - let resolveFirstWrite!: () => void; - const firstWriteGate = new Promise((resolve) => { resolveFirstWrite = resolve; }); - let writeCount = 0; - const store = { - read: async () => persisted, - write: async (entry: StoredEntry) => { - writeCount += 1; - if (writeCount === 1) await firstWriteGate; - persisted = entry; - }, - delete: async () => { persisted = undefined; }, - }; - vi.spyOn(globalThis, "fetch") - .mockResolvedValueOnce(new Response(JSON.stringify([model("stale")]), { status: 200 })) - .mockResolvedValueOnce(new Response(JSON.stringify([model("fresh")]), { status: 200 })); - const provider = withRemoteCatalog(createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - })); - const context = { credential: { type: "api_key" as const }, store, allowNetwork: true, force: true }; - const controller = new AbortController(); - const staleRefresh = provider.refreshModels?.({ ...context, signal: controller.signal }); - await vi.waitFor(() => expect(writeCount).toBe(1)); + it("keeps the etag and overlay after a transient failure", async () => { + const responses = [ + new Response(JSON.stringify({ dynamic: model("dynamic") }), { + headers: { "content-type": "application/json", etag: '"catalog-1"' }, + }), + new Response("rate limited", { status: 429 }), + new Response(null, { status: 304, headers: { etag: '"catalog-1"' } }), + ]; + const fetchSpy = vi.spyOn(globalThis, "fetch").mockImplementation(async () => responses.shift() as Response); + const provider = testProvider(); + const store = new InMemoryModelsStore(); + const refresh = { credential: { type: "api_key" } as const, store: scopedStore(store), allowNetwork: true }; - controller.abort(); - let staleSettled = false; - void staleRefresh?.then(() => { staleSettled = true; }); - await new Promise((resolve) => setTimeout(resolve, 0)); - expect(staleSettled).toBe(false); - const overlappingRetry = provider.refreshModels?.(context); - expect(overlappingRetry).not.toBe(staleRefresh); - expect(writeCount).toBe(1); + await provider.refreshModels?.(refresh); + await expect(provider.refreshModels?.({ ...refresh, force: true })).rejects.toThrow(/429/); - resolveFirstWrite(); - await overlappingRetry; - expect(persisted?.models.map((entry) => entry.id)).toEqual(["fresh"]); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "fresh"]); + const stored = await store.read(provider.id); + expect(stored?.etag).toBe('"catalog-1"'); + expect(stored?.models.map((entry) => entry.id)).toEqual(["dynamic"]); + + await provider.refreshModels?.({ ...refresh, force: true }); + expect(fetchSpy.mock.calls[2]?.[1]?.headers).toMatchObject({ "if-none-match": '"catalog-1"' }); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "dynamic"]); }); - it("escalates a forced network refresh over an in-flight cache restore", async () => { - let resolveCacheRead!: () => void; - const cacheReadGate = new Promise((resolve) => { resolveCacheRead = resolve; }); - const backingStore = new InMemoryModelsStore(); - await backingStore.write("test-provider", { models: [model("cached")], checkedAt: Date.now() }); - let readCount = 0; - const store = { - read: async () => { - readCount += 1; - if (readCount === 1) await cacheReadGate; - return backingStore.read("test-provider"); - }, - write: (entry: Parameters[1]) => backingStore.write("test-provider", entry), - delete: () => backingStore.delete("test-provider"), - }; - const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValue( - new Response(JSON.stringify([model("fresh")]), { status: 200 }), - ); - const provider = withRemoteCatalog(createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - })); - const credential = { type: "api_key" as const }; - const cacheRestore = provider.refreshModels?.({ credential, store, allowNetwork: false }); - await vi.waitFor(() => expect(readCount).toBe(1)); - const forcedRefresh = provider.refreshModels?.({ credential, store, allowNetwork: true, force: true }); + it("treats unimplemented pi.dev catalog routes as an unavailable overlay", async () => { + vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response("not implemented", { status: 501 })); + const provider = testProvider(); + const store = new InMemoryModelsStore(); - resolveCacheRead(); - await cacheRestore; - await forcedRefresh; - expect(fetchSpy).toHaveBeenCalledOnce(); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "fresh"]); + await expect( + provider.refreshModels?.({ + credential: { type: "api_key" }, + store: scopedStore(store), + allowNetwork: true, + }), + ).resolves.toBeUndefined(); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static"]); + expect(await store.read(provider.id)).toMatchObject({ models: [], checkedAt: expect.any(Number) }); }); - it("runs a forced escalation after the weaker cache restore rejects", async () => { - let resolveCacheRead!: () => void; - const cacheReadGate = new Promise((resolve) => { resolveCacheRead = resolve; }); - let readCount = 0; - const store = { - read: async () => { - readCount += 1; - if (readCount === 1) { - await cacheReadGate; - throw new Error("transient cache read"); - } - return undefined; - }, - write: async () => {}, + it("pi publishes refreshed models before surfacing persistence failures", async () => { + vi.spyOn(globalThis, "fetch").mockResolvedValue( + new Response(JSON.stringify({ new: model("new") }), { + status: 200, + headers: { "content-type": "application/json" }, + }), + ); + const provider = testProvider(); + const store: ProviderModelsStore = { + read: async () => ({ models: [model("stale")], checkedAt: 0 }), + write: async () => { throw new Error("disk full"); }, delete: async () => {}, }; - const fetchSpy = vi.spyOn(globalThis, "fetch").mockResolvedValue( - new Response(JSON.stringify([model("fresh")]), { status: 200 }), - ); - const provider = withRemoteCatalog(createProvider({ - id: "test-provider", - auth: { apiKey: { name: "Test", resolve: async () => ({ auth: {} }) } }, - models: [model("static")], - api: { - stream: () => { throw new Error("not used"); }, - streamSimple: () => { throw new Error("not used"); }, - }, - })); - const credential = { type: "api_key" as const }; - const cacheRestore = provider.refreshModels?.({ credential, store, allowNetwork: false }); - await vi.waitFor(() => expect(readCount).toBe(1)); - const forcedRefresh = provider.refreshModels?.({ credential, store, allowNetwork: true, force: true }); - resolveCacheRead(); - await expect(cacheRestore).rejects.toThrow("transient cache read"); - await expect(forcedRefresh).resolves.toBeUndefined(); - expect(fetchSpy).toHaveBeenCalledOnce(); - expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "fresh"]); + await expect(provider.refreshModels?.({ + credential: { type: "api_key" }, + store, + allowNetwork: true, + })).rejects.toThrow("disk full"); + expect(provider.getModels().map((entry) => entry.id)).toEqual(["static", "new"]); }); }); diff --git a/packages/coding-agent/test/resource-loader-01-02.suite.ts b/packages/coding-agent/test/resource-loader-01-02.suite.ts index 696ca2b43..db23ea6bc 100644 --- a/packages/coding-agent/test/resource-loader-01-02.suite.ts +++ b/packages/coding-agent/test/resource-loader-01-02.suite.ts @@ -13,6 +13,7 @@ import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import type { Skill } from "../src/core/skills.ts"; import { createSyntheticSourceInfo } from "../src/core/source-info.ts"; +import { createInMemoryModelRegistry, createModelRegistry } from "./model-runtime-test-utils.ts"; describe("DefaultResourceLoader", () => { let tempDir: string; @@ -182,7 +183,7 @@ Project skill`, const sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage); + const modelRegistry = await createModelRegistry(authStorage); const runner = new ExtensionRunner( extensionsResult.extensions, extensionsResult.runtime, diff --git a/packages/coding-agent/test/resource-loader-05-01.suite.ts b/packages/coding-agent/test/resource-loader-05-01.suite.ts index c15da4a5a..681c3cb9d 100644 --- a/packages/coding-agent/test/resource-loader-05-01.suite.ts +++ b/packages/coding-agent/test/resource-loader-05-01.suite.ts @@ -13,6 +13,7 @@ import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import type { Skill } from "../src/core/skills.ts"; import { createSyntheticSourceInfo } from "../src/core/source-info.ts"; +import { createModelRegistry } from "./model-runtime-test-utils.ts"; describe("DefaultResourceLoader", () => { let tempDir: string; @@ -104,6 +105,7 @@ export default function(pi: ExtensionAPI) { ` import type { ExtensionAPI } from "@bastani/atomic"; import { Type } from "typebox"; + export default function(pi: ExtensionAPI) { pi.registerTool({ name: "duplicate-tool", @@ -130,7 +132,7 @@ export default function(pi: ExtensionAPI) { const sessionManager = SessionManager.inMemory(); const authStorage = AuthStorage.create(join(tempDir, "auth-explicit.json")); - const modelRegistry = ModelRegistry.create(authStorage); + const modelRegistry = await createModelRegistry(authStorage); const runner = new ExtensionRunner( extensionsResult.extensions, extensionsResult.runtime, diff --git a/packages/coding-agent/test/rpc-oauth-login.test.ts b/packages/coding-agent/test/rpc-oauth-login.test.ts index efe595433..96b570dda 100644 --- a/packages/coding-agent/test/rpc-oauth-login.test.ts +++ b/packages/coding-agent/test/rpc-oauth-login.test.ts @@ -1,155 +1,116 @@ -import { afterEach, describe, expect, it } from "vitest"; +import { afterEach, describe, expect, it, vi } from "vitest"; import { AgentSessionRuntime, type CreateAgentSessionRuntimeFactory } from "../src/core/agent-session-runtime.ts"; -import { type AtomicOAuthLoginCallbacks, resetLegacyOAuthProviders } from "../src/core/oauth-provider-bridge.ts"; -import { RemoteModelCatalog } from "../src/modes/interactive-engine/remote-model-catalog.ts"; +import type { AtomicOAuthLoginCallbacks } from "../src/core/oauth-login.ts"; import { loginIsolatedOAuthProvider } from "../src/modes/interactive-engine/isolated-auth.ts"; -import { dispatchRpcOAuthRequest } from "../src/modes/rpc/rpc-oauth-client.ts"; -import { InteractiveModeBase } from "../src/modes/interactive/interactive-mode-base.ts"; -import "../src/modes/interactive/interactive-auth-routing.ts"; -import { resolveLoginProviderReference } from "../src/modes/interactive/login-provider-options.ts"; +import { RemoteModelCatalog } from "../src/modes/interactive-engine/remote-model-catalog.ts"; import { createRpcCommandHandler } from "../src/modes/rpc/rpc-command-handler.ts"; +import { dispatchRpcOAuthRequest } from "../src/modes/rpc/rpc-oauth-client.ts"; import type { RpcPendingExtensionRequests } from "../src/modes/rpc/rpc-extension-ui.ts"; import type { RpcExtensionUIRequest } from "../src/modes/rpc/rpc-types.ts"; import { createHarness, type Harness } from "./suite/harness.ts"; -const createRuntime = (async () => { throw new Error("not used"); }) as CreateAgentSessionRuntimeFactory; +const createRuntime = (async () => { + throw new Error("not used"); +}) as CreateAgentSessionRuntimeFactory; const harnesses: Harness[] = []; + afterEach(() => { - resetLegacyOAuthProviders(); + vi.restoreAllMocks(); while (harnesses.length) harnesses.pop()?.cleanup(); }); async function createRuntimeHarness() { const harness = await createHarness({ withConfiguredAuth: false }); harnesses.push(harness); + harness.session.modelRuntime.registerProvider("corp-oauth", { + baseUrl: "https://provider.test/v1", + api: "openai-completions", + oauth: { + name: "Corp OAuth", + login: async () => ({ access: "new-secret", refresh: "new-refresh", expires: Date.now() + 60_000 }), + refreshToken: async (credential) => credential, + getApiKey: (credential) => credential.access, + }, + models: [{ + id: "corp-model", + name: "Corp Model", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128000, + maxTokens: 4096, + }], + }); + await harness.authStorage.modify("corp-oauth", async () => ({ type: "api_key", key: "previous" })); const runtime = new AgentSessionRuntime( harness.session, { cwd: harness.tempDir, agentDir: harness.tempDir } as never, createRuntime, ); - return { harness, runtime }; + const handler = createRpcCommandHandler({ + runtimeHost: runtime, + getSession: () => harness.session, + rebindSession: async () => {}, + pendingExtensionRequests: new Map(), + output: () => {}, + }); + return { harness, handler }; } -describe("isolated engine OAuth", () => { - it("transports JSON-safe descriptors for engine-only extension OAuth", async () => { - const { harness, runtime } = await createRuntimeHarness(); - harness.session.modelRegistry.registerProvider("corp-oauth", { - oauth: { - name: "Corp OAuth", - usesCallbackServer: true, - login: async () => ({ access: "secret", refresh: "refresh", expires: Date.now() + 60_000 }), - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }, - }); - harness.session.modelRegistry.registerProvider("openrouter", { - oauth: { - name: "Latest OpenRouter Override", - login: async () => ({ access: "override", refresh: "refresh", expires: Date.now() + 60_000 }), - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }, - }); - const handler = createRpcCommandHandler({ - runtimeHost: runtime, getSession: () => harness.session, rebindSession: async () => {}, output: () => {}, - }); +async function expectSuccessfulLoginAndRetainedCredential( + refresh: () => Promise | Promise<{ aborted: boolean; errors: Map }>, +) { + const { harness, handler } = await createRuntimeHarness(); + vi.spyOn(harness.session.modelRuntime, "refresh").mockImplementation(refresh); + const response = await handler({ + id: "login", + type: "login_provider", + provider: "corp-oauth", + authType: "oauth", + }); + expect(response).toMatchObject({ success: true, data: { provider: "corp-oauth", cancelled: false } }); + expect(await harness.authStorage.read("corp-oauth")).toMatchObject({ type: "oauth", access: "new-secret" }); + expect(vi.mocked(harness.session.modelRuntime.refresh).mock.calls[0]?.[0]).not.toHaveProperty("signal"); +} + +describe("RPC OAuth descriptors", () => { + it("serializes provider metadata without OAuth secrets or function-valued fields", async () => { + const { handler } = await createRuntimeHarness(); + const response = await handler({ id: "models", type: "get_available_models" }); - const response = await handler({ id: "catalog", type: "get_available_models" }); - expect(response).toMatchObject({ success: true, data: { oauthProviders: expect.arrayContaining([ - { id: "corp-oauth", name: "Corp OAuth", usesCallbackServer: true }, - { id: "openrouter", name: "Latest OpenRouter Override" }, - expect.objectContaining({ id: "kimi-coding", name: "Kimi Code (subscription)" }), - ]) } }); + expect(response).toMatchObject({ + success: true, + data: { oauthProviders: expect.arrayContaining([{ id: "corp-oauth", name: "Corp OAuth" }]) }, + }); const serialized = JSON.stringify(response); - expect(serialized).not.toContain("secret"); + expect(serialized).not.toContain("new-secret"); expect(serialized).not.toContain("refreshToken"); expect(serialized).not.toContain("getApiKey"); }); +}); - it("exposes transported OAuth descriptors through the frontend auth surface", async () => { - const { harness } = await createRuntimeHarness(); +describe("isolated OAuth frontend transport", () => { + it("exposes transported OAuth descriptors through the frontend runtime", async () => { + const harness = await createHarness({ withConfiguredAuth: false }); + harnesses.push(harness); const catalog = new RemoteModelCatalog({} as never); catalog.apply({ - models: [], scopedModels: [], customAuthProviders: [], + models: [], + scopedModels: [], + customAuthProviders: [], oauthProviders: [{ id: "corp-oauth", name: "Corp OAuth", loginLabel: "Sign in", usesCallbackServer: true }], }); - resetLegacyOAuthProviders(); catalog.patch(harness.session); - expect(harness.session.modelRegistry.authStorage.getOAuthProviders()).toEqual([ + expect(harness.session.modelRuntime.getOAuthProviderMetadata()).toEqual([ { id: "corp-oauth", name: "Corp OAuth", loginLabel: "Sign in", usesCallbackServer: true }, ]); - expect(harness.session.modelRegistry.getProviderDisplayName("corp-oauth")).toBe("Corp OAuth"); - const getOptions = InteractiveModeBase.prototype.getLoginProviderOptions as ( - this: { session: typeof harness.session }, authType?: "oauth" | "api_key", - ) => Array<{ id: string; name: string; authType: "oauth" | "api_key" }>; - const options = getOptions.call({ session: harness.session }); - expect(options).toContainEqual({ id: "corp-oauth", name: "Corp OAuth", authType: "oauth" }); - expect(resolveLoginProviderReference(options, "corp-oauth")).toMatchObject({ - kind: "direct", option: { id: "corp-oauth", authType: "oauth" }, - }); - catalog.apply({ - models: [], scopedModels: [], customAuthProviders: [], - oauthProviders: [{ id: "corp-oauth", name: "Renamed Corp OAuth" }], - }); - expect(harness.session.modelRegistry.getProviderDisplayName("corp-oauth")).toBe("Renamed Corp OAuth"); - }); - - it("runs custom OAuth in the engine, round-trips callbacks, and never returns tokens", async () => { - const { harness, runtime } = await createRuntimeHarness(); - const observed: string[] = []; - harness.session.modelRegistry.registerProvider("corp-oauth", { - oauth: { - name: "Corp OAuth", - login: async (baseCallbacks) => { - const callbacks = baseCallbacks as AtomicOAuthLoginCallbacks; - callbacks.onAuth({ url: "https://login.invalid", instructions: "open" }); - callbacks.onDeviceCode({ userCode: "ABCD", verificationUri: "https://device.invalid" }); - callbacks.onProgress?.("waiting"); - callbacks.onInfo?.("notice", [{ label: "Docs", url: "https://docs.invalid" }]); - const prompt = await callbacks.onPrompt({ message: "Tenant", placeholder: "acme" }); - const selected = await callbacks.onSelect({ message: "Account", options: [{ id: "one", label: "One" }] }); - const manual = await callbacks.onManualCodeInput?.(); - return { access: `engine-secret-${prompt}-${selected}-${manual}`, refresh: "engine-refresh", expires: Date.now() + 60_000 }; - }, - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }, - }); - const pending: RpcPendingExtensionRequests = new Map(); - const output = (frame: object) => { - if (!("type" in frame) || frame.type !== "extension_ui_request") return; - const request = frame as RpcExtensionUIRequest; - observed.push(request.method); - const record = pending.get(request.id); - if (!record) return; - const value = request.method === "oauth_prompt" ? "acme" - : request.method === "oauth_select" ? "one" - : request.method === "oauth_manual_code" ? "manual" : undefined; - if (value !== undefined) queueMicrotask(() => record.resolve({ type: "extension_ui_response", id: request.id, value })); - }; - const handler = createRpcCommandHandler({ - runtimeHost: runtime, - getSession: () => harness.session, - rebindSession: async () => {}, - output, - pendingExtensionRequests: pending, - }); - - const response = await handler({ id: "login", type: "login_provider", provider: "corp-oauth", authType: "oauth" }); - expect(observed).toEqual([ - "oauth_auth", "oauth_device_code", "oauth_progress", "oauth_info", "oauth_prompt", "oauth_select", "oauth_manual_code", - ]); - expect(harness.authStorage.get("corp-oauth")).toMatchObject({ type: "oauth", access: "engine-secret-acme-one-manual" }); - expect(response).toMatchObject({ success: true, data: { provider: "corp-oauth", cancelled: false } }); - expect(JSON.stringify(response)).not.toContain("engine-secret"); - expect(JSON.stringify(response)).not.toContain("engine-refresh"); }); it("dispatches correlated OAuth callbacks in the frontend", async () => { const calls: string[] = []; const responses: object[] = []; - const frontendCallbacks: AtomicOAuthLoginCallbacks = { + const callbacks: AtomicOAuthLoginCallbacks = { onAuth: () => calls.push("auth"), onDeviceCode: () => calls.push("device"), onProgress: () => calls.push("progress"), @@ -170,9 +131,8 @@ describe("isolated engine OAuth", () => { { type: "extension_ui_request", id: "7", method: "oauth_manual_code", provider: "corp", loginId: "login-a" }, { type: "extension_ui_request", id: "8", method: "oauth_manual_code_cancel", provider: "corp", loginId: "login-a" }, ] as const; - await dispatchRpcOAuthRequest("corp", "login-a", frontendCallbacks, { ...requests[0], loginId: "stale" }, respond); - expect(calls).toEqual([]); - for (const request of requests) await dispatchRpcOAuthRequest("corp", "login-a", frontendCallbacks, request, respond); + await dispatchRpcOAuthRequest("corp", "login-a", callbacks, { ...requests[0], loginId: "stale" }, respond); + for (const request of requests) await dispatchRpcOAuthRequest("corp", "login-a", callbacks, request, respond); expect(calls).toEqual(["auth", "device", "progress", "info", "manual-cancel"]); expect(responses).toEqual([ @@ -182,82 +142,153 @@ describe("isolated engine OAuth", () => { ]); }); - it("uses engine-owned acquisition and applies the catalog only after success", async () => { - let reloads = 0; - let applied = 0; - let acquired = 0; - const remoteCatalog = { models: [], scopedModels: [], customAuthProviders: [], oauthProviders: [] }; - const client = { - onExtensionUIRequest: () => () => {}, - respondExtensionUI: async () => {}, - cancelLoginProvider: async () => {}, - requestInternal: async (command: { provider: string }) => { - acquired += 1; - expect(command.provider).toBe("corp"); - return { provider: command.provider, cancelled: false as const, ...remoteCatalog }; - }, - }; - await loginIsolatedOAuthProvider( - { modelRegistry: { authStorage: { reload: () => { reloads += 1; } } } } as never, - client as never, - { apply: () => { applied += 1; } } as never, - "corp", - { onAuth() {}, onDeviceCode() {}, onPrompt: async () => "", onSelect: async () => undefined }, - ); - expect({ acquired, reloads, applied }).toEqual({ acquired: 1, reloads: 1, applied: 1 }); - }); - - it("cancels only the active engine OAuth transaction without overwriting prior credentials", async () => { - const { harness, runtime } = await createRuntimeHarness(); - harness.authStorage.set("corp-oauth", { type: "api_key", key: "previous" }); - harness.session.modelRegistry.registerProvider("corp-oauth", { + it("runs custom OAuth in the engine, round-trips callbacks, and never returns tokens", async () => { + const harness = await createHarness({ withConfiguredAuth: false }); + harnesses.push(harness); + const observed: string[] = []; + harness.session.modelRuntime.registerProvider("callback-oauth", { + baseUrl: "https://callback.test/v1", + api: "openai-completions", oauth: { - name: "Corp OAuth", + name: "Callback OAuth", login: async (callbacks) => { - await callbacks.onPrompt({ message: "Block" }); - return { access: "should-not-save", refresh: "refresh", expires: Date.now() + 60_000 }; + callbacks.onAuth({ url: "https://login.invalid", instructions: "open" }); + callbacks.onDeviceCode({ userCode: "ABCD", verificationUri: "https://device.invalid" }); + callbacks.onProgress?.("waiting"); + callbacks.onInfo?.("notice", [{ label: "Docs", url: "https://docs.invalid" }]); + const prompt = await callbacks.onPrompt({ message: "Tenant", placeholder: "acme" }); + const selected = await callbacks.onSelect({ message: "Account", options: [{ id: "one", label: "One" }] }); + const manual = await callbacks.onManualCodeInput?.(); + return { access: `engine-secret-${prompt}-${selected}-${manual}`, refresh: "engine-refresh", expires: Date.now() + 60_000 }; }, refreshToken: async (credential) => credential, getApiKey: (credential) => credential.access, }, + models: [], }); + const runtime = new AgentSessionRuntime( + harness.session, + { cwd: harness.tempDir, agentDir: harness.tempDir } as never, + createRuntime, + ); const pending: RpcPendingExtensionRequests = new Map(); - let promptReady!: () => void; - const prompt = new Promise((resolve) => { promptReady = resolve; }); const handler = createRpcCommandHandler({ runtimeHost: runtime, getSession: () => harness.session, rebindSession: async () => {}, pendingExtensionRequests: pending, output: (frame) => { - if ("method" in frame && frame.method === "oauth_prompt") promptReady(); + if (!("method" in frame) || !frame.method.startsWith("oauth_")) return; + observed.push(frame.method); + const record = pending.get(frame.id); + if (!record) return; + const value = frame.method === "oauth_prompt" ? "acme" + : frame.method === "oauth_select" ? "one" + : frame.method === "oauth_manual_code" ? "manual" : undefined; + if (value !== undefined) queueMicrotask(() => record.resolve({ type: "extension_ui_response", id: frame.id, value })); }, }); - let resolved = false; - const login = handler({ - id: "login", type: "login_provider", provider: "corp-oauth", authType: "oauth", loginId: "active-login", - }).then((response) => { resolved = true; return response; }); - await prompt; - await handler({ - id: "stale-cancel", type: "cancel_login_provider", provider: "corp-oauth", loginId: "stale-login", - }); - await Promise.resolve(); - expect(resolved).toBe(false); - await handler({ - id: "cancel", type: "cancel_login_provider", provider: "corp-oauth", loginId: "active-login", - }); - const response = await login; - expect(response).toMatchObject({ success: true, data: { provider: "corp-oauth", cancelled: true } }); - expect(harness.authStorage.get("corp-oauth")).toEqual({ type: "api_key", key: "previous" }); - expect(pending.size).toBe(0); + const response = await handler({ id: "login", type: "login_provider", provider: "callback-oauth", authType: "oauth" }); + expect(observed).toEqual([ + "oauth_auth", "oauth_device_code", "oauth_progress", "oauth_info", "oauth_prompt", "oauth_select", "oauth_manual_code", + ]); + expect(await harness.authStorage.read("callback-oauth")) + .toMatchObject({ type: "oauth", access: "engine-secret-acme-one-manual" }); + expect(response).toMatchObject({ success: true, data: { provider: "callback-oauth", cancelled: false } }); + expect(JSON.stringify(response)).not.toContain("engine-secret"); + expect(JSON.stringify(response)).not.toContain("engine-refresh"); }); + it("uses engine-owned acquisition and applies catalog state only after success", async () => { + const harness = await createHarness({ withConfiguredAuth: false }); + harnesses.push(harness); + const reload = vi.spyOn(harness.session.modelRuntime, "reloadCredentials"); + const apply = vi.fn(); + const remoteCatalog = { models: [], scopedModels: [], customAuthProviders: [], oauthProviders: [] }; + const client = { + onExtensionUIRequest: () => () => {}, + respondExtensionUI: async () => {}, + cancelLoginProvider: async () => {}, + requestInternal: async (command: { provider: string }) => ({ provider: command.provider, cancelled: false, ...remoteCatalog }), + }; - it("isolates concurrent provider callbacks and cancellation by login id", async () => { - const { harness, runtime } = await createRuntimeHarness(); + await loginIsolatedOAuthProvider( + harness.session, + client as never, + { apply } as never, + "corp", + { onAuth() {}, onDeviceCode() {}, onPrompt: async () => "", onSelect: async () => undefined }, + ); + expect(apply).toHaveBeenCalledWith(expect.objectContaining({ provider: "corp", cancelled: false })); + expect(reload).toHaveBeenCalledTimes(1); + }); + + it("does not apply or reload cancelled or failed isolated OAuth results", async () => { + const harness = await createHarness({ withConfiguredAuth: false }); + harnesses.push(harness); + const reload = vi.spyOn(harness.session.modelRuntime, "reloadCredentials"); + const apply = vi.fn(); + const baseClient = { + onExtensionUIRequest: () => () => {}, + respondExtensionUI: async () => {}, + cancelLoginProvider: async () => {}, + requestInternal: async () => ({ provider: "corp", cancelled: true }), + }; + const callbacks = { onAuth() {}, onDeviceCode() {}, onPrompt: async () => "", onSelect: async () => undefined }; + + await expect(loginIsolatedOAuthProvider(harness.session, baseClient as never, { apply } as never, "corp", callbacks)) + .rejects.toMatchObject({ message: "Login cancelled" }); + const persistenceFailure = new Error("auth.json is read-only"); + await expect(loginIsolatedOAuthProvider( + harness.session, + { ...baseClient, requestInternal: async () => { throw persistenceFailure; } } as never, + { apply } as never, + "corp", + callbacks, + )).rejects.toBe(persistenceFailure); + expect(apply).not.toHaveBeenCalled(); + expect(reload).not.toHaveBeenCalled(); + }); + + it("normalizes intentional frontend callback cancellation without applying catalog state", async () => { + const harness = await createHarness({ withConfiguredAuth: false }); + harnesses.push(harness); + const abort = new DOMException("dialog closed", "AbortError"); + let listener: ((request: RpcExtensionUIRequest) => void) | undefined; + const apply = vi.fn(); + const reload = vi.spyOn(harness.session.modelRuntime, "reloadCredentials"); + const client = { + onExtensionUIRequest: (next: (request: RpcExtensionUIRequest) => void) => { listener = next; return () => {}; }, + respondExtensionUI: async () => {}, + cancelLoginProvider: async () => {}, + requestInternal: async (command: { provider: string; loginId: string }) => { + listener?.({ type: "extension_ui_request", id: "auth", method: "oauth_auth", provider: command.provider, loginId: command.loginId, info: { url: "https://login.invalid" } }); + return { provider: command.provider, cancelled: true }; + }, + }; + + await expect(loginIsolatedOAuthProvider( + harness.session, + client as never, + { apply } as never, + "corp", + { onAuth: () => { throw abort; }, onDeviceCode() {}, onPrompt: async () => "", onSelect: async () => undefined }, + )).rejects.toMatchObject({ message: "Login cancelled", cause: abort }); + expect(apply).not.toHaveBeenCalled(); + expect(reload).not.toHaveBeenCalled(); + }); +}); + + +describe("RPC OAuth cancellation isolation", () => { + it("cancels one loginId without cancelling a concurrent login for another provider", async () => { + const harness = await createHarness({ withConfiguredAuth: false }); + harnesses.push(harness); for (const provider of ["corp-a", "corp-b"]) { - harness.session.modelRegistry.registerProvider(provider, { + harness.session.modelRuntime.registerProvider(provider, { + baseUrl: `https://${provider}.test/v1`, + api: "openai-completions", oauth: { name: provider, login: async (callbacks) => ({ @@ -268,8 +299,22 @@ describe("isolated engine OAuth", () => { refreshToken: async (credential) => credential, getApiKey: (credential) => credential.access, }, + models: [{ + id: `${provider}-model`, + name: `${provider} model`, + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128_000, + maxTokens: 4_096, + }], }); } + const runtime = new AgentSessionRuntime( + harness.session, + { cwd: harness.tempDir, agentDir: harness.tempDir } as never, + createRuntime, + ); const pending: RpcPendingExtensionRequests = new Map(); const promptIds = new Map(); const handler = createRpcCommandHandler({ @@ -281,6 +326,7 @@ describe("isolated engine OAuth", () => { if ("method" in frame && frame.method === "oauth_prompt") promptIds.set(frame.loginId, frame.id); }, }); + const loginA = handler({ id: "a", type: "login_provider", provider: "corp-a", authType: "oauth", loginId: "login-a" }); let bResolved = false; const loginB = handler({ id: "b", type: "login_provider", provider: "corp-b", authType: "oauth", loginId: "login-b" }) @@ -291,153 +337,64 @@ describe("isolated engine OAuth", () => { expect(await loginA).toMatchObject({ data: { provider: "corp-a", cancelled: true } }); await Promise.resolve(); expect(bResolved).toBe(false); - pending.get(promptIds.get("login-b")!)?.resolve({ type: "extension_ui_response", id: promptIds.get("login-b")!, value: "ok" }); + const promptIdB = promptIds.get("login-b")!; + pending.get(promptIdB)?.resolve({ type: "extension_ui_response", id: promptIdB, value: "ok" }); expect(await loginB).toMatchObject({ data: { provider: "corp-b", cancelled: false } }); - expect(harness.authStorage.get("corp-a")).toBeUndefined(); - expect(harness.authStorage.get("corp-b")).toMatchObject({ type: "oauth", access: "corp-b-ok" }); - }); - it("does not apply or reload cancelled or failed isolated OAuth results", async () => { - let reloads = 0; - let applied = 0; - const client = { - onExtensionUIRequest: () => () => {}, - respondExtensionUI: async () => {}, - cancelLoginProvider: async () => {}, - requestInternal: async () => ({ provider: "corp", cancelled: true as const }), - }; - await expect(loginIsolatedOAuthProvider( - { modelRegistry: { authStorage: { reload: () => { reloads += 1; } } } } as never, - client as never, - { apply: () => { applied += 1; } } as never, - "corp", - { onAuth() {}, onDeviceCode() {}, onPrompt: async () => "", onSelect: async () => undefined }, - )).rejects.toMatchObject({ message: "Login cancelled" }); - expect({ reloads, applied }).toEqual({ reloads: 0, applied: 0 }); - const persistenceFailure = new Error("auth.json is read-only"); - await expect(loginIsolatedOAuthProvider( - { modelRegistry: { authStorage: { reload: () => { reloads += 1; } } } } as never, - { ...client, requestInternal: async () => { throw persistenceFailure; } } as never, - { apply: () => { applied += 1; } } as never, - "corp", - { onAuth() {}, onDeviceCode() {}, onPrompt: async () => "", onSelect: async () => undefined }, - )).rejects.toBe(persistenceFailure); - expect({ reloads, applied }).toEqual({ reloads: 0, applied: 0 }); + expect(await harness.authStorage.read("corp-a")).toBeUndefined(); + expect(await harness.authStorage.read("corp-b")).toMatchObject({ type: "oauth", access: "corp-b-ok" }); }); +}); - - it("normalizes intentional frontend callback cancellation without applying catalog state", async () => { - const abort = new DOMException("dialog closed", "AbortError"); - let listener: ((request: RpcExtensionUIRequest) => void) | undefined; - let applied = 0; - let reloaded = 0; - const client = { - onExtensionUIRequest: (next: (request: RpcExtensionUIRequest) => void) => { listener = next; return () => {}; }, - respondExtensionUI: async () => {}, - cancelLoginProvider: async () => {}, - requestInternal: async (command: { provider: string; loginId: string }) => { - listener?.({ - type: "extension_ui_request", id: "auth", method: "oauth_auth", - provider: command.provider, loginId: command.loginId, info: { url: "https://login.invalid" }, - }); - return { provider: command.provider, cancelled: true as const }; - }, - }; - let caught: unknown; - try { - await loginIsolatedOAuthProvider( - { modelRegistry: { authStorage: { reload: () => { reloaded += 1; } } } } as never, - client as never, - { apply: () => { applied += 1; } } as never, - "corp", - { onAuth: () => { throw abort; }, onDeviceCode() {}, onPrompt: async () => "", onSelect: async () => undefined }, - ); - } catch (error) { - caught = error; - } - expect(caught).toMatchObject({ message: "Login cancelled", cause: abort }); - expect({ applied, reloaded }).toEqual({ applied: 0, reloaded: 0 }); - }); - it("keeps the acquired credential when post-login model refresh reports provider errors", async () => { - const { harness, runtime } = await createRuntimeHarness(); - harness.authStorage.set("corp-oauth", { type: "api_key", key: "previous" }); - harness.session.modelRegistry.registerProvider("corp-oauth", { +describe("RPC OAuth failed transaction isolation", () => { + it("does not refresh or overwrite prior credentials after a failed engine-owned login", async () => { + const { harness, handler } = await createRuntimeHarness(); + harness.session.modelRuntime.registerProvider("corp-oauth", { + baseUrl: "https://provider.test/v1", + api: "openai-completions", oauth: { name: "Corp OAuth", - login: async () => ({ access: "new-secret", refresh: "new-refresh", expires: Date.now() + 60_000 }), + login: async () => { throw new Error("provider denied login"); }, refreshToken: async (credential) => credential, getApiKey: (credential) => credential.access, }, + models: [], }); - harness.session.modelRegistry.refresh = async () => ({ + const refresh = vi.spyOn(harness.session.modelRuntime, "refresh"); + + await expect(handler({ + id: "failed-login", + type: "login_provider", + provider: "corp-oauth", + authType: "oauth", + })).rejects.toThrow("provider denied login"); + expect(await harness.authStorage.read("corp-oauth")).toEqual({ type: "api_key", key: "previous" }); + expect(refresh).not.toHaveBeenCalled(); + }); +}); +describe("RPC OAuth credential survival", () => { + it("keeps the acquired credential when post-login model refresh reports provider errors", async () => { + await expectSuccessfulLoginAndRetainedCredential(async () => ({ aborted: false, errors: new Map([["corp-oauth", new Error("catalog unavailable")]]), - }); - const handler = createRpcCommandHandler({ - runtimeHost: runtime, - getSession: () => harness.session, - rebindSession: async () => {}, - pendingExtensionRequests: new Map(), - output: () => {}, - }); - - const response = await handler({ - id: "login", type: "login_provider", provider: "corp-oauth", authType: "oauth", - }); - expect(response).toMatchObject({ success: true, data: { provider: "corp-oauth", cancelled: false } }); - expect(harness.authStorage.get("corp-oauth")).toMatchObject({ type: "oauth", access: "new-secret" }); + })); }); it("keeps a thrown post-login model refresh failure visible without rolling back", async () => { - const { harness, runtime } = await createRuntimeHarness(); - harness.authStorage.set("corp-oauth", { type: "api_key", key: "previous" }); - harness.session.modelRegistry.registerProvider("corp-oauth", { - oauth: { - name: "Corp OAuth", - login: async () => ({ access: "new-secret", refresh: "new-refresh", expires: Date.now() + 60_000 }), - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }, - }); - const refreshFailure = new DOMException("refresh transport aborted", "AbortError"); - harness.session.modelRegistry.refresh = async () => { throw refreshFailure; }; - const handler = createRpcCommandHandler({ - runtimeHost: runtime, - getSession: () => harness.session, - rebindSession: async () => {}, - pendingExtensionRequests: new Map(), - output: () => {}, - }); + const { harness, handler } = await createRuntimeHarness(); + vi.spyOn(harness.session.modelRuntime, "refresh").mockRejectedValue( + new DOMException("refresh transport aborted", "AbortError"), + ); await expect(handler({ - id: "login", type: "login_provider", provider: "corp-oauth", authType: "oauth", + id: "login", + type: "login_provider", + provider: "corp-oauth", + authType: "oauth", })).rejects.toThrow("refresh transport aborted"); - expect(harness.authStorage.get("corp-oauth")).toMatchObject({ type: "oauth", access: "new-secret" }); + expect(await harness.authStorage.read("corp-oauth")).toMatchObject({ type: "oauth", access: "new-secret" }); }); - it("completes login and keeps the credential when model refresh reports an aborted result", async () => { - const { harness, runtime } = await createRuntimeHarness(); - harness.authStorage.set("corp-oauth", { type: "api_key", key: "previous" }); - harness.session.modelRegistry.registerProvider("corp-oauth", { - oauth: { - name: "Corp OAuth", - login: async () => ({ access: "new-secret", refresh: "new-refresh", expires: Date.now() + 60_000 }), - refreshToken: async (credential) => credential, - getApiKey: (credential) => credential.access, - }, - }); - harness.session.modelRegistry.refresh = async () => ({ aborted: true, errors: new Map() }); - const handler = createRpcCommandHandler({ - runtimeHost: runtime, - getSession: () => harness.session, - rebindSession: async () => {}, - pendingExtensionRequests: new Map(), - output: () => {}, - }); - - const response = await handler({ - id: "login", type: "login_provider", provider: "corp-oauth", authType: "oauth", - }); - expect(response).toMatchObject({ success: true, data: { provider: "corp-oauth", cancelled: false } }); - expect(harness.authStorage.get("corp-oauth")).toMatchObject({ type: "oauth", access: "new-secret" }); + it("keeps the acquired credential when post-login model refresh reports an aborted result", async () => { + await expectSuccessfulLoginAndRetainedCredential(async () => ({ aborted: true, errors: new Map() })); }); }); diff --git a/packages/coding-agent/test/rpc-prompt-response-semantics.test.ts b/packages/coding-agent/test/rpc-prompt-response-semantics.test.ts index 1ba5aa8b6..724fae05e 100644 --- a/packages/coding-agent/test/rpc-prompt-response-semantics.test.ts +++ b/packages/coding-agent/test/rpc-prompt-response-semantics.test.ts @@ -14,10 +14,10 @@ import { afterEach, describe, expect, it, vi } from "vitest"; import { AgentSession } from "../src/core/agent-session.ts"; import type { AgentSessionRuntime } from "../src/core/agent-session-runtime.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; import { runRpcMode } from "../src/modes/rpc/rpc-mode.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; import { withNormalRpcEnvironment } from "./normal-rpc-environment.ts"; import { createTestResourceLoader } from "./utilities.ts"; @@ -97,15 +97,15 @@ function sleep(ms: number): Promise { return new Promise((resolve) => setTimeout(resolve, ms)); } -function createRuntimeHost(options: { +async function createRuntimeHost(options: { withAuth: boolean; responseDelayMs: number; model?: Model; unsupportedFallback?: boolean; -}): { +}): Promise<{ runtimeHost: AgentSessionRuntime; cleanup: () => Promise; -} { +}> { const tempDir = join(tmpdir(), `pi-rpc-prompt-${Date.now()}-${Math.random().toString(36).slice(2)}`); mkdirSync(tempDir, { recursive: true }); @@ -136,17 +136,17 @@ function createRuntimeHost(options: { const sessionManager = SessionManager.inMemory(); const settingsManager = SettingsManager.create(tempDir, tempDir); const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); if (options.withAuth) { - authStorage.setRuntimeApiKey("anthropic", "test-key"); + await authStorage.modify("anthropic", async () => ({ type: "api_key", key: "test-key" })); } + const modelRegistry = await createModelRegistry(authStorage, tempDir); const session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime: getModelRuntime(modelRegistry), resourceLoader: createTestResourceLoader(), }); @@ -155,6 +155,7 @@ function createRuntimeHost(options: { modelFallbackMessage: options.unsupportedFallback ? fallbackWarning : undefined, modelFallbackReason: options.unsupportedFallback ? "configured-provider-unsupported" : undefined, session, + services: { agentDir: tempDir }, newSession: vi.fn(async function(this: { modelFallbackMessage?: string; modelFallbackReason?: string }) { this.modelFallbackMessage = fallbackWarning; this.modelFallbackReason = "configured-provider-unsupported"; @@ -211,7 +212,7 @@ async function startRpcMode(options: { rpcIo.outputLines = []; rpcIo.lineHandler = undefined; - const { runtimeHost, cleanup } = createRuntimeHost(options); + const { runtimeHost, cleanup } = await createRuntimeHost(options); withNormalRpcEnvironment(() => { void runRpcMode(runtimeHost); }); await vi.waitFor(() => expect(rpcIo.lineHandler).toBeDefined()); @@ -263,7 +264,6 @@ describe("RPC prompt response semantics", () => { } }); - it("blocks unsupported prompts but stays live for set_model recovery", async () => { const model = getModel("anthropic", "claude-sonnet-4-5"); if (!model) throw new Error("missing recovery model"); @@ -278,9 +278,9 @@ describe("RPC prompt response semantics", () => { try { lineHandler(JSON.stringify({ id: "blocked", type: "prompt", message: "Do not send" })); await vi.waitFor(() => { - const responses = getPromptResponses(rpcIo.outputLines, "blocked"); - expect(responses).toHaveLength(1); - expect(responses[0]).toMatchObject({ success: false, error: warning }); + expect(getPromptResponses(rpcIo.outputLines, "blocked")).toEqual([ + expect.objectContaining({ success: false, error: warning }), + ]); }); expect(parseOutputLines(rpcIo.outputLines).filter((record) => record.type !== "response")).toEqual([]); expect(rpcIo.outputLines.join("\n")).not.toContain("API key"); @@ -303,7 +303,9 @@ describe("RPC prompt response semantics", () => { lineHandler(JSON.stringify({ id: "replace", type: "new_session" })); await vi.waitFor(() => { - expect(parseOutputLines(rpcIo.outputLines).some((record) => record.id === "replace" && record.success === true)).toBe(true); + expect(parseOutputLines(rpcIo.outputLines).some( + (record) => record.id === "replace" && record.success === true, + )).toBe(true); }); rpcIo.outputLines = []; lineHandler(JSON.stringify({ id: "blocked-again", type: "prompt", message: "blocked again" })); @@ -316,7 +318,8 @@ describe("RPC prompt response semantics", () => { await cleanup(); } }); - it("clears unsupported prompt lock only after a successful changed cycle", async () => { + + it("clears the unsupported prompt lock only after a successful changed cycle", async () => { const initial = getModel("anthropic", "claude-sonnet-4-5"); const selected = getModel("anthropic", "claude-haiku-4-5"); if (!initial || !selected) throw new Error("missing cycle models"); @@ -327,9 +330,9 @@ describe("RPC prompt response semantics", () => { unsupportedFallback: true, }); const lock = (): void => { - (runtimeHost as unknown as { modelFallbackMessage?: string }).modelFallbackMessage = + runtimeHost.modelFallbackMessage = "Configured default model is unavailable or unsupported. Update defaultProvider/defaultModel or use /model."; - (runtimeHost as unknown as { modelFallbackReason?: string }).modelFallbackReason = "configured-provider-unsupported"; + runtimeHost.modelFallbackReason = "configured-provider-unsupported"; }; const cycle = vi.spyOn(runtimeHost.session, "cycleModel"); @@ -366,6 +369,7 @@ describe("RPC prompt response semantics", () => { await cleanup(); } }); + it("emits one success response when prompt preflight succeeds", async () => { const { lineHandler, cleanup } = await startRpcMode({ withAuth: true, responseDelayMs: 0 }); diff --git a/packages/coding-agent/test/rpc-provider-credential-save.test.ts b/packages/coding-agent/test/rpc-provider-credential-save.test.ts new file mode 100644 index 000000000..9c1fc335f --- /dev/null +++ b/packages/coding-agent/test/rpc-provider-credential-save.test.ts @@ -0,0 +1,105 @@ +import { mkdtempSync, rmSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import type { AgentSession } from "../src/core/agent-session.ts"; +import type { AgentSessionRuntime } from "../src/core/agent-session-runtime.ts"; +import { AuthStorage } from "../src/core/auth-storage.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; +import { createRpcCommandHandler } from "../src/modes/rpc/rpc-command-handler.ts"; + +const tempDirs: string[] = []; +beforeEach(() => { + // Saving triggers a catalog refresh. This persistence test does not exercise + // remote catalogs and must not inherit configured provider keys from the host. + vi.stubEnv("ATOMIC_OFFLINE", "1"); +}); +afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllEnvs(); + while (tempDirs.length > 0) rmSync(tempDirs.pop()!, { recursive: true, force: true }); +}); + +async function createHarness() { + const tempDir = mkdtempSync(join(tmpdir(), "atomic-rpc-credential-save-")); + tempDirs.push(tempDir); + const authPath = join(tempDir, "auth.json"); + const modelRuntime = await ModelRuntime.create({ + credentials: AuthStorage.create(authPath), + modelsPath: null, + }); + modelRuntime.registerProvider("persist-probe", { + baseUrl: "https://persist.test/v1", + api: "openai-completions", + models: [ + { + id: "persist-model", + name: "Persist Model", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128_000, + maxTokens: 4096, + }, + ], + }); + await modelRuntime.refresh({ allowNetwork: false }); + const refreshCurrentModelFromRegistry = vi.fn(); + const session = { + modelRuntime, + scopedModels: [], + refreshCurrentModelFromRegistry, + } as unknown as AgentSession; + const handler = createRpcCommandHandler({ + runtimeHost: {} as AgentSessionRuntime, + getSession: () => session, + rebindSession: async () => {}, + output: () => {}, + }); + return { authPath, handler, refreshCurrentModelFromRegistry }; +} + +describe("RPC save_provider_credential", () => { + it("persists an API key and returns the refreshed model catalog", async () => { + const { authPath, handler, refreshCurrentModelFromRegistry } = await createHarness(); + + const response = await handler({ + id: "save", + type: "save_provider_credential", + provider: "persist-probe", + credential: { type: "api_key", key: "persisted-secret" }, + }); + + const freshStorage = AuthStorage.create(authPath); + expect(await freshStorage.read("persist-probe")).toEqual({ type: "api_key", key: "persisted-secret" }); + expect(response).toMatchObject({ + success: true, + data: { + models: expect.arrayContaining([ + expect.objectContaining({ provider: "persist-probe", id: "persist-model" }), + ]), + }, + }); + expect(refreshCurrentModelFromRegistry).toHaveBeenCalledOnce(); + }); + + it("accepts and persists an OAuth credential permitted by the RPC command type", async () => { + const { authPath, handler } = await createHarness(); + const credential = { + type: "oauth" as const, + access: "oauth-access", + refresh: "oauth-refresh", + expires: Date.now() + 60_000, + }; + + const response = await handler({ + id: "save-oauth", + type: "save_provider_credential", + provider: "persist-probe", + credential, + }); + + expect(response).toMatchObject({ success: true }); + expect(await AuthStorage.create(authPath).read("persist-probe")).toEqual(credential); + }); +}); diff --git a/packages/coding-agent/test/runtime-credentials.test.ts b/packages/coding-agent/test/runtime-credentials.test.ts new file mode 100644 index 000000000..a044f715e --- /dev/null +++ b/packages/coding-agent/test/runtime-credentials.test.ts @@ -0,0 +1,42 @@ +import { describe, expect, test } from "vitest"; +import { AuthStorage } from "../src/core/auth-storage.ts"; +import { RuntimeCredentials } from "../src/core/runtime-credentials.ts"; + +describe("RuntimeCredentials", () => { + test("runtime overrides mask stored credentials without persisting", async () => { + const storage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "stored-key" } }); + const credentials = new RuntimeCredentials(storage); + + credentials.setRuntimeApiKey("anthropic", "runtime-key"); + expect(await credentials.read("anthropic")).toEqual({ type: "api_key", key: "runtime-key" }); + expect(await storage.read("anthropic")).toEqual({ type: "api_key", key: "stored-key" }); + + credentials.removeRuntimeApiKey("anthropic"); + expect(await credentials.read("anthropic")).toEqual({ type: "api_key", key: "stored-key" }); + }); + + test("enumeration merges overrides without exposing keys", async () => { + const storage = AuthStorage.inMemory({ + anthropic: { type: "oauth", access: "access", refresh: "refresh", expires: Date.now() + 60_000 }, + }); + const credentials = new RuntimeCredentials(storage); + credentials.setRuntimeApiKey("anthropic", "runtime-key"); + credentials.setRuntimeApiKey("openai", "other-runtime-key"); + + expect(await credentials.list()).toEqual([ + { providerId: "anthropic", type: "api_key" }, + { providerId: "openai", type: "api_key" }, + ]); + }); + + test("delete clears both the override and persisted credential", async () => { + const storage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "stored-key" } }); + const credentials = new RuntimeCredentials(storage); + credentials.setRuntimeApiKey("anthropic", "runtime-key"); + + await credentials.delete("anthropic"); + + expect(await credentials.read("anthropic")).toBeUndefined(); + expect(await credentials.list()).toEqual([]); + }); +}); diff --git a/packages/coding-agent/test/sdk-codex-cache-probe-tool-loop.ts b/packages/coding-agent/test/sdk-codex-cache-probe-tool-loop.ts index c48d890b0..137f16f42 100644 --- a/packages/coding-agent/test/sdk-codex-cache-probe-tool-loop.ts +++ b/packages/coding-agent/test/sdk-codex-cache-probe-tool-loop.ts @@ -33,6 +33,7 @@ import type { ResourceLoader } from "../src/core/resource-loader.ts"; import { createAgentSession } from "../src/core/sdk.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; +import { createInMemoryModelRegistry, createModelRegistry } from "./model-runtime-test-utils.ts"; type Transport = "sse" | "websocket" | "websocket-cached" | "auto"; @@ -275,7 +276,7 @@ async function main(): Promise { mkdirSync(dirname(args.sessionPath), { recursive: true }); const authStorage = AuthStorage.create(); - const modelRegistry = ModelRegistry.create(authStorage); + const modelRegistry = await createModelRegistry(authStorage); const model = getModel("openai-codex", "gpt-5.5"); if (!model) { diff --git a/packages/coding-agent/test/sdk-codex-fast-mode.test.ts b/packages/coding-agent/test/sdk-codex-fast-mode.test.ts index 496a6b76c..56f490f9d 100644 --- a/packages/coding-agent/test/sdk-codex-fast-mode.test.ts +++ b/packages/coding-agent/test/sdk-codex-fast-mode.test.ts @@ -13,7 +13,7 @@ import { ENV_CODEX_FAST_MODE } from "../src/config.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { CODEX_FAST_MODE_SERVICE_TIER } from "../src/core/codex-fast-mode.ts"; import type { OrchestrationContext } from "../src/core/extensions/index.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import { createAgentSession } from "../src/core/sdk.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; @@ -85,7 +85,7 @@ describe("createAgentSession codex fast mode", () => { let tempDir: string; let cwd: string; let agentDir: string; - let registeredProviders: Array<{ registry: ModelRegistry; provider: string }>; + let registeredProviders: Array<{ registry: ModelRuntime; provider: string }>; let previousCodexFastModeEnv: string | undefined; beforeEach(() => { @@ -120,43 +120,44 @@ describe("createAgentSession codex fast mode", () => { orchestrationContext?: OrchestrationContext; payload?: Record; }): Promise { - const api = `codex-fast-capture-${options.provider}-${Math.random().toString(36).slice(2)}` as Api; + const api = "openai-responses" as Api; const model = createModel(options.provider, api); const authStorage = AuthStorage.create(join(agentDir, "auth.json")); - authStorage.setRuntimeApiKey(options.provider, "test-api-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); + await authStorage.modify(options.provider, async () => ({ type: "api_key", key: "test-api-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: join(agentDir, "models.json"), allowModelNetwork: false }); const settingsManager = SettingsManager.inMemory({ codexFastMode: options.settings }); const sessionManager = SessionManager.inMemory(cwd); let capturedOptions: SimpleStreamOptions | undefined; - modelRegistry.registerProvider(options.provider, { + modelRuntime.registerProvider(options.provider, { api, streamSimple: (_model, _context, streamOptions) => { capturedOptions = streamOptions; return createDoneStream(model); }, }); - registeredProviders.push({ registry: modelRegistry, provider: options.provider }); + registeredProviders.push({ registry: modelRuntime, provider: options.provider }); const { session } = await createAgentSession({ cwd, agentDir, model, authStorage, - modelRegistry, + modelRuntime, settingsManager, sessionManager, orchestrationContext: options.orchestrationContext, }); try { - await session.agent.streamFunction(model, { messages: [] }, { sessionId: session.sessionId }); + const stream = await session.agent.streamFunction(model, { messages: [] }, { sessionId: session.sessionId }); + await stream.result(); const payload = await session.agent.onPayload?.(options.payload ?? { model: model.id }, model); return { options: capturedOptions, payload }; } finally { session.dispose(); - modelRegistry.unregisterProvider(options.provider); - registeredProviders = registeredProviders.filter((entry) => entry.registry !== modelRegistry || entry.provider !== options.provider); + modelRuntime.unregisterProvider(options.provider); + registeredProviders = registeredProviders.filter((entry) => entry.registry !== modelRuntime || entry.provider !== options.provider); } } @@ -175,8 +176,8 @@ describe("createAgentSession codex fast mode", () => { it("preserves custom provider streaming for native OpenAI APIs when fast mode is enabled", async () => { const model = createModel("openai", "openai-responses"); const authStorage = AuthStorage.create(join(agentDir, "auth.json")); - authStorage.setRuntimeApiKey("openai", "test-api-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); + await authStorage.modify("openai", async () => ({ type: "api_key", key: "test-api-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: join(agentDir, "models.json"), allowModelNetwork: false }); const settingsManager = SettingsManager.inMemory({ codexFastMode: { chat: true, workflow: false } }); const sessionManager = SessionManager.inMemory(cwd); let capturedOptions: SimpleStreamOptions | undefined; @@ -185,21 +186,21 @@ describe("createAgentSession codex fast mode", () => { }); vi.stubGlobal("fetch", nativeFetch); - modelRegistry.registerProvider("openai", { + modelRuntime.registerProvider("openai", { api: "openai-responses", streamSimple: (_model, _context, streamOptions) => { capturedOptions = streamOptions; return createDoneStream(model); }, }); - registeredProviders.push({ registry: modelRegistry, provider: "openai" }); + registeredProviders.push({ registry: modelRuntime, provider: "openai" }); const { session } = await createAgentSession({ cwd, agentDir, model, authStorage, - modelRegistry, + modelRuntime, settingsManager, sessionManager, }); @@ -215,9 +216,9 @@ describe("createAgentSession codex fast mode", () => { ); } finally { session.dispose(); - modelRegistry.unregisterProvider("openai"); + modelRuntime.unregisterProvider("openai"); registeredProviders = registeredProviders.filter( - (entry) => entry.registry !== modelRegistry || entry.provider !== "openai", + (entry) => entry.registry !== modelRuntime || entry.provider !== "openai", ); } }); @@ -246,7 +247,7 @@ describe("createAgentSession codex fast mode", () => { it("uses the workflow setting for workflow-stage requests", async () => { const disabled = await captureFastModeRequest({ - provider: "openai-codex", + provider: "openai", settings: { chat: true, workflow: false }, orchestrationContext: workflowContext, }); @@ -254,7 +255,7 @@ describe("createAgentSession codex fast mode", () => { expect(disabled.payload).not.toMatchObject({ service_tier: CODEX_FAST_MODE_SERVICE_TIER }); const enabled = await captureFastModeRequest({ - provider: "openai-codex", + provider: "openai", settings: { chat: false, workflow: true }, orchestrationContext: workflowContext, }); @@ -277,8 +278,8 @@ describe("createAgentSession codex fast mode", () => { it("sends priority service tier in native OpenAI Responses request bodies", async () => { const model = createModel("openai", "openai-responses"); const authStorage = AuthStorage.create(join(agentDir, "auth.json")); - authStorage.setRuntimeApiKey("openai", "test-api-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); + await authStorage.modify("openai", async () => ({ type: "api_key", key: "test-api-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: join(agentDir, "models.json"), allowModelNetwork: false }); const settingsManager = SettingsManager.inMemory({ codexFastMode: { chat: true, workflow: false } }); const sessionManager = SessionManager.inMemory(cwd); let capturedPayload: Record | undefined; @@ -312,7 +313,7 @@ describe("createAgentSession codex fast mode", () => { agentDir, model, authStorage, - modelRegistry, + modelRuntime, settingsManager, sessionManager, }); diff --git a/packages/coding-agent/test/sdk-openrouter-attribution.test.ts b/packages/coding-agent/test/sdk-openrouter-attribution.test.ts index 3ba811ada..f54cf8394 100644 --- a/packages/coding-agent/test/sdk-openrouter-attribution.test.ts +++ b/packages/coding-agent/test/sdk-openrouter-attribution.test.ts @@ -6,46 +6,47 @@ import { type AssistantMessage, createAssistantMessageEventStream, type Model, + type ProviderHeaders, type SimpleStreamOptions, -} from "@earendil-works/pi-ai/compat"; +} from "@earendil-works/pi-ai"; import { afterEach, beforeEach, describe, expect, it } from "vitest"; -import { APP_NAME } from "../src/config.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { createAgentSession } from "../src/core/sdk.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; -describe("createAgentSession OpenRouter attribution headers", () => { +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; + +describe("createAgentSession provider attribution headers", () => { let tempDir: string; let cwd: string; let agentDir: string; let originalTelemetryEnv: string | undefined; beforeEach(() => { - tempDir = join(tmpdir(), `atomic-sdk-openrouter-test-${Date.now()}-${Math.random().toString(36).slice(2)}`); + tempDir = join(tmpdir(), `pi-sdk-attribution-test-${Date.now()}-${Math.random().toString(36).slice(2)}`); cwd = join(tempDir, "project"); agentDir = join(tempDir, "agent"); mkdirSync(cwd, { recursive: true }); mkdirSync(agentDir, { recursive: true }); - originalTelemetryEnv = process.env.ATOMIC_TELEMETRY; - delete process.env.ATOMIC_TELEMETRY; + originalTelemetryEnv = process.env.PI_TELEMETRY; + delete process.env.PI_TELEMETRY; }); afterEach(() => { if (originalTelemetryEnv === undefined) { - delete process.env.ATOMIC_TELEMETRY; + delete process.env.PI_TELEMETRY; } else { - process.env.ATOMIC_TELEMETRY = originalTelemetryEnv; + process.env.PI_TELEMETRY = originalTelemetryEnv; } if (tempDir && existsSync(tempDir)) { rmSync(tempDir, { recursive: true, force: true }); } }); - function createModel(provider: string, baseUrl: string): Model { + function createModel(provider: string, baseUrl: string, id = `${provider}-test-model`): Model { return { - id: `${provider}-test-model`, + id, name: `${provider} Test Model`, api: "openai-completions", provider, @@ -89,31 +90,27 @@ describe("createAgentSession OpenRouter attribution headers", () => { requestHeaders?: Record; sessionId?: string; } = {}, - ): Promise | undefined> { + ): Promise { const settingsManager = SettingsManager.create(cwd, agentDir); if (options.telemetryEnabled === false) { settingsManager.setEnableInstallTelemetry(false); } const authStorage = AuthStorage.create(join(agentDir, "auth.json")); - authStorage.setRuntimeApiKey(model.provider, "test-api-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const registeredProviders = ["capture-provider"]; + await authStorage.modify(model.provider, async () => ({ type: "api_key", key: "test-api-key" })); + const modelRegistry = await createModelRegistry(authStorage, join(agentDir, "models.json")); let capturedOptions: SimpleStreamOptions | undefined; - modelRegistry.registerProvider("capture-provider", { - api: "openai-completions", + modelRegistry.registerProvider(model.provider, { + api: model.api, + headers: options.providerHeaders, streamSimple: (_model, _context, providerOptions) => { capturedOptions = providerOptions; return createDoneStream(); }, }); - if (options.providerHeaders) { - modelRegistry.registerProvider(model.provider, { headers: options.providerHeaders }); - registeredProviders.push(model.provider); - } - + const modelRuntime = getModelRuntime(modelRegistry); const sessionManager = SessionManager.inMemory(cwd); if (options.sessionId) { sessionManager.newSession({ id: options.sessionId }); @@ -123,14 +120,13 @@ describe("createAgentSession OpenRouter attribution headers", () => { cwd, agentDir, model, - authStorage, - modelRegistry, + modelRuntime, settingsManager, sessionManager, }); try { - await session.agent.streamFunction( + const stream = await session.agent.streamFunction( model, { messages: [] }, { @@ -138,12 +134,11 @@ describe("createAgentSession OpenRouter attribution headers", () => { ...(options.requestHeaders ? { headers: options.requestHeaders } : {}), }, ); + await stream.result(); return capturedOptions?.headers; } finally { session.dispose(); - for (const provider of registeredProviders.reverse()) { - modelRegistry.unregisterProvider(provider); - } + modelRegistry.unregisterProvider(model.provider); } } @@ -151,7 +146,7 @@ describe("createAgentSession OpenRouter attribution headers", () => { const headers = await captureHeaders(createModel("openrouter", "https://openrouter.ai/api/v1")); expect(headers?.["HTTP-Referer"]).toBe("https://atomic.sh"); - expect(headers?.["X-OpenRouter-Title"]).toBe(APP_NAME); + expect(headers?.["X-OpenRouter-Title"]).toBe("atomic"); expect(headers?.["X-OpenRouter-Categories"]).toBe("cli-agent"); }); @@ -169,24 +164,18 @@ describe("createAgentSession OpenRouter attribution headers", () => { const headers = await captureHeaders(createModel("custom-openrouter", "https://openrouter.ai/api/v1")); expect(headers?.["HTTP-Referer"]).toBe("https://atomic.sh"); - expect(headers?.["X-OpenRouter-Title"]).toBe(APP_NAME); + expect(headers?.["X-OpenRouter-Title"]).toBe("atomic"); expect(headers?.["X-OpenRouter-Categories"]).toBe("cli-agent"); }); - it("does not add OpenRouter attribution headers for substring-matched custom hosts", async () => { - const headers = await captureHeaders(createModel("custom-openrouter-like", "https://openrouter.ai.evil.test/v1")); + it("does not attribute a different OpenRouter subdomain as the OpenRouter API", async () => { + const headers = await captureHeaders(createModel("custom-openrouter", "https://proxy.openrouter.ai/v1")); expect(headers?.["HTTP-Referer"]).toBeUndefined(); expect(headers?.["X-OpenRouter-Title"]).toBeUndefined(); expect(headers?.["X-OpenRouter-Categories"]).toBeUndefined(); }); - it("adds Atomic attribution headers for NVIDIA NIM models", async () => { - const headers = await captureHeaders(createModel("nvidia", "https://integrate.api.nvidia.com/v1")); - - expect(headers?.["X-BILLING-INVOKE-ORIGIN"]).toBe("Atomic"); - }); - it("lets provider and request headers override the defaults", async () => { const headers = await captureHeaders(createModel("openrouter", "https://openrouter.ai/api/v1"), { providerHeaders: { @@ -203,13 +192,63 @@ describe("createAgentSession OpenRouter attribution headers", () => { expect(headers?.["X-OpenRouter-Categories"]).toBe("provider-category"); }); + it("adds default attribution headers for direct NVIDIA NIM endpoints", async () => { + const headers = await captureHeaders(createModel("custom-nim", "https://integrate.api.nvidia.com/v1")); + + expect(headers?.["X-BILLING-INVOKE-ORIGIN"]).toBe("Atomic"); + }); + + it("adds default attribution headers for the NVIDIA provider", async () => { + const headers = await captureHeaders(createModel("nvidia", "https://example.test/v1")); + + expect(headers?.["X-BILLING-INVOKE-ORIGIN"]).toBe("Atomic"); + }); + + it("does not add NVIDIA NIM attribution headers when telemetry is disabled", async () => { + const headers = await captureHeaders(createModel("nvidia", "https://integrate.api.nvidia.com/v1"), { + telemetryEnabled: false, + }); + + expect(headers?.["X-BILLING-INVOKE-ORIGIN"]).toBeUndefined(); + }); + + it("lets provider and request headers override NVIDIA NIM defaults", async () => { + const headers = await captureHeaders(createModel("nvidia", "https://integrate.api.nvidia.com/v1"), { + providerHeaders: { + "X-BILLING-INVOKE-ORIGIN": "Provider", + }, + requestHeaders: { + "X-BILLING-INVOKE-ORIGIN": "Request", + }, + }); + + expect(headers?.["X-BILLING-INVOKE-ORIGIN"]).toBe("Request"); + }); + + it("does not add NVIDIA NIM attribution headers for NVIDIA models routed through OpenRouter", async () => { + const headers = await captureHeaders( + createModel("openrouter", "https://openrouter.ai/api/v1", "nvidia/nemotron-3-super-120b-a12b"), + ); + + expect(headers?.["HTTP-Referer"]).toBe("https://atomic.sh"); + expect(headers?.["X-BILLING-INVOKE-ORIGIN"]).toBeUndefined(); + }); + + it("does not add NVIDIA NIM attribution headers for NVIDIA models routed through Vercel AI Gateway", async () => { + const headers = await captureHeaders( + createModel("vercel-ai-gateway", "https://ai-gateway.vercel.sh/v1", "nvidia/nemotron-3-super-120b-a12b"), + ); + + expect(headers?.["X-BILLING-INVOKE-ORIGIN"]).toBeUndefined(); + }); + it("adds OpenCode session headers", async () => { const headers = await captureHeaders(createModel("opencode", "https://opencode.ai/zen/v1"), { sessionId: "opencode-session", }); expect(headers?.["x-opencode-session"]).toBe("opencode-session"); - expect(headers?.["x-opencode-client"]).toBe(APP_NAME); + expect(headers?.["x-opencode-client"]).toBe("atomic"); }); it("lets configured OpenCode headers override the defaults", async () => { diff --git a/packages/coding-agent/test/sdk-session-manager.test.ts b/packages/coding-agent/test/sdk-session-manager.test.ts index c91b0199f..bfdd3c2f9 100644 --- a/packages/coding-agent/test/sdk-session-manager.test.ts +++ b/packages/coding-agent/test/sdk-session-manager.test.ts @@ -1,20 +1,67 @@ import { existsSync, mkdirSync, realpathSync, rmSync } from "node:fs"; import { tmpdir } from "node:os"; -import { join, sep } from "node:path"; +import { join } from "node:path"; import { getModel } from "@earendil-works/pi-ai/compat"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; 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 { createAgentSession } from "../src/core/sdk.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; +const missingApiKeyEnv = "ATOMIC_SDK_SESSION_MANAGER_MISSING_KEY"; +const testModel = (id: string) => ({ + id, + name: id, + reasoning: false, + input: ["text" as const], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 8192, + maxTokens: 1024, +}); + +async function createSelectionRuntime(): Promise { + const runtime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory(), + modelsPath: null, + allowModelNetwork: false, + }); + runtime.registerProvider("test-ready", { + api: "openai-completions", + baseUrl: "https://ready.invalid/v1", + apiKey: "test-key", + models: [testModel("ready-model")], + }); + runtime.registerProvider("test-locked", { + api: "openai-completions", + baseUrl: "https://locked.invalid/v1", + apiKey: `$${missingApiKeyEnv}`, + models: [testModel("locked-model")], + }); + await runtime.refresh({ allowNetwork: false }); + return runtime; +} + +function persistedSession(cwd: string, provider: string, modelId: string): SessionManager { + const manager = SessionManager.inMemory(cwd); + manager.appendModelChange(provider, modelId); + manager.appendMessage({ + role: "user", + content: [{ type: "text", text: "restore" }], + timestamp: Date.now(), + }); + return manager; +} + describe("createAgentSession session manager defaults", () => { let tempDir: string; let cwd: string; let agentDir: string; + let previousMissingApiKey: string | undefined; beforeEach(() => { + previousMissingApiKey = process.env[missingApiKeyEnv]; + delete process.env[missingApiKeyEnv]; tempDir = join(tmpdir(), `pi-sdk-session-test-${Date.now()}-${Math.random().toString(36).slice(2)}`); cwd = join(tempDir, "project"); agentDir = join(tempDir, "agent"); @@ -23,9 +70,10 @@ describe("createAgentSession session manager defaults", () => { }); afterEach(() => { - if (tempDir && existsSync(tempDir)) { - rmSync(tempDir, { recursive: true, force: true }); - } + if (previousMissingApiKey === undefined) delete process.env[missingApiKeyEnv]; + else process.env[missingApiKeyEnv] = previousMissingApiKey; + vi.restoreAllMocks(); + if (tempDir && existsSync(tempDir)) rmSync(tempDir, { recursive: true, force: true }); }); it("uses agentDir for the default persisted session path", async () => { @@ -44,7 +92,7 @@ describe("createAgentSession session manager defaults", () => { const sessionFile = session.sessionManager.getSessionFile(); expect(sessionDir).toBe(expectedSessionDir); - expect(sessionFile?.startsWith(`${expectedSessionDir}${sep}`)).toBe(true); + expect(sessionFile?.startsWith(`${expectedSessionDir}/`)).toBe(true); session.dispose(); }); @@ -81,11 +129,11 @@ describe("createAgentSession session manager defaults", () => { }); expect(session.sessionManager).toBe(sessionManager); - expect(session.systemPrompt).toContain(`Current working directory: ${sessionCwd.replaceAll("\\", "/")}`); + expect(session.systemPrompt).toContain(`Current working directory: ${sessionCwd}`); const bashTool = session.agent.state.tools.find((tool) => tool.name === "bash"); expect(bashTool).toBeTruthy(); - const result = await bashTool!.execute("test", { command: 'bun -e "console.log(process.cwd())"' }); + const result = await bashTool!.execute("test", { command: "pwd" }); const output = result.content .filter((item): item is { type: "text"; text: string } => item.type === "text") .map((item) => item.text) @@ -96,7 +144,7 @@ describe("createAgentSession session manager defaults", () => { session.dispose(); }); - it("enables ask_user_question and todo by default", async () => { + it("exposes current session state to the built-in bash tool", async () => { const model = getModel("anthropic", "claude-sonnet-4-5"); expect(model).toBeTruthy(); @@ -104,126 +152,121 @@ describe("createAgentSession session manager defaults", () => { cwd, agentDir, model: model!, + thinkingLevel: "high", }); + expect(session.sessionFile).toBeTruthy(); + expect(session.systemPrompt).toContain( + "Inspect ATOMIC_* or PI_* environment variables for current model and session details.", + ); + + const bashTool = session.agent.state.tools.find((tool) => tool.name === "bash"); + expect(bashTool).toBeTruthy(); + const result = await bashTool!.execute("test", { + command: `printf '%s\\n' "$PI_SESSION_ID" "$PI_SESSION_FILE" "$PI_PROVIDER" "$PI_MODEL" "$PI_REASONING_LEVEL"`, + }); + const output = result.content + .filter((item): item is { type: "text"; text: string } => item.type === "text") + .map((item) => item.text) + .join(""); + + expect(output.trim().split("\n")).toEqual([ + session.sessionId, + session.sessionFile, + model!.provider, + model!.id, + session.thinkingLevel, + ]); + + session.dispose(); + }); + + it("enables ask_user_question and todo by default", async () => { + const model = getModel("anthropic", "claude-sonnet-4-5"); + expect(model).toBeTruthy(); + const { session } = await createAgentSession({ cwd, agentDir, model: model! }); expect(session.getActiveToolNames()).toEqual( expect.arrayContaining(["read", "bash", "edit", "write", "ask_user_question", "todo"]), ); - session.dispose(); }); - it("synthesizes an absent custom model id only while restoring persisted session state", async () => { - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openrouter", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const sessionManager = SessionManager.inMemory(cwd); - sessionManager.appendModelChange("openrouter", "future/custom-restored-model"); - sessionManager.appendMessage({ - role: "user", - content: [{ type: "text", text: "restore me" }], - timestamp: Date.now(), - }); + it("restores an absent model id for a registered OpenAI-compatible provider", async () => { + const modelRuntime = await createSelectionRuntime(); + const restoredModelId = "future/custom-restored-model"; - const { session, modelFallbackMessage } = await createAgentSession({ + const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ cwd, agentDir, - authStorage, - modelRegistry, + modelRuntime, settingsManager: SettingsManager.inMemory(), - sessionManager, + sessionManager: persistedSession(cwd, "test-ready", restoredModelId), }); - expect(session.model?.provider).toBe("openrouter"); - expect(session.model?.id).toBe("future/custom-restored-model"); + expect(session.model?.provider).toBe("test-ready"); + expect(session.model?.id).toBe(restoredModelId); expect(modelFallbackMessage).toBeUndefined(); + expect(modelFallbackReason).toBeUndefined(); session.dispose(); }); it("does not synthesize an exact unauthenticated model during SDK session restoration", async () => { - const authStorage = AuthStorage.inMemory(); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const openRouterModels = modelRegistry.getAll().filter((model) => model.provider === "openrouter"); - const exactModel = openRouterModels[0]; - const sameProviderTemplate = openRouterModels[1]; - expect(exactModel).toBeDefined(); - expect(sameProviderTemplate).toBeDefined(); - - vi.spyOn(modelRegistry, "hasConfiguredAuth").mockImplementation((model) => model !== exactModel); - expect(modelRegistry.getAvailable()).toContain(sameProviderTemplate); - expect(modelRegistry.getAvailable()).not.toContain(exactModel); - const sessionManager = SessionManager.inMemory(cwd); - sessionManager.appendModelChange(exactModel!.provider, exactModel!.id); - sessionManager.appendMessage({ - role: "user", - content: [{ type: "text", text: "restore without auth" }], - timestamp: Date.now(), - }); + const modelRuntime = await createSelectionRuntime(); + const locked = modelRuntime.getModel("test-locked", "locked-model"); + expect(locked).toBeDefined(); + expect(modelRuntime.hasConfiguredAuth("test-locked")).toBe(false); const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ cwd, agentDir, - authStorage, - modelRegistry, + modelRuntime, settingsManager: SettingsManager.inMemory(), - sessionManager, + sessionManager: persistedSession(cwd, "test-locked", "locked-model"), }); - expect(session.model).not.toBe(exactModel); - expect(session.model?.id).not.toBe(exactModel!.id); - expect(modelFallbackMessage).toContain(`${exactModel!.provider}/${exactModel!.id}`); + expect(session.model).not.toEqual(locked); + expect(modelFallbackMessage).toContain("test-locked/locked-model"); expect(modelFallbackReason).toBe("session-restore"); session.dispose(); }); it("propagates a generic warning for an unusable complete saved default without switching providers", async () => { - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); + const modelRuntime = await createSelectionRuntime(); const settingsManager = SettingsManager.inMemory({ - defaultProvider: ["cur", "sor"].join(""), - defaultModel: ["composer", "-2"].join(""), + defaultProvider: "unsupported-provider", + defaultModel: "unsupported-model", }); - const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ cwd, agentDir, - authStorage, - modelRegistry, + modelRuntime, settingsManager, sessionManager: SessionManager.inMemory(cwd), }); - expect(session.model?.provider).toBe("unknown"); - expect(session.model?.provider).not.toBe("openai"); - expect(typeof modelFallbackMessage).toBe("string"); + expect(session.model?.provider).not.toBe("test-ready"); expect(modelFallbackMessage).toBe( "Configured default model is unavailable or unsupported. Update defaultProvider/defaultModel or use /model.", ); expect(modelFallbackReason).toBe("configured-provider-unsupported"); - expect(settingsManager.getDefaultProvider()).toBe(["cur", "sor"].join("")); - expect(settingsManager.getDefaultModel()).toBe(["composer", "-2"].join("")); + expect(settingsManager.getDefaultProvider()).toBe("unsupported-provider"); + expect(settingsManager.getDefaultModel()).toBe("unsupported-model"); session.dispose(); }); - it("keeps normal automatic selection for an unknown model on a supported provider", async () => { - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const settingsManager = SettingsManager.inMemory({ - defaultProvider: "openai", - defaultModel: "unknown-saved-model", - }); + it("keeps normal automatic selection for an unknown model on a supported provider", async () => { + const modelRuntime = await createSelectionRuntime(); const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ cwd, agentDir, - authStorage, - modelRegistry, - settingsManager, + modelRuntime, + settingsManager: SettingsManager.inMemory({ + defaultProvider: "test-ready", + defaultModel: "unknown-saved-model", + }), sessionManager: SessionManager.inMemory(cwd), }); - expect(session.model?.provider).toBe("openai"); expect(session.model?.id).not.toBe("unknown-saved-model"); expect(modelFallbackMessage).toBeUndefined(); expect(modelFallbackReason).toBeUndefined(); @@ -231,55 +274,39 @@ describe("createAgentSession session manager defaults", () => { }); it("keeps normal automatic selection when a supported exact default lacks auth", async () => { - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const savedModel = modelRegistry.getAll().find((model) => model.provider === "anthropic"); - if (!savedModel) throw new Error("missing Anthropic model fixture"); - const settingsManager = SettingsManager.inMemory({ - defaultProvider: savedModel.provider, - defaultModel: savedModel.id, - }); - - const { session, modelFallbackMessage } = await createAgentSession({ + const modelRuntime = await createSelectionRuntime(); + const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ cwd, agentDir, - authStorage, - modelRegistry, - settingsManager, + modelRuntime, + settingsManager: SettingsManager.inMemory({ + defaultProvider: "test-locked", + defaultModel: "locked-model", + }), sessionManager: SessionManager.inMemory(cwd), }); - expect(session.model?.provider).toBe("openai"); - expect(session.model).not.toBe(savedModel); + expect(session.model?.provider).not.toBe("test-locked"); expect(modelFallbackMessage).toBeUndefined(); + expect(modelFallbackReason).toBeUndefined(); session.dispose(); }); it("gives an unsupported saved provider precedence over failed persisted-session restoration", async () => { - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const removedProvider = ["cur", "sor"].join(""); - const removedModel = ["composer", "-2"].join(""); - const sessionManager = SessionManager.inMemory(cwd); - sessionManager.appendModelChange(removedProvider, removedModel); - sessionManager.appendMessage({ - role: "user", - content: [{ type: "text", text: "persisted stale model" }], - timestamp: Date.now(), - }); - + const modelRuntime = await createSelectionRuntime(); const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ cwd, agentDir, - authStorage, - modelRegistry, - settingsManager: SettingsManager.inMemory({ defaultProvider: removedProvider, defaultModel: removedModel }), - sessionManager, + modelRuntime, + settingsManager: SettingsManager.inMemory({ + defaultProvider: "unsupported-provider", + defaultModel: "unsupported-model", + }), + sessionManager: persistedSession(cwd, "absent-session-provider", "absent-session-model"), }); expect(session.model?.provider).toBe("unknown"); + expect(session.model?.provider).not.toBe("test-ready"); expect(modelFallbackMessage).toBe( "Configured default model is unavailable or unsupported. Update defaultProvider/defaultModel or use /model.", ); @@ -287,75 +314,55 @@ describe("createAgentSession session manager defaults", () => { session.dispose(); }); - it("preserves failed restoration guidance when a valid saved default is selected", async () => { - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const savedModel = modelRegistry.getAvailable().find((model) => model.provider === "openai"); - if (!savedModel) throw new Error("missing OpenAI model fixture"); - const sessionManager = SessionManager.inMemory(cwd); - sessionManager.appendModelChange("absent-session-provider", "absent-session-model"); - sessionManager.appendMessage({ role: "user", content: [{ type: "text", text: "restore" }], timestamp: Date.now() }); - - const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ - cwd, agentDir, authStorage, modelRegistry, sessionManager, - settingsManager: SettingsManager.inMemory({ defaultProvider: savedModel.provider, defaultModel: savedModel.id }), - }); - - expect(session.model).toBe(savedModel); - expect(modelFallbackMessage).toContain("Could not restore model absent-session-provider/absent-session-model"); - expect(modelFallbackMessage).toContain(`Using ${savedModel.provider}/${savedModel.id}`); - expect(modelFallbackReason).toBe("session-restore"); - session.dispose(); - }); - it("preserves restoration guidance with supported unknown and unauthenticated saved defaults", async () => { - for (const defaultKind of ["unknown", "unauthenticated"] as const) { - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("openai", "test-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - const unauthenticated = modelRegistry.getAll().find((model) => model.provider === "anthropic"); - if (!unauthenticated) throw new Error("missing Anthropic model fixture"); - const sessionManager = SessionManager.inMemory(cwd); - sessionManager.appendModelChange("absent-session-provider", "absent-session-model"); - sessionManager.appendMessage({ role: "user", content: [{ type: "text", text: defaultKind }], timestamp: Date.now() }); - const settingsManager = SettingsManager.inMemory(defaultKind === "unknown" - ? { defaultProvider: "openai", defaultModel: "unknown-saved-model" } - : { defaultProvider: unauthenticated.provider, defaultModel: unauthenticated.id }); - + for (const savedDefault of [ + { defaultProvider: "test-ready", defaultModel: "unknown-saved-model" }, + { defaultProvider: "test-locked", defaultModel: "locked-model" }, + ]) { + const modelRuntime = await createSelectionRuntime(); const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ - cwd, agentDir, authStorage, modelRegistry, settingsManager, sessionManager, + cwd, + agentDir, + modelRuntime, + settingsManager: SettingsManager.inMemory(savedDefault), + sessionManager: persistedSession(cwd, "absent-session-provider", "absent-session-model"), }); - expect(session.model?.provider).toBe("openai"); - expect(modelFallbackMessage).toContain("Could not restore model absent-session-provider/absent-session-model"); + expect(modelFallbackMessage).toContain( + "Could not restore model absent-session-provider/absent-session-model", + ); expect(modelFallbackMessage).toContain(`Using ${session.model?.provider}/${session.model?.id}`); expect(modelFallbackReason).toBe("session-restore"); session.dispose(); } }); + it("classifies ordinary empty catalogs separately from unsupported providers", async () => { - const authStorage = AuthStorage.inMemory(); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); - vi.spyOn(modelRegistry, "getAvailable").mockReturnValue([]); + const modelRuntime = await createSelectionRuntime(); + vi.spyOn(modelRuntime, "getAvailableSnapshot").mockReturnValue([]); const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ - cwd, agentDir, authStorage, modelRegistry, + cwd, + agentDir, + modelRuntime, settingsManager: SettingsManager.inMemory(), sessionManager: SessionManager.inMemory(cwd), }); + expect(modelFallbackMessage).toContain("No models available"); expect(modelFallbackReason).toBe("no-models-available"); session.dispose(); }); - it("marks the session header internal when a workflow-stage orchestration context is supplied", async () => { + it("marks workflow-stage sessions internal with their orchestration identity", async () => { const model = getModel("anthropic", "claude-sonnet-4-5"); expect(model).toBeTruthy(); + const sessionManager = SessionManager.inMemory(cwd); const { session } = await createAgentSession({ cwd, agentDir, model: model!, + sessionManager, orchestrationContext: { kind: "workflow-stage", workflowRunId: "run-42", @@ -365,10 +372,47 @@ describe("createAgentSession session manager defaults", () => { }, }); - const header = session.sessionManager.getHeader(); - expect(header?.internal).toBe(true); - expect(header?.workflow).toEqual({ runId: "run-42", stageId: "stage-7", stageName: "build" }); + expect(sessionManager.getHeader()?.internal).toBe(true); + expect(sessionManager.getHeader()?.workflow).toEqual({ + runId: "run-42", + stageId: "stage-7", + stageName: "build", + }); + session.dispose(); + }); + + it("reports session-restore fallback reason and selected replacement model", async () => { + const modelRuntime = await ModelRuntime.create({ + credentials: AuthStorage.inMemory({ openai: { type: "api_key", key: "test-key" } }), + modelsPath: null, + allowModelNetwork: false, + }); + const replacement = modelRuntime.getAvailableSnapshot().find((model) => model.provider === "openai"); + expect(replacement).toBeDefined(); + const sessionManager = SessionManager.inMemory(cwd); + sessionManager.appendModelChange("absent-session-provider", "absent-session-model"); + sessionManager.appendMessage({ + role: "user", + content: [{ type: "text", text: "restore" }], + timestamp: Date.now(), + }); + + const { session, modelFallbackMessage, modelFallbackReason } = await createAgentSession({ + cwd, + agentDir, + modelRuntime, + sessionManager, + settingsManager: SettingsManager.inMemory({ + defaultProvider: replacement!.provider, + defaultModel: replacement!.id, + }), + }); + expect(session.model).toEqual(replacement); + expect(modelFallbackReason).toBe("session-restore"); + expect(modelFallbackMessage).toBe( + `Could not restore model absent-session-provider/absent-session-model. Using ${replacement!.provider}/${replacement!.id}`, + ); session.dispose(); }); }); diff --git a/packages/coding-agent/test/sdk-shared-model-registry.test.ts b/packages/coding-agent/test/sdk-shared-model-registry.test.ts index f142f0fbe..67317b7cd 100644 --- a/packages/coding-agent/test/sdk-shared-model-registry.test.ts +++ b/packages/coding-agent/test/sdk-shared-model-registry.test.ts @@ -1,17 +1,8 @@ /** - * Regression: createAgentSession must reuse a supplied ModelRegistry (and its - * AuthStorage) instead of eagerly constructing a fresh AuthStorage. - * - * Workflow stages reuse one ModelRegistry across model-fallback candidates so a - * successfully-loaded primary session's credentials are not discarded and - * re-loaded per candidate. Re-loading under auth.json lock contention can fail - * and leave an empty in-memory credential set, misreporting configured - * providers as "No API key found" (issue #1431). A fresh AuthStorage also calls - * reload() in its constructor, so even building one only to throw it away takes - * the same contended file lock — createAgentSession must avoid that when a - * registry is provided. - * - * cross-ref: packages/coding-agent/src/core/sdk.ts (createAgentSession) + * Regression: createAgentSession must reuse a supplied ModelRuntime instead of + * eagerly constructing a fresh AuthStorage. Workflow stages share the runtime + * so credentials survive fallback-candidate session creation without another + * contended auth.json read (issue #1431). */ import { mkdirSync, rmSync } from "node:fs"; import { tmpdir } from "node:os"; @@ -20,10 +11,10 @@ import { getModel } from "@earendil-works/pi-ai/compat"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { AuthStorage } from "../src/core/auth-storage.ts"; import { createExtensionRuntime } from "../src/core/extensions/loader.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import type { ResourceLoader } from "../src/core/resource-loader.ts"; import { createAgentSession } from "../src/core/sdk.ts"; import { SessionManager } from "../src/core/session-manager.ts"; +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; function emptyResourceLoader(): ResourceLoader { return { @@ -39,7 +30,7 @@ function emptyResourceLoader(): ResourceLoader { }; } -describe("createAgentSession shared ModelRegistry (#1431)", () => { +describe("createAgentSession shared ModelRuntime (#1431)", () => { let tempDir: string; beforeEach(() => { @@ -54,11 +45,12 @@ describe("createAgentSession shared ModelRegistry (#1431)", () => { } }); - it("reuses a supplied modelRegistry and never constructs a fresh AuthStorage", async () => { + it("reuses a supplied modelRuntime and never constructs a fresh AuthStorage", async () => { const authStorage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "test-key" }, }); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); + const modelRegistry = await createModelRegistry(authStorage, tempDir); + const modelRuntime = getModelRuntime(modelRegistry); const createSpy = vi.spyOn(AuthStorage, "create"); const model = getModel("anthropic", "claude-sonnet-4-5"); @@ -66,14 +58,13 @@ describe("createAgentSession shared ModelRegistry (#1431)", () => { cwd: tempDir, agentDir: tempDir, model, - modelRegistry, + modelRuntime, sessionManager: SessionManager.inMemory(), resourceLoader: emptyResourceLoader(), }); - // The whole point: no fresh AuthStorage (and thus no extra contended - // auth.json reload) when a registry is already supplied. + // Supplying the runtime must avoid another credential-store construction. expect(createSpy).not.toHaveBeenCalled(); - expect(session.modelRegistry).toBe(modelRegistry); + expect(session.modelRuntime).toBe(modelRuntime); }); }); diff --git a/packages/coding-agent/test/sdk-stream-options.test.ts b/packages/coding-agent/test/sdk-stream-options.test.ts index 3d25169d1..a937cb170 100644 --- a/packages/coding-agent/test/sdk-stream-options.test.ts +++ b/packages/coding-agent/test/sdk-stream-options.test.ts @@ -1,39 +1,45 @@ -import { mkdirSync, mkdtempSync, rmSync } from "node:fs"; +import { mkdirSync, mkdtempSync, rmSync, writeFileSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; import { type Api, type AssistantMessage, + type AuthResult, createAssistantMessageEventStream, type Model, type SimpleStreamOptions, -} from "@earendil-works/pi-ai/compat"; -import { afterEach, beforeEach, describe, expect, it } from "vitest"; +} from "@earendil-works/pi-ai"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { ENV_CODEX_FAST_MODE } from "../src/config.ts"; import { AuthStorage } from "../src/core/auth-storage.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; import { createAgentSession } from "../src/core/sdk.ts"; import { SessionManager } from "../src/core/session-manager.ts"; -import { SettingsManager } from "../src/core/settings-manager.ts"; +import { type Settings, SettingsManager } from "../src/core/settings-manager.ts"; + +import { createModelRegistry, getModelRuntime } from "./model-runtime-test-utils.ts"; describe("createAgentSession stream options", () => { let tempDir: string; let cwd: string; let agentDir: string; + let previousCodexFastModeEnv: string | undefined; beforeEach(() => { - tempDir = mkdtempSync(join(tmpdir(), "atomic-sdk-stream-options-")); + tempDir = mkdtempSync(join(tmpdir(), "pi-sdk-stream-options-")); + previousCodexFastModeEnv = process.env[ENV_CODEX_FAST_MODE]; + delete process.env[ENV_CODEX_FAST_MODE]; cwd = join(tempDir, "project"); agentDir = join(tempDir, "agent"); mkdirSync(cwd, { recursive: true }); mkdirSync(agentDir, { recursive: true }); }); - afterEach(() => { - if (tempDir) { - rmSync(tempDir, { recursive: true, force: true }); - } + vi.restoreAllMocks(); + vi.unstubAllGlobals(); + if (tempDir) rmSync(tempDir, { recursive: true, force: true }); + if (previousCodexFastModeEnv === undefined) delete process.env[ENV_CODEX_FAST_MODE]; + else process.env[ENV_CODEX_FAST_MODE] = previousCodexFastModeEnv; }); - function createModel(api: Api): Model { return { id: "capture-model", @@ -46,6 +52,7 @@ describe("createAgentSession stream options", () => { cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, contextWindow: 128000, maxTokens: 4096, + headers: { "x-model": "model" }, }; } @@ -74,38 +81,48 @@ describe("createAgentSession stream options", () => { async function captureStreamOptions( api: Api, - settings: { httpIdleTimeoutMs?: number; websocketConnectTimeoutMs?: number }, + settings: Partial, requestOptions: SimpleStreamOptions = {}, + extensionSource?: string, + authResult?: AuthResult, ): Promise { const model = createModel(api); const settingsManager = SettingsManager.inMemory(settings); + if (extensionSource) { + const extensionsDir = join(agentDir, "extensions"); + mkdirSync(extensionsDir, { recursive: true }); + writeFileSync(join(extensionsDir, "headers.ts"), extensionSource); + } const authStorage = AuthStorage.create(join(agentDir, "auth.json")); - authStorage.setRuntimeApiKey(model.provider, "test-api-key"); - const modelRegistry = ModelRegistry.create(authStorage, join(agentDir, "models.json")); + await authStorage.modify(model.provider, async () => ({ type: "api_key", key: "test-api-key" })); + const modelRegistry = await createModelRegistry(authStorage, join(agentDir, "models.json")); let capturedOptions: SimpleStreamOptions | undefined; modelRegistry.registerProvider(model.provider, { api, - streamSimple: (_model, _context, providerOptions) => { + headers: { "x-provider": "provider" }, + streamSimple: (_requestModel, _context, providerOptions) => { capturedOptions = providerOptions; return createDoneStream(api); }, }); + const modelRuntime = getModelRuntime(modelRegistry); + if (authResult !== undefined) vi.spyOn(modelRuntime, "getAuth").mockResolvedValue(authResult); const sessionManager = SessionManager.inMemory(cwd); const { session } = await createAgentSession({ cwd, agentDir, model, - authStorage, - modelRegistry, + modelRuntime, settingsManager, sessionManager, }); try { - await session.agent.streamFunction(model, { messages: [] }, requestOptions); + const stream = await session.agent.streamFunction(model, { messages: [] }, requestOptions); + await stream.result(); return capturedOptions; } finally { session.dispose(); @@ -151,114 +168,166 @@ describe("createAgentSession stream options", () => { expect(options?.websocketConnectTimeoutMs).toBe(0); }); - it("dispatches with the credential-specific Copilot baseUrl", async () => { - const authStorage = AuthStorage.inMemory({ - "github-copilot": { - type: "oauth", - refresh: "github-token", - access: "tid=example;proxy-ep=proxy.enterprise.example.com;", - expires: Date.now() + 60_000, - }, + it("forwards provider retry settings", async () => { + const options = await captureStreamOptions("openai-completions", { + retry: { provider: { maxRetries: 2, maxRetryDelayMs: 3000 } }, }); - const modelRegistry = ModelRegistry.inMemory(authStorage); - const model = modelRegistry.getAll().find((candidate) => candidate.provider === "github-copilot")!; - let dispatchedBaseUrl: string | undefined; - modelRegistry.registerProvider(model.provider, { - api: model.api, - streamSimple: (requestModel) => { - dispatchedBaseUrl = requestModel.baseUrl; - return createDoneStream(model.api); + + expect(options?.maxRetries).toBe(2); + expect(options?.maxRetryDelayMs).toBe(3000); + }); + + it("runs before_provider_headers on assembled headers without forwarding the transform", async () => { + const options = await captureStreamOptions( + "openai-completions", + {}, + { headers: { "x-explicit": "explicit" } }, + `export default function (pi) { + pi.on("before_provider_headers", (event) => { + event.headers["x-hook"] = [ + event.headers["x-provider"], + event.headers["x-model"], + event.headers["x-explicit"], + ].join(":"); + }); + }`, + ); + + expect(options?.headers).toMatchObject({ + "x-provider": "provider", + "x-model": "model", + "x-explicit": "explicit", + "x-hook": "provider:model:explicit", + }); + expect(options).not.toHaveProperty("transformHeaders"); + }); + + it("preserves null credential headers through extension-provider dispatch", async () => { + const options = await captureStreamOptions( + "openai-completions", + {}, + {}, + undefined, + { + auth: { + apiKey: "credential-key", + headers: { Authorization: null, "x-api-key": null, "x-credential": "present" }, + }, }, + ); + + expect(options?.apiKey).toBe("credential-key"); + expect(options?.headers).toMatchObject({ + Authorization: null, + "x-api-key": null, + "x-credential": "present", }); + }); + + it("uses a credential-derived baseUrl for native Codex fast-mode dispatch", async () => { + const model: Model = { ...createModel("openai-responses"), provider: "openai" }; + const modelRuntime = getModelRuntime( + await createModelRegistry(AuthStorage.inMemory(), join(agentDir, "models.json")), + ); + vi.spyOn(modelRuntime, "getAuth").mockResolvedValue({ + auth: { apiKey: "credential-key", baseUrl: "https://credential.example/v1" }, + }); + let dispatchedUrl: string | undefined; + vi.stubGlobal( + "fetch", + vi.fn(async (input: RequestInfo | URL): Promise => { + dispatchedUrl = String(input); + const completed = { + type: "response.completed", + response: { + id: "resp_test", + status: "completed", + usage: { + input_tokens: 0, + input_tokens_details: { cached_tokens: 0 }, + output_tokens: 0, + total_tokens: 0, + }, + }, + }; + return new Response(`data: ${JSON.stringify(completed)}\n\ndata: [DONE]\n\n`, { + status: 200, + headers: { "content-type": "text/event-stream" }, + }); + }), + ); const { session } = await createAgentSession({ cwd, agentDir, model, - authStorage, - modelRegistry, - settingsManager: SettingsManager.inMemory(), + modelRuntime, + settingsManager: SettingsManager.inMemory({ codexFastMode: { chat: true, workflow: false } }), sessionManager: SessionManager.inMemory(cwd), }); try { - await session.agent.streamFunction(model, { messages: [] }); - expect(dispatchedBaseUrl).toBe("https://api.enterprise.example.com"); + const stream = await session.agent.streamFunction(model, { messages: [] }); + await stream.result(); + expect(dispatchedUrl).toMatch(/^https:\/\/credential\.example\/v1\//u); } finally { session.dispose(); - modelRegistry.unregisterProvider(model.provider); } }); - it("preserves an empty credential baseUrl during SDK dispatch", async () => { - const authStorage = AuthStorage.inMemory(); - const modelRegistry = ModelRegistry.inMemory(authStorage); - const model = modelRegistry.getAll()[0]!; - let dispatchedBaseUrl: string | undefined; - modelRegistry.getApiKeyAndHeaders = async () => ({ ok: true, apiKey: "key", baseUrl: "" }); - modelRegistry.registerProvider(model.provider, { + it("rejects authHeader providers before Codex fast-mode dispatch when credentials are missing", async () => { + const model: Model = { ...createModel("openai-responses"), provider: "openai" }; + const modelRuntime = getModelRuntime( + await createModelRegistry(AuthStorage.inMemory(), join(agentDir, "models.json")), + ); + vi.spyOn(modelRuntime, "getAuth").mockResolvedValue(undefined); + const streamSimple = vi.fn(() => createDoneStream(model.api)); + modelRuntime.registerProvider("openai", { api: model.api, - streamSimple: (requestModel) => { - dispatchedBaseUrl = requestModel.baseUrl; - return createDoneStream(model.api); - }, + baseUrl: model.baseUrl, + authHeader: true, + streamSimple, }); const { session } = await createAgentSession({ cwd, agentDir, model, - authStorage, - modelRegistry, - settingsManager: SettingsManager.inMemory(), + modelRuntime, + settingsManager: SettingsManager.inMemory({ codexFastMode: { chat: true, workflow: false } }), sessionManager: SessionManager.inMemory(cwd), }); try { - await session.agent.streamFunction(model, { messages: [] }); - expect(dispatchedBaseUrl).toBe(""); + await expect(session.agent.streamFunction(model, { messages: [] })).rejects.toThrow( + `No API key found for "${model.provider}"`, + ); + expect(streamSimple).not.toHaveBeenCalled(); } finally { session.dispose(); - modelRegistry.unregisterProvider(model.provider); + modelRuntime.unregisterProvider("openai"); } }); - it("resolves provider-owned null headers from a runtime API key through stream dispatch", async () => { - const previousAccount = process.env.CLOUDFLARE_ACCOUNT_ID; - const previousGateway = process.env.CLOUDFLARE_GATEWAY_ID; - process.env.CLOUDFLARE_ACCOUNT_ID = "account-id"; - process.env.CLOUDFLARE_GATEWAY_ID = "gateway-id"; - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("cloudflare-ai-gateway", "runtime-cf-key"); - const modelRegistry = ModelRegistry.inMemory(authStorage); - const model = modelRegistry.getAll().find((candidate) => candidate.provider === "cloudflare-ai-gateway")!; - let dispatchedHeaders: SimpleStreamOptions["headers"]; - modelRegistry.registerProvider(model.provider, { - api: model.api, - streamSimple: (_requestModel, _context, options) => { - dispatchedHeaders = options?.headers; - return createDoneStream(model.api); - }, - }); + it("rejects native Codex fast-mode dispatch when provider auth is unresolved", async () => { + const model: Model = { ...createModel("openai-responses"), provider: "openai" }; + const modelRuntime = getModelRuntime( + await createModelRegistry(AuthStorage.inMemory(), join(agentDir, "models.json")), + ); + vi.spyOn(modelRuntime, "getAuth").mockResolvedValue(undefined); const { session } = await createAgentSession({ cwd, agentDir, model, - authStorage, - modelRegistry, - settingsManager: SettingsManager.inMemory(), + modelRuntime, + settingsManager: SettingsManager.inMemory({ codexFastMode: { chat: true, workflow: false } }), sessionManager: SessionManager.inMemory(cwd), }); try { - await session.agent.streamFunction(model, { messages: [] }); - expect(dispatchedHeaders?.["cf-aig-authorization"]).toBe("Bearer runtime-cf-key"); - expect(dispatchedHeaders?.Authorization).toBeNull(); - expect(dispatchedHeaders?.["x-api-key"]).toBeNull(); + await expect(session.agent.streamFunction(model, { messages: [] })).rejects.toThrow( + `No API key found for "${model.provider}"`, + ); } finally { session.dispose(); - if (previousAccount === undefined) delete process.env.CLOUDFLARE_ACCOUNT_ID; - else process.env.CLOUDFLARE_ACCOUNT_ID = previousAccount; - if (previousGateway === undefined) delete process.env.CLOUDFLARE_GATEWAY_ID; - else process.env.CLOUDFLARE_GATEWAY_ID = previousGateway; } }); }); diff --git a/packages/coding-agent/test/suite/agent-session-auth.test.ts b/packages/coding-agent/test/suite/agent-session-auth.test.ts index cef8de932..99fea4da1 100644 --- a/packages/coding-agent/test/suite/agent-session-auth.test.ts +++ b/packages/coding-agent/test/suite/agent-session-auth.test.ts @@ -1,10 +1,10 @@ -import { afterEach, describe, expect, it } from "vitest"; +import { afterEach, describe, expect, it, vi } from "vitest"; import { AgentSessionRuntime, type CreateAgentSessionRuntimeFactory, } from "../../src/core/agent-session-runtime.ts"; -import { createRpcCommandHandler } from "../../src/modes/rpc/rpc-command-handler.ts"; import { createHarness, type Harness } from "./harness.ts"; +import { createRpcCommandHandler } from "../../src/modes/rpc/rpc-command-handler.ts"; const createRuntime = (async () => { throw new Error("not used"); @@ -25,11 +25,11 @@ describe("provider-metadata authentication runtime", () => { while (harnesses.length > 0) harnesses.pop()?.cleanup(); }); - it("logs in through ModelRegistry provider metadata and persists the acquired credential", async () => { + it("logs in through provider-owned runtime metadata and persists the acquired credential", async () => { const harness = await createHarness({ withConfiguredAuth: false }); harnesses.push(harness); const providerId = "openrouter"; - harness.session.modelRegistry.registerProvider(providerId, { + harness.session.modelRuntime.registerProvider(providerId, { oauth: { name: "Faux subscription", login: async () => ({ @@ -41,6 +41,7 @@ describe("provider-metadata authentication runtime", () => { getApiKey: (credentials) => credentials.access, }, }); + vi.spyOn(harness.session.modelRuntime, "refresh").mockResolvedValue({ ok: true, providers: [] }); await runtimeFor(harness).loginOAuthProvider(providerId, { onAuth: () => {}, @@ -49,7 +50,7 @@ describe("provider-metadata authentication runtime", () => { onSelect: async () => undefined, }); - expect(harness.authStorage.get(providerId)).toMatchObject({ + expect(await harness.authStorage.read(providerId)).toMatchObject({ type: "oauth", access: "access-token", refresh: "refresh-token", @@ -85,6 +86,6 @@ describe("provider-metadata authentication runtime", () => { command: "save_provider_credential", success: true, }); - expect(harness.authStorage.get(harness.getModel().provider)).toEqual(credential); + expect(await harness.authStorage.read(harness.getModel().provider)).toEqual(credential); }); }); diff --git a/packages/coding-agent/test/suite/agent-session-runtime-01.suite.ts b/packages/coding-agent/test/suite/agent-session-runtime-01.suite.ts index 25ef0f883..d22431f04 100644 --- a/packages/coding-agent/test/suite/agent-session-runtime-01.suite.ts +++ b/packages/coding-agent/test/suite/agent-session-runtime-01.suite.ts @@ -51,7 +51,7 @@ describe("AgentSessionRuntime characterization", () => { faux.setResponses([fauxAssistantMessage("one"), fauxAssistantMessage("two"), fauxAssistantMessage("three")]); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await authStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); const runtimeOptions = { agentDir: tempDir, @@ -335,7 +335,7 @@ describe("AgentSessionRuntime characterization", () => { faux.setResponses([fauxAssistantMessage("one"), fauxAssistantMessage("two"), fauxAssistantMessage("three")]); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await authStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); const runtimeOptions = { agentDir: tempDir, diff --git a/packages/coding-agent/test/suite/agent-session-runtime-02.suite.ts b/packages/coding-agent/test/suite/agent-session-runtime-02.suite.ts index 177fac6b1..bae17b37f 100644 --- a/packages/coding-agent/test/suite/agent-session-runtime-02.suite.ts +++ b/packages/coding-agent/test/suite/agent-session-runtime-02.suite.ts @@ -51,7 +51,7 @@ describe("AgentSessionRuntime characterization", () => { faux.setResponses([fauxAssistantMessage("one"), fauxAssistantMessage("two"), fauxAssistantMessage("three")]); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await authStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); const runtimeOptions = { agentDir: tempDir, @@ -126,7 +126,7 @@ describe("AgentSessionRuntime characterization", () => { mkdirSync(secondDir, { recursive: true }); const { runtime, faux, tempDir } = await createRuntimeForTest(() => {}, { cwd: firstDir }); const otherAuthStorage = AuthStorage.inMemory(); - otherAuthStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await otherAuthStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); const otherRuntimeOptions = { agentDir: tempDir, authStorage: otherAuthStorage, @@ -198,7 +198,7 @@ describe("AgentSessionRuntime characterization", () => { const otherDir = join(tempDir, "other"); mkdirSync(otherDir, { recursive: true }); const otherAuthStorage = AuthStorage.inMemory(); - otherAuthStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await otherAuthStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); const otherRuntimeOptions = { agentDir: tempDir, authStorage: otherAuthStorage, diff --git a/packages/coding-agent/test/suite/harness.ts b/packages/coding-agent/test/suite/harness.ts index 233f445ce..7afae1c76 100644 --- a/packages/coding-agent/test/suite/harness.ts +++ b/packages/coding-agent/test/suite/harness.ts @@ -13,7 +13,8 @@ import { AgentSession, type AgentSessionEvent } from "../../src/core/agent-sessi import { AuthStorage } from "../../src/core/auth-storage.ts"; import type { ExtensionRunner } from "../../src/core/extensions/index.ts"; import { convertToLlm } from "../../src/core/messages.ts"; -import { ModelRegistry } from "../../src/core/model-registry.ts"; +import { InMemoryCodingAgentModelsStore } from "../../src/core/models-store.ts"; +import { ModelRuntime } from "../../src/core/model-runtime.ts"; import { SessionManager } from "../../src/core/session-manager.ts"; import type { Settings } from "../../src/core/settings-manager.ts"; import { SettingsManager } from "../../src/core/settings-manager.ts"; @@ -111,11 +112,11 @@ export async function createHarness(options: HarnessOptions = {}): Promise ({ type: "api_key", key: "faux-key" })); } - const modelRegistry = ModelRegistry.inMemory(authStorage); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null, modelsStore: new InMemoryCodingAgentModelsStore() }); if (withConfiguredAuth) { - modelRegistry.registerProvider(model.provider, { + modelRuntime.registerProvider(model.provider, { baseUrl: model.baseUrl, apiKey: "faux-key", api: fauxProvider.api, @@ -177,7 +178,7 @@ export async function createHarness(options: HarnessOptions = {}): Promise { models: [{ id: "faux-1", reasoning: false }], }); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await authStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); const createRuntime: CreateAgentSessionRuntimeFactory = async ({ cwd, sessionManager, sessionStartEvent }) => { const services = await createAgentSessionServices({ diff --git a/packages/coding-agent/test/suite/regressions/2860-replaced-session-context.test.ts b/packages/coding-agent/test/suite/regressions/2860-replaced-session-context.test.ts index b6f7002e8..5d0afdaa8 100644 --- a/packages/coding-agent/test/suite/regressions/2860-replaced-session-context.test.ts +++ b/packages/coding-agent/test/suite/regressions/2860-replaced-session-context.test.ts @@ -45,7 +45,7 @@ describe("regression #2860: replaced session callbacks", () => { faux.setResponses(responses.map((response) => fauxAssistantMessage(response))); const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + await authStorage.modify(faux.getModel().provider, async () => ({ type: "api_key", key: "faux-key" })); const createRuntime: CreateAgentSessionRuntimeFactory = async ({ cwd, sessionManager, sessionStartEvent }) => { const services = await createAgentSessionServices({ diff --git a/packages/coding-agent/test/suite/regressions/3217-scoped-model-order.test.ts b/packages/coding-agent/test/suite/regressions/3217-scoped-model-order.test.ts index 6e453a356..9fc1de48a 100644 --- a/packages/coding-agent/test/suite/regressions/3217-scoped-model-order.test.ts +++ b/packages/coding-agent/test/suite/regressions/3217-scoped-model-order.test.ts @@ -83,7 +83,7 @@ describe("issue #3217 scoped model ordering", () => { createFakeTui(), modelOne, harness.settingsManager, - harness.session.modelRegistry, + harness.session.modelRuntime, [{ model: modelTwo }, { model: modelOne }, { model: modelThree }], () => {}, () => {}, diff --git a/packages/coding-agent/test/test-harness.test.ts b/packages/coding-agent/test/test-harness.test.ts index c314a1be0..27801a522 100644 --- a/packages/coding-agent/test/test-harness.test.ts +++ b/packages/coding-agent/test/test-harness.test.ts @@ -17,7 +17,7 @@ describe("test harness", () => { }); it("simple text response", async () => { - harness = createHarness({ responses: ["hello world"] }); + harness = await createHarness({ responses: ["hello world"] }); await harness.session.prompt("hi"); @@ -32,7 +32,7 @@ describe("test harness", () => { }); it("response sequence", async () => { - harness = createHarness({ responses: ["first", "second", "third"] }); + harness = await createHarness({ responses: ["first", "second", "third"] }); await harness.session.prompt("a"); await harness.session.prompt("b"); @@ -60,7 +60,7 @@ describe("test harness", () => { }, }; - harness = createHarness({ + harness = await createHarness({ responses: [{ toolCalls: [{ name: "echo", args: { text: "hi" } }] }, "done after tool"], tools: [echoTool], baseToolsOverride: { echo: echoTool }, @@ -76,7 +76,7 @@ describe("test harness", () => { }); it("error response", async () => { - harness = createHarness({ + harness = await createHarness({ responses: [{ error: "something broke" }], }); @@ -89,7 +89,7 @@ describe("test harness", () => { }); it("retry on transient error", async () => { - harness = createHarness({ + harness = await createHarness({ responses: [{ error: "overloaded_error" }, "recovered"], settings: { retry: { enabled: true, maxRetries: 3, baseDelayMs: 1 } }, }); @@ -107,7 +107,7 @@ describe("test harness", () => { }); it("custom usage numbers", async () => { - harness = createHarness({ + harness = await createHarness({ responses: [{ text: "big response", usage: { input: 100000, output: 5000 } }], }); @@ -119,7 +119,7 @@ describe("test harness", () => { }); it("event capture", async () => { - harness = createHarness({ responses: ["hello"] }); + harness = await createHarness({ responses: ["hello"] }); await harness.session.prompt("hi"); @@ -134,7 +134,7 @@ describe("test harness", () => { }); it("context capture", async () => { - harness = createHarness({ responses: ["reply"] }); + harness = await createHarness({ responses: ["reply"] }); await harness.session.prompt("my question"); @@ -145,7 +145,7 @@ describe("test harness", () => { }); it("wraps around when more calls than responses", async () => { - harness = createHarness({ responses: ["a", "b"] }); + harness = await createHarness({ responses: ["a", "b"] }); await harness.session.prompt("1"); await harness.session.prompt("2"); @@ -161,7 +161,7 @@ describe("test harness", () => { }); it("streams text deltas", async () => { - harness = createHarness({ responses: ["hello world"] }); + harness = await createHarness({ responses: ["hello world"] }); await harness.session.prompt("hi"); @@ -175,7 +175,7 @@ describe("test harness", () => { }); it("streams thinking deltas", async () => { - harness = createHarness({ + harness = await createHarness({ responses: [{ thinking: "let me think about this", text: "answer" }], }); @@ -203,7 +203,7 @@ describe("test harness", () => { execute: async () => ({ content: [{ type: "text", text: "echoed" }], details: {} }), }; - harness = createHarness({ + harness = await createHarness({ responses: [{ toolCalls: [{ name: "echo", args: { text: "hi" } }] }, "done"], tools: [echoTool], baseToolsOverride: { echo: echoTool }, @@ -230,7 +230,7 @@ describe("test harness", () => { execute: async () => ({ content: [{ type: "text", text: "echoed" }], details: {} }), }; - harness = createHarness({ + harness = await createHarness({ responses: [ { thinking: "hmm", @@ -310,7 +310,7 @@ describe("test harness", () => { }); it("session persistence works", async () => { - harness = createHarness({ responses: ["persisted"] }); + harness = await createHarness({ responses: ["persisted"] }); await harness.session.prompt("hi"); diff --git a/packages/coding-agent/test/test-harness.ts b/packages/coding-agent/test/test-harness.ts index 48df9ef18..14563f834 100644 --- a/packages/coding-agent/test/test-harness.ts +++ b/packages/coding-agent/test/test-harness.ts @@ -28,7 +28,7 @@ import type { import { createAssistantMessageEventStream } from "@earendil-works/pi-ai/compat"; import { AgentSession, type AgentSessionEvent } from "../src/core/agent-session.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 { SessionManager } from "../src/core/session-manager.ts"; import type { Settings } from "../src/core/settings-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; @@ -361,11 +361,11 @@ function createTempDir(): string { return tempDir; } -function createHarnessWithResourceLoader( +async function createHarnessWithResourceLoader( options: HarnessOptions, resourceLoader: ResourceLoader, tempDir: string, -): Harness { +): Promise { const baseModel = options.model ?? fauxModel; const model: Model = options.contextWindow ? { ...baseModel, contextWindow: options.contextWindow } : baseModel; @@ -389,15 +389,20 @@ function createHarnessWithResourceLoader( } const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - authStorage.setRuntimeApiKey(model.provider, "faux-key"); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); - + await authStorage.modify(model.provider, async () => ({ type: "api_key", key: "faux-key" })); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); + modelRuntime.registerProvider(model.provider, { + baseUrl: model.baseUrl, + apiKey: "faux-key", + api: model.api, + models: [model], + }); const session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader, baseToolsOverride: options.baseToolsOverride, }); @@ -429,7 +434,7 @@ function createHarnessWithResourceLoader( }; } -export function createHarness(options: HarnessOptions = {}): Harness { +export async function createHarness(options: HarnessOptions = {}): Promise { if (options.extensionFactories?.length) { throw new Error("createHarness does not support extensionFactories. Use createHarnessWithExtensions()."); } diff --git a/packages/coding-agent/test/utilities.ts b/packages/coding-agent/test/utilities.ts index b28108198..ea1c11974 100644 --- a/packages/coding-agent/test/utilities.ts +++ b/packages/coding-agent/test/utilities.ts @@ -8,11 +8,11 @@ import { join } from "node:path"; import { Agent } from "@earendil-works/pi-agent-core"; import { getModel, streamSimple } from "@earendil-works/pi-ai/compat"; import { AgentSession } from "../src/core/agent-session.ts"; -import { AuthStorage } from "../src/core/auth-storage.ts"; +import { readStoredCredential, AuthStorage } from "../src/core/auth-storage.ts"; import { createEventBus } from "../src/core/event-bus.ts"; import type { Extension, ExtensionFactory, LoadExtensionsResult } from "../src/core/extensions/index.ts"; import { createExtensionRuntime, loadExtensionFromFactory } from "../src/core/extensions/loader.ts"; -import { ModelRegistry } from "../src/core/model-registry.ts"; +import { ModelRuntime } from "../src/core/model-runtime.ts"; import type { ResourceLoader } from "../src/core/resource-loader.ts"; import { SessionManager } from "../src/core/session-manager.ts"; import { SettingsManager } from "../src/core/settings-manager.ts"; @@ -39,14 +39,15 @@ const AUTH_PATH = join(homedir(), ".pi", "agent", "auth.json"); * */ export async function resolveApiKey(provider: string): Promise { - return AuthStorage.create(AUTH_PATH).getApiKey(provider); + const credential = await AuthStorage.create(AUTH_PATH).read(provider); + return credential?.type === "api_key" ? credential.key : credential?.access; } /** * Check if a provider has credentials in ~/.pi/agent/auth.json */ export function hasAuthForProvider(provider: string): boolean { - return AuthStorage.create(AUTH_PATH).has(provider); + return readStoredCredential(provider, AUTH_PATH) !== undefined; } /** Path to the real pi agent config directory */ @@ -167,7 +168,7 @@ export function createTestResourceLoader(options: CreateTestResourceLoaderOption * Create an AgentSession for testing with proper setup and cleanup. * Use this for e2e tests that need real LLM calls. */ -export function createTestSession(options: TestSessionOptions = {}): TestSessionContext { +export async function createTestSession(options: TestSessionOptions = {}): Promise { const tempDir = join(tmpdir(), `pi-test-${Date.now()}-${Math.random().toString(36).slice(2)}`); mkdirSync(tempDir, { recursive: true }); @@ -190,14 +191,14 @@ export function createTestSession(options: TestSessionOptions = {}): TestSession } const authStorage = AuthStorage.create(join(tempDir, "auth.json")); - const modelRegistry = ModelRegistry.create(authStorage, tempDir); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const session = new AgentSession({ agent, sessionManager, settingsManager, cwd: tempDir, - modelRegistry, + modelRuntime, resourceLoader: createTestResourceLoader(), }); diff --git a/packages/mcp/sampling-handler.ts b/packages/mcp/sampling-handler.ts index 9f45726af..bcb189c4f 100644 --- a/packages/mcp/sampling-handler.ts +++ b/packages/mcp/sampling-handler.ts @@ -152,11 +152,7 @@ async function resolveSamplingModel( errors.push(`${model.provider}/${model.id}: ${auth.error}`); continue; } - return { - model: auth.baseUrl === undefined ? model : { ...model, baseUrl: auth.baseUrl }, - apiKey: auth.apiKey, - headers: auth.headers, - }; + return { model, apiKey: auth.apiKey, headers: auth.headers }; } if (errors.length > 0) { diff --git a/packages/workflows/src/extension/extension-runtime-state.ts b/packages/workflows/src/extension/extension-runtime-state.ts index f4345ffd5..78fc4dfd3 100644 --- a/packages/workflows/src/extension/extension-runtime-state.ts +++ b/packages/workflows/src/extension/extension-runtime-state.ts @@ -1,12 +1,9 @@ -import type { CreateAgentSessionOptions } from "@bastani/atomic"; import type { StageAdapters } from "../runs/foreground/stage-runner.js"; import type { SessionManager } from "../shared/persistence-restore.js"; import type { RunSnapshot } from "../shared/store-types.js"; import type { WorkflowExecutionPolicy, WorkflowMcpPort, - WorkflowModelCatalogPort, - WorkflowModelInfo, WorkflowPersistencePort, WorkflowRuntimeConfig, } from "../shared/types.js"; @@ -43,6 +40,7 @@ import { workflowReloadDiagnostics, type WorkflowReloadReport, } from "./workflow-reload-report.js"; +import { workflowModelCatalogFromContext } from "./workflow-model-catalog.js"; export interface WorkflowExtensionRuntimeState { persistenceRef: { current: WorkflowPersistencePort | undefined }; @@ -171,27 +169,6 @@ export function createWorkflowExtensionRuntimeState( }, }; - function workflowModelCatalogFromContext(ctx?: PiModelContext): WorkflowModelCatalogPort | undefined { - if (ctx?.modelRegistry === undefined && ctx?.model === undefined) return undefined; - return { - listModels: async (): Promise => { - const available = ctx.modelRegistry?.getAvailable() ?? (ctx.model === undefined ? [] : [ctx.model]); - return available.map((model) => ({ - provider: String(model.provider), - id: model.id, - fullId: `${String(model.provider)}/${model.id}`, - model: model as NonNullable, - })); - }, - ...(ctx.model !== undefined - ? { - currentModel: ctx.model as NonNullable, - preferredProvider: String(ctx.model.provider), - } - : {}), - }; - } - function runtimeForContext(ctx?: PiModelContext): ExtensionRuntime { const models = workflowModelCatalogFromContext(ctx); if (models === undefined) return runtimeProxy; diff --git a/packages/workflows/src/extension/workflow-model-catalog.ts b/packages/workflows/src/extension/workflow-model-catalog.ts new file mode 100644 index 000000000..b1f24b982 --- /dev/null +++ b/packages/workflows/src/extension/workflow-model-catalog.ts @@ -0,0 +1,26 @@ +import type { CreateAgentSessionOptions } from "@bastani/atomic"; +import type { WorkflowModelCatalogPort, WorkflowModelInfo } from "../shared/types.js"; +import type { PiModelContext } from "./public-types.js"; + +export function workflowModelCatalogFromContext( + ctx?: PiModelContext, +): WorkflowModelCatalogPort | undefined { + if (ctx?.modelRegistry === undefined && ctx?.model === undefined) return undefined; + return { + listModels: async (): Promise => { + const available = ctx.modelRegistry?.getAvailable() ?? (ctx.model === undefined ? [] : [ctx.model]); + return available.map((model) => ({ + provider: String(model.provider), + id: model.id, + fullId: `${String(model.provider)}/${model.id}`, + model: model as NonNullable, + })); + }, + ...(ctx.model !== undefined + ? { + currentModel: ctx.model as NonNullable, + preferredProvider: String(ctx.model.provider), + } + : {}), + }; +} diff --git a/packages/workflows/src/runs/foreground/stage-runner-controller.ts b/packages/workflows/src/runs/foreground/stage-runner-controller.ts index d9b371c79..1b8e6fa77 100644 --- a/packages/workflows/src/runs/foreground/stage-runner-controller.ts +++ b/packages/workflows/src/runs/foreground/stage-runner-controller.ts @@ -35,7 +35,7 @@ export class StageSessionController { private candidatesPromise: Promise | undefined; private activeCandidateIndex: number | undefined; private selectedModel: string | undefined; - private sharedModelRegistry: CreateAgentSessionOptions["modelRegistry"]; + private sharedModelRuntime: CreateAgentSessionOptions["modelRuntime"]; private sharedOrchestrationContext: CreateAgentSessionOptions["orchestrationContext"]; private resumeCurrentSession = false; private readonly modelAttempts: WorkflowModelAttempt[] = []; @@ -299,7 +299,7 @@ export class StageSessionController { candidate, restoreSavedModel: resumeOptions?.restoreSavedModel, reattachSessionFile: this.reattachSessionFile, - sharedModelRegistry: this.sharedModelRegistry, + sharedModelRuntime: this.sharedModelRuntime, }); const created = this.opts.adapters.agentSession ? await this.opts.adapters.agentSession.create( @@ -325,9 +325,9 @@ export class StageSessionController { if (this.generationSealed) result.session.sealWorkflowStageGeneration?.(); this.replacement.adopt(result.session); this.session = result.session; - if (this.sharedModelRegistry === undefined) { - const withRegistry = result.session as Partial>; - if (withRegistry.modelRegistry !== undefined) this.sharedModelRegistry = withRegistry.modelRegistry; + if (this.sharedModelRuntime === undefined) { + const withRegistry = result.session as Partial>; + if (withRegistry.modelRuntime !== undefined) this.sharedModelRuntime = withRegistry.modelRuntime; } this.sessionSettingsManager = result.settingsManager ?? result.session.settingsManager; if (this.pendingThinkingLevel !== undefined) result.session.setThinkingLevel(this.pendingThinkingLevel); diff --git a/packages/workflows/src/runs/foreground/stage-runner-session-options.ts b/packages/workflows/src/runs/foreground/stage-runner-session-options.ts index e2d9cc183..96e9a0a12 100644 --- a/packages/workflows/src/runs/foreground/stage-runner-session-options.ts +++ b/packages/workflows/src/runs/foreground/stage-runner-session-options.ts @@ -10,7 +10,7 @@ interface StageSessionOptionsInput { readonly candidate: WorkflowResolvedModelCandidate | undefined; readonly restoreSavedModel?: boolean; readonly reattachSessionFile: string | undefined; - readonly sharedModelRegistry: CreateAgentSessionOptions["modelRegistry"]; + readonly sharedModelRuntime: CreateAgentSessionOptions["modelRuntime"]; } export function buildStageSessionOptions(input: StageSessionOptionsInput): StageOptions | undefined { @@ -31,8 +31,8 @@ export function buildStageSessionOptions(input: StageSessionOptionsInput): Stage options.context = undefined; options.forkFromSessionFile = undefined; } - if (input.sharedModelRegistry !== undefined && options.modelRegistry === undefined) { - options.modelRegistry = input.sharedModelRegistry; + if (input.sharedModelRuntime !== undefined && options.modelRuntime === undefined) { + options.modelRuntime = input.sharedModelRuntime; } return Object.keys(options).length === 0 ? undefined : options; } diff --git a/packages/workflows/src/runs/foreground/stage-runner-session.ts b/packages/workflows/src/runs/foreground/stage-runner-session.ts index c9bbff6ca..1637efa65 100644 --- a/packages/workflows/src/runs/foreground/stage-runner-session.ts +++ b/packages/workflows/src/runs/foreground/stage-runner-session.ts @@ -35,11 +35,11 @@ export async function disposeStageSession(current: StageSessionRuntime | undefin export function asAgentSession(activeSession: StageSessionRuntime | undefined): AgentSession | undefined { if (!activeSession) return undefined; - const candidate = activeSession as StageSessionRuntime & Partial>; + const candidate = activeSession as StageSessionRuntime & Partial>; if ( candidate.state !== undefined && candidate.sessionManager !== undefined && - candidate.modelRegistry !== undefined && + candidate.modelRuntime !== undefined && typeof candidate.getContextUsage === "function" ) { return candidate as AgentSession; diff --git a/test/integration/compaction-fallback-rungs.test.ts b/test/integration/compaction-fallback-rungs.test.ts index 03cc16809..c5689be4b 100644 --- a/test/integration/compaction-fallback-rungs.test.ts +++ b/test/integration/compaction-fallback-rungs.test.ts @@ -18,7 +18,7 @@ function boundary(manager: { getBranch(): Array<{ type: string }> }): Compaction test("rate-limit rescue: a healthy fallback keeps the boundary at rung planned and the UI quiet", async () => { const { streamFn, calls } = plannerScript({ anthropic: [THROTTLED], openai: [{ text: "3,9\n" }] }); - const built = createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); + const built = await createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); try { const sessionModelBefore = built.session.model; const result = await built.session.compact({ preserve_recent: 2 }); @@ -51,7 +51,7 @@ test("starvation rescue: exactly two planner requests, both at the inherited rea anthropic: [{ text: "", stopReason: "length", reasoningTokens: 4096 }], openai: [{ text: "3,9\n" }], }); - const built = createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"], thinkingLevel: "high" }); + const built = await createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"], thinkingLevel: "high" }); try { const result = await built.session.compact({ preserve_recent: 2 }); assert.equal(result.rung, "planned"); @@ -68,7 +68,7 @@ test("overflow trimming retries the same model on a smaller region before advanc const { streamFn, calls } = plannerScript({ anthropic: [{ errorMessage: "prompt is too long: context_length_exceeded" }, { text: "1,4\n" }], }); - const built = createRungSession({ streamFn }); + const built = await createRungSession({ streamFn }); try { const result = await built.session.compact({ preserve_recent: 2 }); assert.equal(result.rung, "planned"); @@ -90,7 +90,7 @@ test("overflow trimming continues until no smaller region exists, then advances" anthropic: [OVERFLOW], openai: [{ text: "3,9\n" }], }); - const built = createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); + const built = await createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); try { const result = await built.session.compact({ preserve_recent: 2 }); assert.equal(result.rung, "planned"); @@ -121,7 +121,7 @@ test("a trimmed region that succeeds past the old three-trim cap is still accept anthropic: [OVERFLOW, OVERFLOW, OVERFLOW, OVERFLOW, { text: "1,2\n" }], openai: [{ text: "3,9\n" }], }); - const built = createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); + const built = await createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); try { const result = await built.session.compact({ preserve_recent: 2 }); assert.equal(result.rung, "planned"); @@ -136,7 +136,7 @@ test("a trimmed region that succeeds past the old three-trim cap is still accept test("each trimmed view derives its own keep target from compression_ratio", async () => { const { streamFn, calls } = plannerScript({ anthropic: [OVERFLOW, OVERFLOW, { text: "1,2\n" }] }); - const built = createRungSession({ streamFn }); + const built = await createRungSession({ streamFn }); try { await built.session.compact({ preserve_recent: 2 }); assert.equal(calls.length, 3); @@ -167,7 +167,7 @@ test("a 20-line region halved to 10 still asks for deletions", async () => { // The exact shape from the finding: ratio 0.5 over 20 lines asks to keep 10; // carrying that target into the 10-line view would request zero deletions. const { streamFn, calls } = plannerScript({ anthropic: [OVERFLOW, { text: "1,2\n" }] }); - const built = createRungSession({ streamFn, turns: 4 }); + const built = await createRungSession({ streamFn, turns: 4 }); try { await built.session.compact({ preserve_recent: 2 }); assert.ok(calls.length >= 2); @@ -183,7 +183,7 @@ test("a 20-line region halved to 10 still asks for deletions", async () => { test("manual honesty: every model rate limited writes no boundary and leaves the session usable", async () => { const { streamFn, calls } = plannerScript({ anthropic: [THROTTLED], openai: [THROTTLED] }); - const built = createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); + const built = await createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); try { await assert.rejects(() => built.session.compact({ preserve_recent: 2 }), /429 Too Many Requests/); assert.deepEqual(calls.map((call) => call.provider), ["anthropic", "openai"]); @@ -211,7 +211,7 @@ test("manual honesty: every model rate limited writes no boundary and leaves the test("threshold auto-compaction is recoverable and cannot clear context", async () => { const { streamFn } = plannerScript({ anthropic: [THROTTLED] }); - const built = createRungSession({ streamFn }); + const built = await createRungSession({ streamFn }); try { const runAutoCompaction = ( built.session as unknown as { @@ -230,7 +230,7 @@ test("threshold auto-compaction is recoverable and cannot clear context", async test("overflow recovery is load-bearing and completes on the fresh rung when every model fails", async () => { const { streamFn } = plannerScript({ anthropic: [THROTTLED], openai: [THROTTLED] }); - const built = createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); + const built = await createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"] }); try { const runAutoCompaction = ( built.session as unknown as { @@ -258,7 +258,7 @@ test("overflow recovery is load-bearing and completes on the fresh rung when eve test("post-tool preflight is load-bearing: it completes without scheduling a continuation", async () => { const { streamFn } = plannerScript({ anthropic: [THROTTLED] }); - const built = createRungSession({ streamFn, contextWindow: 4_000, reserveTokens: 3_999, turns: 12 }); + const built = await createRungSession({ streamFn, contextWindow: 4_000, reserveTokens: 3_999, turns: 12 }); try { const preflight = ( built.session as unknown as { diff --git a/test/integration/compaction-manual-honesty.test.ts b/test/integration/compaction-manual-honesty.test.ts index 1decf6828..03ac5b03e 100644 --- a/test/integration/compaction-manual-honesty.test.ts +++ b/test/integration/compaction-manual-honesty.test.ts @@ -27,7 +27,7 @@ function diagnostics(directory: string): string[] { test("manual compaction with every model rate limited reports a real diagnostic path", async () => { const directory = mkdtempSync(join(tmpdir(), "compaction-manual-honesty-")); const { streamFn, calls } = plannerScript({ anthropic: [THROTTLED], openai: [THROTTLED] }); - const built = createRungSession({ + const built = await createRungSession({ streamFn, fallbackModels: ["openai/gpt-5.1"], sessionDir: directory, @@ -79,7 +79,7 @@ test("an over-limit post-tool context with a tiny region persists a fresh bounda // Two seeded turns leave fewer than the planner minimum of compactable lines, // while the reported usage puts the whole context past the hard input limit. const { streamFn, calls } = plannerScript({ default: [{ text: "1,2\n" }] }); - const built = createRungSession({ + const built = await createRungSession({ streamFn, turns: 2, contextWindow: 4_000, @@ -122,7 +122,7 @@ test("a fitting post-tool threshold crossing with a tiny region is a safe no-op" // Clearing a sub-minimum region here would destroy conversation for nothing; // the follow-up provider request can be sent unchanged. const { streamFn, calls } = plannerScript({ default: [{ text: "1,2\n" }] }); - const built = createRungSession({ + const built = await createRungSession({ streamFn, turns: 2, contextWindow: 1_000_000, @@ -161,7 +161,7 @@ test("a fitting post-tool threshold crossing with a tiny region is a safe no-op" test("an over-hard-limit post-tool crossing with a tiny region reaches fresh", async () => { const { streamFn, calls } = plannerScript({ default: [{ text: "1,2\n" }] }); - const built = createRungSession({ + const built = await createRungSession({ streamFn, turns: 2, contextWindow: 4_000, @@ -190,7 +190,7 @@ test("an over-hard-limit post-tool crossing with a tiny region reaches fresh", a test("real overflow recovery keeps fresh reachable for a tiny region", async () => { const { streamFn, calls } = plannerScript({ default: [{ text: "1,2\n" }] }); - const built = createRungSession({ streamFn, turns: 2, contextWindow: 1_000_000 }); + const built = await createRungSession({ streamFn, turns: 2, contextWindow: 1_000_000 }); try { const runAutoCompaction = ( built.session as unknown as { @@ -210,7 +210,7 @@ test("real overflow recovery keeps fresh reachable for a tiny region", async () test("a small recoverable region is still refused, never cleared", async () => { const { streamFn, calls } = plannerScript({ default: [{ text: "1,2\n" }] }); - const built = createRungSession({ streamFn, turns: 2 }); + const built = await createRungSession({ streamFn, turns: 2 }); try { const apply = ( built.session as unknown as { @@ -238,7 +238,7 @@ test("a caller cannot inject load_bearing urgency through the public compact doo // `session.compact()` projects only compaction parameters, so a widened or // cast object cannot smuggle urgency in and reach the destructive rung. const { streamFn, calls } = plannerScript({ anthropic: [THROTTLED] }); - const built = createRungSession({ streamFn }); + const built = await createRungSession({ streamFn }); try { const injected = { preserve_recent: 2, urgency: "load_bearing" } as unknown as { preserve_recent: number }; await assert.rejects(() => built.session.compact(injected), /429 Too Many Requests/); diff --git a/test/integration/compaction-post-tool-chain.test.ts b/test/integration/compaction-post-tool-chain.test.ts index 5c6384d7f..1a7b3911a 100644 --- a/test/integration/compaction-post-tool-chain.test.ts +++ b/test/integration/compaction-post-tool-chain.test.ts @@ -18,7 +18,7 @@ import { Type } from "typebox"; import { AgentSession, type AgentSessionEvent } from "../../packages/coding-agent/src/core/agent-session.js"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.js"; import { convertToLlm } from "../../packages/coding-agent/src/core/messages.js"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.js"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.js"; import { SessionManager } from "../../packages/coding-agent/src/core/session-manager.js"; import { SettingsManager } from "../../packages/coding-agent/src/core/settings-manager.js"; import { RANGE_PLANNER_SYSTEM_PROMPT } from "../../packages/coding-agent/src/core/compaction/range-planner.js"; @@ -102,8 +102,10 @@ test("post-tool preflight: every model fails, the turn still completes on the fr return realContinue(...args); }) as typeof agent.continue; - const authStorage = AuthStorage.inMemory(); - for (const provider of ["anthropic", "openai", "google"]) authStorage.setRuntimeApiKey(provider, `${provider}-key`); + const authStorage = AuthStorage.inMemory(Object.fromEntries( + ["anthropic", "openai", "google"].map((provider) => [provider, { type: "api_key" as const, key: `${provider}-key` }]), + )); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const settingsManager = SettingsManager.inMemory(); settingsManager.applyOverrides({ retry: { enabled: false, maxRetries: 0, baseDelayMs: 0 }, @@ -114,7 +116,7 @@ test("post-tool preflight: every model fails, the turn still completes on the fr sessionManager: manager, settingsManager, cwd: process.cwd(), - modelRegistry: ModelRegistry.create(authStorage), + modelRuntime, resourceLoader: createTestResourceLoader(), fallbackModels: ["openai/gpt-5.1", "google/gemini-2.5-pro"], baseToolsOverride: { large_result: largeResultTool }, diff --git a/test/integration/compaction-rung-session.ts b/test/integration/compaction-rung-session.ts index 12453cc38..ca52a7c7b 100644 --- a/test/integration/compaction-rung-session.ts +++ b/test/integration/compaction-rung-session.ts @@ -11,7 +11,7 @@ import type { Api, AssistantMessage, Model } from "@earendil-works/pi-ai/compat" import { getModel } from "@earendil-works/pi-ai/compat"; import { AgentSession, type AgentSessionEvent } from "../../packages/coding-agent/src/core/agent-session.js"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.js"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.js"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.js"; import { SessionManager } from "../../packages/coding-agent/src/core/session-manager.js"; import { SettingsManager } from "../../packages/coding-agent/src/core/settings-manager.js"; import { createTestResourceLoader } from "../../packages/coding-agent/test/utilities.js"; @@ -122,7 +122,7 @@ function assistantTurn(text: string, timestamp: number, totalTokens: number): As } as AssistantMessage; } -export function createRungSession(options: RungSessionOptions): RungSession { +export async function createRungSession(options: RungSessionOptions): Promise { const model = options.contextWindow === undefined ? SESSION_MODEL : ({ ...SESSION_MODEL, contextWindow: options.contextWindow } as Model); @@ -146,9 +146,11 @@ export function createRungSession(options: RungSessionOptions): RungSession { return realContinue(...args); }) as typeof agent.continue; - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("anthropic", "anthropic-key"); - if (options.authenticateFallback !== false) authStorage.setRuntimeApiKey("openai", "openai-key"); + const authStorage = AuthStorage.inMemory({ + anthropic: { type: "api_key", key: "anthropic-key" }, + ...(options.authenticateFallback === false ? {} : { openai: { type: "api_key" as const, key: "openai-key" } }), + }); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const settingsManager = SettingsManager.inMemory(); settingsManager.applyOverrides({ retry: { enabled: false, maxRetries: 0, baseDelayMs: 0 }, @@ -164,7 +166,7 @@ export function createRungSession(options: RungSessionOptions): RungSession { sessionManager: manager, settingsManager, cwd: process.cwd(), - modelRegistry: ModelRegistry.create(authStorage), + modelRuntime, resourceLoader: createTestResourceLoader(), fallbackModels: options.fallbackModels ?? [], }); diff --git a/test/unit/auth-logout-invalidation.test.ts b/test/unit/auth-logout-invalidation.test.ts index 2ee562a07..d0ed55092 100644 --- a/test/unit/auth-logout-invalidation.test.ts +++ b/test/unit/auth-logout-invalidation.test.ts @@ -5,9 +5,13 @@ import { tmpdir } from "node:os"; import { join } from "node:path"; import lockfile from "proper-lockfile"; import { RpcClient } from "../../packages/coding-agent/src/modes/rpc/rpc-client.ts"; +import { loginIsolatedOAuthProvider } from "../../packages/coding-agent/src/modes/interactive-engine/isolated-auth.ts"; +import { RemoteModelCatalog } from "../../packages/coding-agent/src/modes/interactive-engine/remote-model-catalog.ts"; import { IsolatedInteractiveRuntime } from "../../packages/coding-agent/src/modes/interactive-engine/isolated-runtime.ts"; +import { InteractiveModeBase } from "../../packages/coding-agent/src/modes/interactive/interactive-mode-base.ts"; import { formatLogoutStatus } from "../../packages/coding-agent/src/modes/interactive/interactive-auth-routing.ts"; import { AuthStorage, type AuthStorageData } from "../../packages/coding-agent/src/core/auth-storage.ts"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.ts"; function readAuth(path: string): AuthStorageData { return JSON.parse(readFileSync(path, "utf8")) as AuthStorageData; @@ -30,9 +34,9 @@ describe("logout credential invalidation (#1919)", () => { try { const storage = AuthStorage.create([primary, legacy]); - assert.equal(storage.has("github-copilot"), true); - await storage.logoutAsync("github-copilot"); - assert.equal(storage.has("github-copilot"), false); + assert.notEqual(await storage.read("github-copilot"), undefined); + await storage.delete("github-copilot"); + assert.equal(await storage.read("github-copilot"), undefined); assert.deepEqual(readAuth(primary), { openai: { type: "api_key", key: "primary-openai" }, @@ -42,8 +46,8 @@ describe("logout credential invalidation (#1919)", () => { }); const restarted = AuthStorage.create([primary, legacy]); - assert.equal(restarted.has("github-copilot"), false); - assert.deepEqual(restarted.list(), ["anthropic", "openai"]); + assert.equal(await restarted.read("github-copilot"), undefined); + assert.deepEqual((await restarted.list()).map(({ providerId }) => providerId).sort(), ["anthropic", "openai"]); } finally { rmSync(directory, { recursive: true, force: true }); } @@ -57,7 +61,7 @@ describe("logout credential invalidation (#1919)", () => { await Bun.write(primary, JSON.stringify({ "github-copilot": { type: "api_key", key: "primary" } })); await Bun.write(legacy, JSON.stringify({ "github-copilot": { type: "api_key", key: "legacy" } })); const releaseLegacy = await lockfile.lock(legacy, { realpath: false }); - const logout = AuthStorage.create([primary, legacy]).logoutAsync("github-copilot"); + const logout = AuthStorage.create([primary, legacy]).delete("github-copilot"); try { for (let attempt = 0; attempt < 50 && !existsSync(`${primary}.lock`); attempt += 1) await Bun.sleep(10); @@ -198,24 +202,79 @@ describe("logout credential invalidation (#1919)", () => { } }); - test("the isolated host applies the child logout catalog and only reloads its local credential view", async () => { + test("isolated login reloads the frontend credential snapshot after the engine persists OAuth", async () => { + const directory = mkdtempSync(join(tmpdir(), "atomic-isolated-login-reload-")); + const authPath = join(directory, "auth.json"); + await Bun.write(authPath, "{}\n"); + const modelRuntime = await ModelRuntime.create({ authPath, modelsPath: null }); + const session = { modelRuntime, scopedModels: [] }; + const remoteCatalog = new RemoteModelCatalog({} as never); + remoteCatalog.patch(session as never); + const client = { + onExtensionUIRequest: () => () => {}, + respondExtensionUI: async () => {}, + cancelLoginProvider: async () => {}, + requestInternal: async () => { + await Bun.write(authPath, JSON.stringify({ + "corp-oauth": { + type: "oauth", + access: "engine-token", + refresh: "refresh-token", + expires: Date.now() + 60_000, + }, + })); + return { + provider: "corp-oauth", + cancelled: false, + models: [{ provider: "corp-oauth", id: "corp-model" }], + scopedModels: [], + customAuthProviders: [], + oauthProviders: [{ id: "corp-oauth", name: "Corp OAuth" }], + }; + }, + }; + + try { + assert.deepEqual(modelRuntime.getProviderAuthStatus("corp-oauth"), { configured: false }); + await loginIsolatedOAuthProvider( + session as never, + client as never, + remoteCatalog, + "corp-oauth", + {} as never, + ); + assert.deepEqual(modelRuntime.getProviderAuthStatus("corp-oauth"), { + configured: true, + source: "stored", + }); + assert.equal(modelRuntime.hasConfiguredAuth("corp-oauth"), true); + const logoutOptions = InteractiveModeBase.prototype.getLogoutProviderOptions.call({ session } as never); + assert.deepEqual(logoutOptions, [{ id: "corp-oauth", name: "Corp OAuth", authType: "oauth" }]); + } finally { + rmSync(directory, { recursive: true, force: true }); + } + }); + + test("the isolated host applies the child logout catalog without persisting frontend credentials", async () => { const model = { provider: "github-copilot", id: "claude-haiku-4.5" }; const authStorage = AuthStorage.inMemory({ "github-copilot": { type: "api_key", key: "controlled-fake-key" }, }); - let reloadCalls = 0; - authStorage.reload = () => { reloadCalls += 1; }; - authStorage.logoutAsync = async () => { - throw new Error("frontend must not persist isolated logout"); + let deleteCalls = 0; + const originalDelete = authStorage.delete.bind(authStorage); + authStorage.delete = async (...args) => { + deleteCalls += 1; + return originalDelete(...args); }; - const registry = { - authStorage, - getAvailable: () => [model], - find: () => model, - hasConfiguredAuth: () => true, + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); + let reloadCalls = 0; + const originalReloadCredentials = modelRuntime.reloadCredentials.bind(modelRuntime); + modelRuntime.reloadCredentials = async () => { + reloadCalls += 1; + await originalReloadCredentials(); }; const session = { - modelRegistry: registry, + modelRuntime, agent: { state: { model, thinkingLevel: "medium", messages: [] }, steeringMode: "all", @@ -223,6 +282,7 @@ describe("logout credential invalidation (#1919)", () => { }, scopedModels: [], sessionManager: {}, + refreshCurrentModelFromRegistry: () => {}, }; const client = { onEvent: () => () => {}, @@ -254,10 +314,12 @@ describe("logout credential invalidation (#1919)", () => { ); await runtime.initializeFromEngine(); - assert.deepEqual(await runtime.session.modelRegistry.getAvailable(), [model]); + assert.deepEqual(runtime.session.modelRuntime.getAvailableSnapshot(), [model]); await runtime.logoutProvider("github-copilot"); - assert.deepEqual(await runtime.session.modelRegistry.getAvailable(), []); + assert.deepEqual(runtime.session.modelRuntime.getAvailableSnapshot(), []); + assert.equal(deleteCalls, 0); assert.equal(reloadCalls, 1); + assert.notEqual(await authStorage.read("github-copilot"), undefined); }); test("logout status names remaining environment auth without changing ordinary success text", () => { diff --git a/test/unit/compaction-borrowing-purity-persisted.test.ts b/test/unit/compaction-borrowing-purity-persisted.test.ts index ce616c7bc..d620703fa 100644 --- a/test/unit/compaction-borrowing-purity-persisted.test.ts +++ b/test/unit/compaction-borrowing-purity-persisted.test.ts @@ -28,7 +28,7 @@ import type { Api, AssistantMessage, Model } from "@earendil-works/pi-ai/compat" import { getModel } from "@earendil-works/pi-ai/compat"; import { AgentSession, type AgentSessionEvent } from "../../packages/coding-agent/src/core/agent-session.js"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.js"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.js"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.js"; import { SessionManager } from "../../packages/coding-agent/src/core/session-manager.js"; import { SettingsManager } from "../../packages/coding-agent/src/core/settings-manager.js"; import { createTestResourceLoader } from "../../packages/coding-agent/test/utilities.js"; @@ -85,7 +85,7 @@ interface PersistedHarness { dispose: () => void; } -function createPersistedSession(directory: string): PersistedHarness { +async function createPersistedSession(directory: string): Promise { const { streamFn, models } = rescueStream(); const manager = SessionManager.create(directory, directory); const agent = new Agent({ @@ -100,9 +100,11 @@ function createPersistedSession(directory: string): PersistedHarness { return realContinue(...args); }) as typeof agent.continue; - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("anthropic", PRIMARY_KEY); - authStorage.setRuntimeApiKey("openai", FALLBACK_KEY); + const authStorage = AuthStorage.inMemory({ + anthropic: { type: "api_key", key: PRIMARY_KEY }, + openai: { type: "api_key", key: FALLBACK_KEY }, + }); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const settingsManager = SettingsManager.inMemory(); // One transport attempt per candidate, so the ladder alone explains the calls. settingsManager.applyOverrides({ retry: { enabled: false, maxRetries: 0, baseDelayMs: 0 } }); @@ -111,7 +113,7 @@ function createPersistedSession(directory: string): PersistedHarness { sessionManager: manager, settingsManager, cwd: directory, - modelRegistry: ModelRegistry.create(authStorage), + modelRuntime, resourceLoader: createTestResourceLoader(), fallbackModels: ["openai/gpt-5.1"], }); @@ -169,7 +171,7 @@ function recordsOf(raw: string): string[] { test("a rescued compaction appends exactly one durable boundary and nothing else", async () => { const directory = mkdtempSync(join(tmpdir(), "compaction-purity-persisted-")); - const harness = createPersistedSession(directory); + const harness = await createPersistedSession(directory); try { const before = { model: harness.session.model, @@ -241,7 +243,7 @@ test("a rescued compaction appends exactly one durable boundary and nothing else test("no borrowed or session credential reaches the durable session artifacts", async () => { const directory = mkdtempSync(join(tmpdir(), "compaction-purity-secrets-")); - const harness = createPersistedSession(directory); + const harness = await createPersistedSession(directory); try { const result = await harness.session.compact({ preserve_recent: 2 }); assert.equal(result.plannerModel?.id, "gpt-5.1"); diff --git a/test/unit/compaction-borrowing-purity.test.ts b/test/unit/compaction-borrowing-purity.test.ts index 7f262a446..0f23bdde1 100644 --- a/test/unit/compaction-borrowing-purity.test.ts +++ b/test/unit/compaction-borrowing-purity.test.ts @@ -15,7 +15,7 @@ import { getModel } from "@earendil-works/pi-ai/compat"; import { AgentSession } from "../../packages/coding-agent/src/core/agent-session.js"; import type { AgentSessionEvent } from "../../packages/coding-agent/src/core/agent-session.js"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.js"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.js"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.js"; import { SessionManager } from "../../packages/coding-agent/src/core/session-manager.js"; import { SettingsManager } from "../../packages/coding-agent/src/core/settings-manager.js"; import { runVerbatimCompaction } from "../../packages/coding-agent/src/core/compaction/compaction-runner.js"; @@ -123,9 +123,11 @@ test("a fallback-rescued compaction leaves session model, thinking level, histor return realContinue(...args); }) as typeof agent.continue; - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey("anthropic", "anthropic-key"); - authStorage.setRuntimeApiKey("openai", "openai-key"); + const authStorage = AuthStorage.inMemory({ + anthropic: { type: "api_key", key: "anthropic-key" }, + openai: { type: "api_key", key: "openai-key" }, + }); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const settingsManager = SettingsManager.inMemory(); // One transport attempt per candidate, so the ladder alone explains the calls. settingsManager.applyOverrides({ retry: { enabled: false, maxRetries: 0, baseDelayMs: 0 } }); @@ -134,7 +136,7 @@ test("a fallback-rescued compaction leaves session model, thinking level, histor sessionManager: manager, settingsManager, cwd: process.cwd(), - modelRegistry: ModelRegistry.create(authStorage), + modelRuntime, resourceLoader: createTestResourceLoader(), fallbackModels: ["openai/gpt-5.1"], }); diff --git a/test/unit/compaction-fallback-planner.test.ts b/test/unit/compaction-fallback-planner.test.ts index ef26753c1..13929ebd9 100644 --- a/test/unit/compaction-fallback-planner.test.ts +++ b/test/unit/compaction-fallback-planner.test.ts @@ -206,8 +206,8 @@ test("an entry that could not resolve to a model is never revisited", async () = // later configured candidate. const known = [primary, tertiary]; const registry = { - getAvailable: () => known.filter((model) => model.id !== "planner-b"), - find: (provider: string, id: string) => known.find((m) => m.provider === provider && m.id === id), + getAvailableSnapshot: () => known.filter((model) => model.id !== "planner-b"), + getModel: (provider: string, id: string) => known.find((m) => m.provider === provider && m.id === id), hasConfiguredAuth: () => true, }; const borrow = borrower({ fallbackModels: ["backup/planner-b", "spare/planner-c"], registry }); diff --git a/test/unit/compaction-rung-support.ts b/test/unit/compaction-rung-support.ts index e8a9edade..211900810 100644 --- a/test/unit/compaction-rung-support.ts +++ b/test/unit/compaction-rung-support.ts @@ -138,12 +138,11 @@ export function borrowed( export function registryOf(models: Model[], unauthenticated: string[] = []) { const authenticated = (model: Model) => !unauthenticated.includes(model.id); return { - getAvailable: () => models.filter(authenticated), - find: (provider: string, id: string) => models.find((model) => model.provider === provider && model.id === id), - hasConfiguredAuth: authenticated, + getAvailableSnapshot: () => models.filter(authenticated), + getModel: (provider: string, id: string) => models.find((model) => model.provider === provider && model.id === id), + hasConfiguredAuth: (provider: string) => models.some((model) => model.provider === provider && authenticated(model)), }; } - export function runRequest(overrides: Partial & { streamFn: StreamFn }): CompactionRunRequest { return { resolveAuth: async () => ({ apiKey: "primary-key" }), diff --git a/test/unit/executor-queued-message-helpers.ts b/test/unit/executor-queued-message-helpers.ts index 04897150a..e88af5032 100644 --- a/test/unit/executor-queued-message-helpers.ts +++ b/test/unit/executor-queued-message-helpers.ts @@ -71,7 +71,7 @@ export function streamingTurnSession(recorder: QueuedMessageRecorder): Streaming // Makes the fake visible through handle.agentSession, matching the SDK path. state: {} as never, sessionManager: {} as never, - modelRegistry: {} as never, + modelRuntime: {} as never, getContextUsage: (() => undefined) as never, get isStreaming() { return streaming; }, get pendingMessageCount() { return steering.length + followUp.length; }, diff --git a/test/unit/interactive-engine-model-refresh.test.ts b/test/unit/interactive-engine-model-refresh.test.ts index 67476a4e1..b5bee98af 100644 --- a/test/unit/interactive-engine-model-refresh.test.ts +++ b/test/unit/interactive-engine-model-refresh.test.ts @@ -12,21 +12,22 @@ import { type CreateAgentSessionRuntimeFactory, } from "../../packages/coding-agent/src/core/agent-session-runtime.ts"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.ts"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.ts"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.ts"; import { SessionManager } from "../../packages/coding-agent/src/core/session-manager.ts"; import { IsolatedInteractiveRuntime } from "../../packages/coding-agent/src/modes/interactive-engine/isolated-runtime.ts"; import type { RpcClient } from "../../packages/coding-agent/src/modes/rpc/rpc-client.ts"; import type { RpcModelRefreshResult } from "../../packages/coding-agent/src/modes/rpc/rpc-types.ts"; -function kimiModel(): Model { +async function kimiModel(): Promise> { const auth = AuthStorage.inMemory({ "kimi-coding": { type: "api_key", key: "fake-kimi-key" } }); - const model = ModelRegistry.inMemory(auth).getAvailable().find((candidate) => candidate.provider === "kimi-coding"); + const runtime = await ModelRuntime.create({ credentials: auth, modelsPath: null }); + const model = runtime.getAvailableSnapshot().find((candidate) => candidate.provider === "kimi-coding"); assert.ok(model); return model; } -test("isolated host refresh atomically applies the engine model catalog without restart", async () => { - const model = kimiModel(); +test("isolated host treats an engine-published extension model as configured after refresh", async () => { + const model = { ...(await kimiModel()), provider: "auth-free-extension", id: "published-model" }; const scopedModels = [{ model, thinkingLevel: "high" as const }]; let observedOptions: { timeoutMs?: number; force?: boolean; allowNetwork?: boolean } | undefined; const refreshResult: RpcModelRefreshResult = { @@ -66,9 +67,9 @@ test("isolated host refresh atomically applies the engine model catalog without cancelLoginProvider: async () => {}, getCommands: async () => [], } as unknown as RpcClient; - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const modelRuntime = await ModelRuntime.create({ modelsPath: null }); const session = { - modelRegistry: registry, + modelRuntime, scopedModels: [], sessionFile: undefined, agent: { @@ -83,21 +84,16 @@ test("isolated host refresh atomically applies the engine model catalog without const runtime = new IsolatedInteractiveRuntime(localRuntime, createRuntime, client); await runtime.initializeFromEngine(); - assert.deepEqual(registry.getAvailable(), []); - assert.deepEqual(registry.getCustomApiKeyAuthProviders(), [{ id: "extension-provider", name: "Extension Provider" }]); - assert.equal(registry.getProviderDisplayName("extension-provider"), "Extension Provider"); - const remoteAuth = registry.getCustomApiKeyAuth("extension-provider"); - assert.ok(remoteAuth); - assert.equal(remoteAuth.name, "Extension Provider"); - assert.deepEqual(await remoteAuth.login({ signal: new AbortController().signal, prompt: async () => "unused" }), { - type: "api_key", key: "remote-key", - }); - const result = await registry.refresh({ allowNetwork: false, force: true, timeoutMs: 321 }); + assert.deepEqual(modelRuntime.getAvailableSnapshot(), []); + assert.equal(modelRuntime.hasConfiguredAuth(model.provider), false); + const result = await (modelRuntime as unknown as { + refresh(options: { allowNetwork?: boolean; force?: boolean; timeoutMs?: number }): ReturnType; + }).refresh({ allowNetwork: false, force: true, timeoutMs: 321 }); assert.deepEqual(observedOptions, { allowNetwork: false, force: true, timeoutMs: 321 }); - assert.deepEqual(registry.getAvailable(), [model]); - assert.equal(registry.find(model.provider, model.id), model); - assert.equal(registry.hasConfiguredAuth(model), true); + assert.deepEqual(modelRuntime.getAvailableSnapshot(), [model]); + assert.equal(modelRuntime.getModel(model.provider, model.id), model); + assert.equal(modelRuntime.hasConfiguredAuth(model.provider), true); assert.deepEqual(session.scopedModels, scopedModels); assert.equal(result.aborted, false); assert.ok(result.errors instanceof Map); @@ -105,7 +101,7 @@ test("isolated host refresh atomically applies the engine model catalog without }); test("an aborted isolated refresh does not replace the current model catalog", async () => { - const model = kimiModel(); + const model = await kimiModel(); let resolveRefresh!: (result: RpcModelRefreshResult) => void; const pending = new Promise((resolve) => { resolveRefresh = resolve; }); const client = { @@ -119,9 +115,9 @@ test("an aborted isolated refresh does not replace the current model catalog", a refreshModels: async () => pending, getCommands: async () => [], } as unknown as RpcClient; - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const modelRuntime = await ModelRuntime.create({ modelsPath: null }); const session = { - modelRegistry: registry, scopedModels: [], sessionFile: undefined, + modelRuntime, scopedModels: [], sessionFile: undefined, agent: { state: { model: undefined, thinkingLevel: "off", messages: [] }, steeringMode: "all", followUpMode: "all" }, } as unknown as AgentSession; const createRuntime = (async () => { throw new Error("not used"); }) as CreateAgentSessionRuntimeFactory; @@ -129,18 +125,18 @@ test("an aborted isolated refresh does not replace the current model catalog", a const runtime = new IsolatedInteractiveRuntime(localRuntime, createRuntime, client); await runtime.initializeFromEngine(); const controller = new AbortController(); - const refresh = registry.refresh({ signal: controller.signal }); + const refresh = modelRuntime.refresh({ signal: controller.signal }); controller.abort(); assert.deepEqual(await refresh, { aborted: true, errors: new Map() }); resolveRefresh({ aborted: false, errors: [], models: [model], scopedModels: [{ model }], customAuthProviders: [] }); await Bun.sleep(0); - assert.deepEqual(registry.getAvailable(), []); + assert.deepEqual(modelRuntime.getAvailableSnapshot(), []); assert.deepEqual(session.scopedModels, []); }); test("isolated host synchronizes authoritative engine fallback state and clears it after remote model selection", async () => { - const model = kimiModel(); + const model = await kimiModel(); type EngineState = Awaited>; let state: EngineState = { model, @@ -162,9 +158,9 @@ test("isolated host synchronizes authoritative engine fallback state and clears setModel: async () => model, getCommands: async () => [], } as unknown as RpcClient; - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const modelRuntime = await ModelRuntime.create({ modelsPath: null }); const session = { - modelRegistry: registry, + modelRuntime, sessionManager: SessionManager.inMemory(process.cwd()), scopedModels: [], sessionFile: undefined, @@ -216,7 +212,7 @@ test("isolated host synchronizes authoritative engine fallback state and clears }); test("isolated explicit cycle clears fallback only for a changed model despite an intervening model event", async () => { - const previous = kimiModel(); + const previous = await kimiModel(); const selected = { ...previous, id: `${previous.id}-next`, name: `${previous.name} Next` }; type CycleBehavior = "changed" | "same" | "null" | "throw"; let behavior: CycleBehavior = "changed"; @@ -232,9 +228,9 @@ test("isolated explicit cycle clears fallback only for a changed model despite a }, getCommands: async () => [], } as unknown as RpcClient; - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const modelRuntime = await ModelRuntime.create({ modelsPath: null }); const sessionFixture = { - modelRegistry: registry, + modelRuntime, sessionManager: SessionManager.inMemory(process.cwd()), scopedModels: [], sessionFile: undefined, @@ -296,7 +292,7 @@ test("isolated session synchronization replaces each engine-selected session exa const services = await createAgentSessionServices({ cwd: options.cwd, agentDir: options.agentDir, - authStorage, + modelRuntime: await ModelRuntime.create({ credentials: authStorage, modelsPath: null }), resourceLoaderOptions: { noExtensions: true, noSkills: true, noPromptTemplates: true, noThemes: true }, }); return { @@ -376,9 +372,9 @@ test("isolated queue pause reaches the engine before abort and resumes remotely" abort: async () => { calls.push("abort"); }, getCommands: async () => [], } as unknown as RpcClient; - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + const modelRuntime = await ModelRuntime.create({ modelsPath: null }); const session = { - modelRegistry: registry, scopedModels: [], sessionFile: undefined, sessionManager: {}, + modelRuntime, scopedModels: [], sessionFile: undefined, sessionManager: {}, agent: { state: { model: undefined, thinkingLevel: "off", messages: [] }, steeringMode: "all", followUpMode: "all" }, } as unknown as AgentSession; const createRuntime = (async () => { throw new Error("not used"); }) as CreateAgentSessionRuntimeFactory; diff --git a/test/unit/interactive-engine-oauth.test.ts b/test/unit/interactive-engine-oauth.test.ts index 15fa00cf6..bca613f11 100644 --- a/test/unit/interactive-engine-oauth.test.ts +++ b/test/unit/interactive-engine-oauth.test.ts @@ -1,9 +1,9 @@ import { test } from "bun:test"; import assert from "node:assert/strict"; -import { mkdtempSync, readFileSync, rmSync } from "node:fs"; +import { existsSync, mkdtempSync, readFileSync, rmSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; -import type { AtomicOAuthLoginCallbacks } from "../../packages/coding-agent/src/core/oauth-provider-bridge.ts"; +import type { AtomicOAuthLoginCallbacks } from "../../packages/coding-agent/src/core/oauth-login.ts"; import { RpcClient } from "../../packages/coding-agent/src/modes/rpc/rpc-client.ts"; import { loginRpcOAuthProvider } from "../../packages/coding-agent/src/modes/rpc/rpc-oauth-client.ts"; import type { RpcModelCatalog } from "../../packages/coding-agent/src/modes/rpc/rpc-types.ts"; @@ -38,7 +38,7 @@ test.serial("real isolated child discovers and acquires engine-only custom OAuth assert.equal(before.oauthProviders?.some(({ id }) => id === "openrouter"), true); assert.equal(before.oauthProviders?.some(({ id }) => id === "kimi-coding"), true); assert.equal(JSON.stringify(before).includes("engine-token"), false); - const refreshesBeforeLogin = readFileSync(logFile, "utf8").match(/^refresh:/gm)?.length ?? 0; + const refreshesBeforeLogin = existsSync(logFile) ? (readFileSync(logFile, "utf8").match(/^refresh:/gm)?.length ?? 0) : 0; const callbacks: AtomicOAuthLoginCallbacks = { onAuth: ({ url }) => callbackLog.push(`auth:${url}`), onDeviceCode: ({ userCode }) => callbackLog.push(`device:${userCode}`), diff --git a/test/unit/interactive-mode-keybinding-lifecycle.test.ts b/test/unit/interactive-mode-keybinding-lifecycle.test.ts index f95253c8c..f75841f96 100644 --- a/test/unit/interactive-mode-keybinding-lifecycle.test.ts +++ b/test/unit/interactive-mode-keybinding-lifecycle.test.ts @@ -13,6 +13,7 @@ import { } from "../../packages/coding-agent/src/core/agent-session-runtime.ts"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.ts"; import type { AgentSessionReloadOptions } from "../../packages/coding-agent/src/core/agent-session-types.ts"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.ts"; import { SessionManager } from "../../packages/coding-agent/src/core/session-manager.ts"; import type { ExtensionFactory } from "../../packages/coding-agent/src/core/extensions/types.ts"; import { keyText } from "../../packages/coding-agent/src/modes/interactive/components/keybinding-hints.ts"; @@ -55,13 +56,13 @@ function writeExpandBinding(agentDir: string, binding: string): void { async function createMode(agentDir: string, extensionFactory?: ExtensionFactory): Promise { const cwd = mkdtempSync(join(tmpdir(), "atomic-local-mode-cwd-")); const faux = registerFauxProvider(); - const authStorage = AuthStorage.inMemory(); - authStorage.setRuntimeApiKey(faux.getModel().provider, "faux-key"); + const authStorage = AuthStorage.inMemory({ [faux.getModel().provider]: { type: "api_key", key: "faux-key" } }); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const createRuntime: CreateAgentSessionRuntimeFactory = async ({ cwd: runtimeCwd, agentDir: runtimeAgentDir, sessionManager, sessionStartEvent }) => { const services = await createAgentSessionServices({ cwd: runtimeCwd, agentDir: runtimeAgentDir, - authStorage, + modelRuntime, resourceLoaderOptions: { extensionFactories: extensionFactory ? [extensionFactory] : [], noSkills: true, diff --git a/test/unit/llama-extension-parity.test.ts b/test/unit/llama-extension-parity.test.ts index 8762b3fab..dcb63da95 100644 --- a/test/unit/llama-extension-parity.test.ts +++ b/test/unit/llama-extension-parity.test.ts @@ -6,12 +6,12 @@ import { join } from "node:path"; import { InMemoryModelsStore } from "@earendil-works/pi-ai"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.js"; import { DefaultResourceLoader } from "../../packages/coding-agent/src/core/resource-loader.js"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.js"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.js"; import { FileModelsStore } from "../../packages/coding-agent/src/core/models-store.js"; import { SettingsManager } from "../../packages/coding-agent/src/core/settings-manager.js"; import { builtInExtensions } from "../../packages/coding-agent/src/extensions/index.js"; import { LlamaClient, llamaInferenceUrl, normalizeLlamaServerUrl } from "../../packages/coding-agent/src/extensions/llama/client.js"; -import { createLlamaProvider, toLlamaModel } from "../../packages/coding-agent/src/extensions/llama/provider.js"; +import { createLlamaProvider } from "../../packages/coding-agent/src/extensions/llama/provider.js"; const originalFetch = globalThis.fetch; afterEach(() => { globalThis.fetch = originalFetch; }); @@ -39,59 +39,62 @@ describe("llama.cpp router client", () => { describe("llama.cpp provider", () => { test("maps loaded model metadata to context-sized, zero-cost OpenAI models", () => { - const mapped = toLlamaModel({ - id: "vision.gguf", - status: { value: "loaded" }, - architecture: { input_modalities: ["text", "image"] }, - meta: { n_ctx: 8192 }, - }, "http://localhost:8080"); + const controller = createLlamaProvider(); + controller.setCatalog( + [ + { + id: "vision.gguf", + status: { value: "loaded" }, + architecture: { input_modalities: ["text", "image"] }, + meta: { n_ctx: 8192 }, + }, + { id: "large", status: { value: "loaded" }, meta: { n_ctx: 32768 } }, + { id: "fallback", status: { value: "loaded" } }, + { id: "idle", status: { value: "unloaded" } }, + ], + "http://localhost:8080", + ); + const models = controller.provider.getModels(); + assert.deepEqual(models.map((model) => model.id), ["vision.gguf", "large", "fallback"]); + const mapped = models[0]!; assert.equal(mapped.baseUrl, "http://localhost:8080/v1"); assert.deepEqual(mapped.input, ["text", "image"]); assert.equal(mapped.contextWindow, 8192); assert.equal(mapped.maxTokens, 8192); assert.deepEqual(mapped.cost, { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }); - assert.equal(toLlamaModel({ id: "large", status: { value: "loaded" }, meta: { n_ctx: 32768 } }, "http://localhost").maxTokens, 32768); - assert.equal(toLlamaModel({ id: "fallback", status: { value: "loaded" } }, "http://localhost").contextWindow, 128000); + assert.equal(models[1]?.maxTokens, 32768); + assert.equal(models[2]?.contextWindow, 128000); }); - test("custom login validates the router and stores normalized URL plus optional key", async () => { - let requestedUrl = ""; - globalThis.fetch = mockFetch(async (input) => { - requestedUrl = String(input); - return jsonResponse({ data: [{ id: "local", status: { value: "loaded" } }] }); - }); - const auth = createLlamaProvider().config.auth?.apiKey; - assert.ok(auth); - const answers = ["http://localhost:9000/v1/", "secret"]; - const credential = await auth.login({ - signal: new AbortController().signal, - prompt: async () => answers.shift() ?? "", - }); - assert.equal(requestedUrl, "http://localhost:9000/models"); - assert.deepEqual(credential, { - type: "api_key", - key: "secret", - env: { LLAMA_BASE_URL: "http://localhost:9000" }, - }); + test("stays dormant in refresh without a configured server URL", async () => { + let fetched = false; + globalThis.fetch = mockFetch(async () => { fetched = true; return jsonResponse({ data: [] }); }); + const runtime = await ModelRuntime.create({ credentials: AuthStorage.inMemory({}), modelsPath: null }); + runtime.registerNativeProvider(createLlamaProvider().provider); + const result = await runtime.refresh({ allowNetwork: true }); + assert.equal(result.errors.has("llama.cpp"), false); + assert.equal(fetched, false); + assert.equal(runtime.getModels("llama.cpp").length, 0); }); - test("dynamic refresh exposes only loaded models and resolves stored URL metadata", async () => { + test("dynamic refresh exposes only loaded models through the configured provider", async () => { globalThis.fetch = mockFetch(async () => jsonResponse({ data: [ { id: "loaded", status: { value: "loaded" }, meta: { n_ctx: 4096 } }, { id: "idle", status: { value: "unloaded" } }, ] })); const storage = AuthStorage.inMemory({ - "llama.cpp": { type: "api_key", env: { LLAMA_BASE_URL: "http://localhost:8080" } }, + "llama.cpp": { type: "api_key", key: "local", env: { LLAMA_BASE_URL: "http://localhost:8080" } }, }); - const registry = ModelRegistry.inMemory(storage); - registry.registerProvider("llama.cpp", createLlamaProvider().config); - assert.deepEqual(registry.getCustomApiKeyAuthProviders(), [{ id: "llama.cpp", name: "llama.cpp server" }]); - const result = await registry.refresh(); + const runtime = await ModelRuntime.create({ credentials: storage, modelsPath: null }); + runtime.registerNativeProvider(createLlamaProvider().provider); + await runtime.refresh({ allowNetwork: false }); + const result = await runtime.refresh({ allowNetwork: true }); assert.equal(result.errors.size, 0); - const models = registry.getAll().filter((model) => model.provider === "llama.cpp"); + const models = runtime.getModels("llama.cpp"); assert.deepEqual(models.map((model) => model.id), ["loaded"]); - const auth = await registry.getApiKeyAndHeaders(models[0]!); - assert.deepEqual(auth, { ok: true, apiKey: "local", headers: undefined, baseUrl: "http://localhost:8080/v1" }); + const auth = await runtime.getAuth(models[0]!); + assert.equal(auth?.auth.apiKey, "local"); + assert.equal(models[0]?.baseUrl, "http://localhost:8080/v1"); }); test("persists the loaded catalog and restores it before a later network refresh", async () => { @@ -105,16 +108,16 @@ describe("llama.cpp provider", () => { delete: () => backing.delete("llama.cpp"), }; const credential = { type: "api_key" as const, key: "local", env: { LLAMA_BASE_URL: "http://localhost:8080" } }; - const first = createLlamaProvider().config; - const refreshed = await first.refreshModels?.({ credential, store, allowNetwork: true }); - assert.equal(refreshed?.[0]?.id, "cached"); + const first = createLlamaProvider(); + await first.provider.refreshModels?.({ credential, store, allowNetwork: true }); + assert.equal(first.provider.getModels()[0]?.id, "cached"); assert.equal((await backing.read("llama.cpp"))?.models[0]?.maxTokens, 65536); globalThis.fetch = mockFetch(async () => { throw new Error("offline"); }); - const restarted = createLlamaProvider().config; - const restored = await restarted.refreshModels?.({ credential, store, allowNetwork: false }); - assert.equal(restored?.[0]?.id, "cached"); - assert.equal(restored?.[0]?.maxTokens, 65536); + const restarted = createLlamaProvider(); + await restarted.provider.refreshModels?.({ credential, store, allowNetwork: false }); + assert.equal(restarted.provider.getModels()[0]?.id, "cached"); + assert.equal(restarted.provider.getModels()[0]?.maxTokens, 65536); }); test("keeps a validated cached catalog visible when the first online refresh after restart fails", async () => { @@ -126,31 +129,33 @@ describe("llama.cpp provider", () => { globalThis.fetch = mockFetch(async () => jsonResponse({ data: [ { id: "cached-on-restart", status: { value: "loaded" }, meta: { n_ctx: 65536 } }, ] })); - const first = ModelRegistry.create(auth, [], root); - first.registerProvider("llama.cpp", createLlamaProvider().config); + const first = await ModelRuntime.create({ credentials: auth, modelsPath: join(root, "models.json"), modelsStorePath: join(root, "models-store.json") }); + first.registerNativeProvider(createLlamaProvider().provider); + await first.refresh({ allowNetwork: false }); assert.equal((await first.refresh({ allowNetwork: true })).errors.size, 0); globalThis.fetch = mockFetch(async () => { throw new Error("router offline"); }); - const restarted = ModelRegistry.create(auth, [], root); - restarted.registerProvider("llama.cpp", createLlamaProvider().config); + const restarted = await ModelRuntime.create({ credentials: auth, modelsPath: join(root, "models.json"), modelsStorePath: join(root, "models-store.json") }); + restarted.registerNativeProvider(createLlamaProvider().provider); + await restarted.refresh({ allowNetwork: false }); const result = await restarted.refresh({ allowNetwork: true }); assert.equal(result.errors.get("llama.cpp")?.message, "router offline"); assert.deepEqual( - restarted.getAll().filter((model) => model.provider === "llama.cpp").map((model) => model.id), + restarted.getModels().filter((model) => model.provider === "llama.cpp").map((model) => model.id), ["cached-on-restart"], ); - assert.equal(restarted.find("llama.cpp", "cached-on-restart")?.maxTokens, 65536); + assert.equal(restarted.getModel("llama.cpp", "cached-on-restart")?.maxTokens, 65536); globalThis.fetch = mockFetch(async () => jsonResponse({ data: [ { id: "fresh-after-restart", status: { value: "loaded" }, meta: { n_ctx: 32768 } }, ] })); assert.equal((await restarted.refresh({ allowNetwork: true })).errors.size, 0); assert.deepEqual( - restarted.getAll().filter((model) => model.provider === "llama.cpp").map((model) => model.id), + restarted.getModels().filter((model) => model.provider === "llama.cpp").map((model) => model.id), ["fresh-after-restart"], ); - assert.equal(restarted.find("llama.cpp", "cached-on-restart"), undefined); + assert.equal(restarted.getModel("llama.cpp", "cached-on-restart"), undefined); } finally { rmSync(root, { recursive: true, force: true }); } @@ -163,13 +168,14 @@ describe("llama.cpp provider", () => { "llama.cpp": { type: "api_key", key: "local", env: { LLAMA_BASE_URL: "http://localhost:8080" } }, }); globalThis.fetch = mockFetch(async () => { throw new Error("router offline"); }); - const registry = ModelRegistry.create(auth, [], root); - registry.registerProvider("llama.cpp", createLlamaProvider().config); + const registry = await ModelRuntime.create({ credentials: auth, modelsPath: join(root, "models.json"), modelsStorePath: join(root, "models-store.json") }); + registry.registerNativeProvider(createLlamaProvider().provider); + await registry.refresh({ allowNetwork: false }); const result = await registry.refresh({ allowNetwork: true }); assert.equal(result.errors.get("llama.cpp")?.message, "router offline"); - assert.equal(registry.getAll().some((model) => model.provider === "llama.cpp"), false); + assert.equal(registry.getModels().some((model) => model.provider === "llama.cpp"), false); } finally { rmSync(root, { recursive: true, force: true }); } @@ -178,7 +184,9 @@ describe("llama.cpp provider", () => { test("rejects cached entries for the wrong provider or API during failed-refresh recovery", async () => { const root = mkdtempSync(join(tmpdir(), "atomic-llama-invalid-")); try { - const valid = toLlamaModel({ id: "invalid", status: { value: "loaded" } }, "http://localhost:8080"); + const mapper = createLlamaProvider(); + mapper.setCatalog([{ id: "invalid", status: { value: "loaded" } }], "http://localhost:8080"); + const valid = mapper.provider.getModels()[0]!; await new FileModelsStore(join(root, "models-store.json")).write("llama.cpp", { models: [{ ...valid, provider: "other" }, { ...valid, api: "anthropic-messages" }], checkedAt: 0, @@ -187,13 +195,14 @@ describe("llama.cpp provider", () => { "llama.cpp": { type: "api_key", key: "local", env: { LLAMA_BASE_URL: "http://localhost:8080" } }, }); globalThis.fetch = mockFetch(async () => { throw new Error("router offline"); }); - const registry = ModelRegistry.create(auth, [], root); - registry.registerProvider("llama.cpp", createLlamaProvider().config); + const registry = await ModelRuntime.create({ credentials: auth, modelsPath: join(root, "models.json"), modelsStorePath: join(root, "models-store.json") }); + registry.registerNativeProvider(createLlamaProvider().provider); + await registry.refresh({ allowNetwork: false }); const result = await registry.refresh({ allowNetwork: true }); assert.equal(result.errors.get("llama.cpp")?.message, "router offline"); - assert.equal(registry.getAll().some((model) => model.provider === "llama.cpp"), false); + assert.equal(registry.getModels().some((model) => model.provider === "llama.cpp"), false); } finally { rmSync(root, { recursive: true, force: true }); } @@ -236,8 +245,8 @@ describe("built-in inline extension", () => { }); // Keep the auth-storage import exercised against the newly optional key shape. -test("provider metadata credentials remain backward compatible", () => { +test("provider metadata credentials remain backward compatible", async () => { const storage = AuthStorage.inMemory({ old: { type: "api_key", key: "legacy" }, local: { type: "api_key", env: { LLAMA_BASE_URL: "http://localhost" } } }); - assert.equal(storage.get("old")?.type, "api_key"); - assert.deepEqual(storage.get("local"), { type: "api_key", env: { LLAMA_BASE_URL: "http://localhost" } }); + assert.equal((await storage.read("old"))?.type, "api_key"); + assert.deepEqual(await storage.read("local"), { type: "api_key", env: { LLAMA_BASE_URL: "http://localhost" } }); }); diff --git a/test/unit/login-provider-parity.test.ts b/test/unit/login-provider-parity.test.ts index 0e81bd267..86d33e1c2 100644 --- a/test/unit/login-provider-parity.test.ts +++ b/test/unit/login-provider-parity.test.ts @@ -4,9 +4,9 @@ import type { TUI } from "@earendil-works/pi-tui"; import { builtinProviders } from "@earendil-works/pi-ai/providers/all"; import { defaultModelPerProvider } from "../../packages/coding-agent/src/core/model-resolver-defaults.ts"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.ts"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.ts"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.ts"; import { BUILT_IN_PROVIDER_DISPLAY_NAMES } from "../../packages/coding-agent/src/core/provider-display-names.ts"; -import { createAuthInteraction } from "../../packages/coding-agent/src/core/oauth-provider-bridge.ts"; +import { createAuthInteraction } from "../../packages/coding-agent/src/core/oauth-login.ts"; import { BUILTIN_SLASH_COMMANDS } from "../../packages/coding-agent/src/core/slash-commands.ts"; import { LoginDialogComponent } from "../../packages/coding-agent/src/modes/interactive/components/login-dialog.ts"; import { InteractiveModeBase } from "../../packages/coding-agent/src/modes/interactive/interactive-mode-base.ts"; @@ -89,24 +89,25 @@ test("every installed builtin provider has a preferred default", () => { assert.equal(Object.hasOwn(defaultModelPerProvider, "cursor"), false); }); -test("stale Cursor authentication cannot restore the removed provider", () => { +test("stale Cursor authentication cannot restore the removed provider", async () => { const authStorage = AuthStorage.inMemory({ cursor: { type: "api_key", key: "stale-token" }, }); - const registry = ModelRegistry.inMemory(authStorage); + const runtime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); - assert.equal(authStorage.has("cursor"), true); - assert.equal(registry.getAll().some((model) => model.provider === "cursor"), false); + assert.notEqual(await authStorage.read("cursor"), undefined); + assert.equal(runtime.getModels().some((model) => model.provider === "cursor"), false); assert.equal(BUILT_IN_PROVIDER_DISPLAY_NAMES.cursor, undefined); }); -test("logout options ignore credentials for removed providers", () => { +test("logout options ignore credentials for removed providers", async () => { const authStorage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "active-token" }, cursor: { type: "api_key", key: "stale-token" }, }); + const modelRuntime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); const context = { - session: { modelRegistry: { authStorage, getProviderDisplayName: (id: string) => id } }, + session: { modelRuntime }, getLoginProviderOptions: () => [ { id: "anthropic", name: "Anthropic", authType: "api_key" as const }, ], diff --git a/test/unit/main-chat-model-fallback.test.ts b/test/unit/main-chat-model-fallback.test.ts index e576fdf12..1fbba6d94 100644 --- a/test/unit/main-chat-model-fallback.test.ts +++ b/test/unit/main-chat-model-fallback.test.ts @@ -69,10 +69,10 @@ test("main-chat fallback switches models after same-model retry exhaustion", asy getDefaultThinkingLevel: () => "high" as ThinkingLevel, getDefaultProvider: () => "openai-codex", }, - _modelRegistry: { - getAvailable: () => [primary, fallback], - find: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, - hasConfiguredAuth: (candidate: Model) => candidate.provider === fallback.provider || candidate.provider === primary.provider, + _modelRuntime: { + getAvailableSnapshot: () => [primary, fallback], + getModel: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, + hasConfiguredAuth: (provider: string) => provider === fallback.provider || provider === primary.provider, }, agent: { state: { @@ -126,9 +126,9 @@ test("main-chat fallback can change reasoning on the same provider/model", async getDefaultThinkingLevel: () => "high" as ThinkingLevel, getDefaultProvider: () => "openai", }, - _modelRegistry: { - getAvailable: () => [primary], - find: (provider: string, id: string) => provider === primary.provider && id === primary.id ? primary : undefined, + _modelRuntime: { + getAvailableSnapshot: () => [primary], + getModel: (provider: string, id: string) => provider === primary.provider && id === primary.id ? primary : undefined, hasConfiguredAuth: () => true, }, agent: { @@ -178,10 +178,10 @@ test("main-chat retry-disabled fallback keeps prompt waiting for fallback comple getDefaultThinkingLevel: () => "high" as ThinkingLevel, getDefaultProvider: () => "openai-codex", }, - _modelRegistry: { - getAvailable: () => [primary, fallback], - find: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, - hasConfiguredAuth: (candidate: Model) => candidate.provider === fallback.provider || candidate.provider === primary.provider, + _modelRuntime: { + getAvailableSnapshot: () => [primary, fallback], + getModel: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, + hasConfiguredAuth: (provider: string) => provider === fallback.provider || provider === primary.provider, }, agent: { state: { model: primary, thinkingLevel: "high" as ThinkingLevel, messages: [message] }, @@ -246,10 +246,10 @@ test("main-chat fallback rejection settles the retry wait", async () => { getDefaultThinkingLevel: () => "high" as ThinkingLevel, getDefaultProvider: () => "openai-codex", }, - _modelRegistry: { - getAvailable: () => [primary, fallback], - find: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, - hasConfiguredAuth: (candidate: Model) => candidate.provider === fallback.provider || candidate.provider === primary.provider, + _modelRuntime: { + getAvailableSnapshot: () => [primary, fallback], + getModel: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, + hasConfiguredAuth: (provider: string) => provider === fallback.provider || provider === primary.provider, }, agent: { state: { model: primary, thinkingLevel: "high" as ThinkingLevel, messages: [retryableMessage()] }, @@ -299,10 +299,10 @@ test("main-chat fallback continuation resolution does not mark assistant errors getDefaultThinkingLevel: () => "high" as ThinkingLevel, getDefaultProvider: () => "openai-codex", }, - _modelRegistry: { - getAvailable: () => [primary, fallback], - find: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, - hasConfiguredAuth: (candidate: Model) => candidate.provider === fallback.provider || candidate.provider === primary.provider, + _modelRuntime: { + getAvailableSnapshot: () => [primary, fallback], + getModel: (provider: string, id: string) => provider === fallback.provider && id === fallback.id ? fallback : undefined, + hasConfiguredAuth: (provider: string) => provider === fallback.provider || provider === primary.provider, }, agent: { state: { model: primary, thinkingLevel: "high" as ThinkingLevel, messages: [retryableMessage()] }, diff --git a/test/unit/pi-0.82.1-model-catalog.test.ts b/test/unit/pi-0.82.1-model-catalog.test.ts index 37c561184..7c4a3cee4 100644 --- a/test/unit/pi-0.82.1-model-catalog.test.ts +++ b/test/unit/pi-0.82.1-model-catalog.test.ts @@ -1,19 +1,18 @@ import { describe, test } from "bun:test"; import assert from "node:assert/strict"; import { builtinImagesModels } from "@earendil-works/pi-ai/providers/all"; -import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.js"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.js"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.js"; -function requiredModel(registry: ModelRegistry, provider: string, id: string) { - const model = registry.find(provider, id); +function requiredModel(runtime: ModelRuntime, provider: string, id: string) { + const model = runtime.getModel(provider, id); assert.ok(model, `missing ${provider}/${id}`); return model; } describe("Pi 0.82.1 generated catalogs through Atomic", () => { - test("exposes Claude Opus 5 adaptive xhigh metadata", () => { - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); - const anthropic = requiredModel(registry, "anthropic", "claude-opus-5"); + test("exposes Claude Opus 5 adaptive xhigh metadata", async () => { + const runtime = await ModelRuntime.create({ modelsPath: null }); + const anthropic = requiredModel(runtime, "anthropic", "claude-opus-5"); assert.equal(anthropic.reasoning, true); assert.equal(anthropic.thinkingLevelMap?.xhigh, "xhigh"); assert.equal(anthropic.compat && "forceAdaptiveThinking" in anthropic.compat ? anthropic.compat.forceAdaptiveThinking : undefined, true); @@ -21,7 +20,7 @@ describe("Pi 0.82.1 generated catalogs through Atomic", () => { assert.equal(anthropic.compat && "supportsStrictTools" in anthropic.compat ? anthropic.compat.supportsStrictTools : undefined, true); for (const region of ["au", "eu", "global", "jp", "us"]) { - const bedrock = requiredModel(registry, "amazon-bedrock", `${region}.anthropic.claude-opus-5`); + const bedrock = requiredModel(runtime, "amazon-bedrock", `${region}.anthropic.claude-opus-5`); assert.equal(bedrock.api, "bedrock-converse-stream"); assert.equal(bedrock.reasoning, true); assert.equal(bedrock.thinkingLevelMap?.xhigh, "xhigh"); diff --git a/test/unit/pi-parity-group-3.test.ts b/test/unit/pi-parity-group-3.test.ts index 07e72c046..79adaa50b 100644 --- a/test/unit/pi-parity-group-3.test.ts +++ b/test/unit/pi-parity-group-3.test.ts @@ -6,7 +6,7 @@ import { join } from "node:path"; import type { OpenAICompletionsCompat } from "@earendil-works/pi-ai/compat"; import { parseConfigCommand } from "../../packages/coding-agent/src/config-command-parser.ts"; import { applyHttpProxySettings, parseHttpIdleTimeoutMs } from "../../packages/coding-agent/src/core/http-dispatcher.ts"; -import { loadCustomModelsFromPaths } from "../../packages/coding-agent/src/core/model-registry-custom-loader.ts"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.ts"; import { DefaultPackageManager } from "../../packages/coding-agent/src/core/package-manager.ts"; import { SettingsManager } from "../../packages/coding-agent/src/core/settings-manager.ts"; import { buildSelfUpdatePlan } from "../../packages/coding-agent/src/self-update-plan.ts"; @@ -60,14 +60,13 @@ test("models.json preserves deferredToolsMode and constructs configured Radius p }, corporate: { name: "Corporate Radius", baseUrl: "https://radius.example/v1", oauth: "radius" }, } })); - const loaded = loadCustomModelsFromPaths([path]); - assert.equal(loaded.error, undefined); - const compat = loaded.models.find((model) => model.provider === "kimi")?.compat as OpenAICompletionsCompat | undefined; + const runtime = await ModelRuntime.create({ modelsPath: path }); + const compat = runtime.getModels("kimi")[0]?.compat as OpenAICompletionsCompat | undefined; assert.equal(compat?.deferredToolsMode, "kimi"); - const radius = loaded.configuredProviders.get("corporate"); + const radius = runtime.getProvider("corporate"); assert.equal(radius?.id, "corporate"); assert.equal(radius?.name, "Corporate Radius"); - assert.equal(loaded.overrides.get("corporate")?.baseUrl, undefined); + assert.equal(radius?.baseUrl, "https://radius.example/v1"); }); test("project autoload:false is a delta over a global package including workflows", async () => { diff --git a/test/unit/pi-parity-group-4.test.ts b/test/unit/pi-parity-group-4.test.ts index 8216e44f2..c5d848f28 100644 --- a/test/unit/pi-parity-group-4.test.ts +++ b/test/unit/pi-parity-group-4.test.ts @@ -1,6 +1,6 @@ import { test } from "bun:test"; import assert from "node:assert/strict"; -import { mkdtempSync, rmSync } from "node:fs"; +import { mkdtempSync, rmSync, writeFileSync } from "node:fs"; import { tmpdir } from "node:os"; import { join } from "node:path"; import { Text } from "@earendil-works/pi-tui"; @@ -26,7 +26,7 @@ import { ExtensionRunner } from "../../packages/coding-agent/src/core/extensions import { SessionManager } from "../../packages/coding-agent/src/core/session-manager.ts"; import { DefaultResourceLoader } from "../../packages/coding-agent/src/core/resource-loader.ts"; import { SettingsManager } from "../../packages/coding-agent/src/core/settings-manager.ts"; -import { createAuthInteraction } from "../../packages/coding-agent/src/core/oauth-provider-bridge.ts"; +import { createAuthInteraction } from "../../packages/coding-agent/src/core/oauth-login.ts"; import { initTheme } from "../../packages/coding-agent/src/modes/interactive/theme/theme.ts"; import { stripAnsi } from "../../packages/coding-agent/src/utils/ansi.ts"; @@ -39,8 +39,8 @@ test("before_provider_headers handlers mutate headers sequentially and report er api.on("before_provider_headers", (event) => { event.headers["x-second"] = `${event.headers["x-first"]}-two`; }); api.on("before_provider_headers", () => { throw new Error("header failure"); }); }, process.cwd(), createEventBus(), runtime); - const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); - const runner = new ExtensionRunner([extension], runtime, process.cwd(), SessionManager.inMemory(), registry); + const modelRuntime = await ModelRuntime.create({ modelsPath: null }); + const runner = new ExtensionRunner([extension], runtime, process.cwd(), SessionManager.inMemory(), new ModelRegistry(modelRuntime)); const errors: string[] = []; runner.onError((error) => errors.push(error.error)); @@ -53,7 +53,8 @@ test("entry renderer registration is discoverable", async () => { const runtime = createExtensionRuntime(); const renderer: EntryRenderer = () => new Text("entry", 0, 0); const extension = await loadExtensionFromFactory((api) => api.registerEntryRenderer("state", renderer), process.cwd(), createEventBus(), runtime); - const runner = new ExtensionRunner([extension], runtime, process.cwd(), SessionManager.inMemory(), ModelRegistry.inMemory(AuthStorage.inMemory())); + const modelRuntime = await ModelRuntime.create({ modelsPath: null }); + const runner = new ExtensionRunner([extension], runtime, process.cwd(), SessionManager.inMemory(), new ModelRegistry(modelRuntime)); assert.equal(runner.getEntryRenderer("state"), renderer); assert.equal(runner.getEntryRenderer("missing"), undefined); }); @@ -80,15 +81,14 @@ test("CustomEntryComponent renders, suppresses empty output, propagates expansio }); test("native provider registration round-trips through ModelRegistry and ModelRuntime", async () => { - const auth = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(auth); + const runtime = await ModelRuntime.create({ modelsPath: null }); + const registry = new ModelRegistry(runtime); const builtin = registry.getProvider("anthropic"); assert.ok(builtin); const provider = { ...builtin, id: "native-test", name: "Native Test" }; registry.registerProvider(provider); assert.equal(registry.getProvider("native-test"), provider); - const runtime = new ModelRuntime(registry, auth); assert.equal(runtime.getProvider("native-test"), provider); assert.ok(runtime.getProviders().some((candidate) => candidate.id === "native-test")); assert.equal(await runtime.checkAuth("native-test"), undefined); @@ -99,8 +99,7 @@ test("native provider registration round-trips through ModelRegistry and ModelRu test("ModelRuntime delegates enumeration, runtime auth, snapshots, and refresh", async () => { const auth = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(auth); - const runtime = new ModelRuntime(registry, auth); + const runtime = await ModelRuntime.create({ credentials: auth, modelsPath: null }); const model = runtime.getModels("anthropic")[0]; assert.ok(model); await runtime.setRuntimeApiKey("anthropic", "runtime-secret", { allowNetwork: false }); @@ -113,16 +112,19 @@ test("ModelRuntime delegates enumeration, runtime auth, snapshots, and refresh", test("ModelRuntime delegates provider login and logout", async () => { const auth = AuthStorage.inMemory(); - const runtime = new ModelRuntime(ModelRegistry.inMemory(auth), auth); + const previousOffline = process.env.ATOMIC_OFFLINE; + process.env.ATOMIC_OFFLINE = "1"; + const runtime = await ModelRuntime.create({ credentials: auth, modelsPath: null }); + if (previousOffline === undefined) delete process.env.ATOMIC_OFFLINE; else process.env.ATOMIC_OFFLINE = previousOffline; const interaction = createAuthInteraction({ onAuth() {}, onDeviceCode() {}, onProgress() {}, async onPrompt() { return "sk-test"; }, async onSelect() { return ""; }, }); assert.deepEqual(await runtime.login("anthropic", "api_key", interaction), { type: "api_key", key: "sk-test" }); - assert.deepEqual(auth.get("anthropic"), { type: "api_key", key: "sk-test" }); + assert.deepEqual(await auth.read("anthropic"), { type: "api_key", key: "sk-test" }); await runtime.logout("anthropic"); - assert.equal(auth.get("anthropic"), undefined); + assert.equal(await auth.read("anthropic"), undefined); }); test("session context helpers preserve custom-message and compaction semantics", () => { @@ -135,12 +137,12 @@ test("session context helpers preserve custom-message and compaction semantics", assert.equal(sessionEntryToContextMessages(entries[1]).length, 0); }); -test("readStoredCredential reads a single credential from an explicit Atomic auth path", () => { +test("readStoredCredential reads a single credential from an explicit Atomic auth path", async () => { const dir = mkdtempSync(join(tmpdir(), "atomic-auth-")); const path = join(dir, "auth.json"); try { const storage = AuthStorage.create(path); - storage.set("example", { type: "api_key", key: "secret" }); + await storage.modify("example", async () => ({ type: "api_key", key: "secret" })); assert.deepEqual(readStoredCredential("example", path), { type: "api_key", key: "secret" }); assert.equal(readStoredCredential("missing", path), undefined); } finally { @@ -148,6 +150,17 @@ test("readStoredCredential reads a single credential from an explicit Atomic aut } }); +test("readStoredCredential returns undefined for malformed auth.json", () => { + const dir = mkdtempSync(join(tmpdir(), "atomic-auth-malformed-")); + const path = join(dir, "auth.json"); + try { + writeFileSync(path, "{not valid json", "utf8"); + assert.equal(readStoredCredential("example", path), undefined); + } finally { + rmSync(dir, { recursive: true, force: true }); + } +}); + test("entry_appended is emitted and agent_settled follows agent_end", async () => { const dir = mkdtempSync(join(tmpdir(), "atomic-group4-session-")); const provider = "group4-fixture"; @@ -180,8 +193,8 @@ test("entry_appended is emitted and agent_settled follows agent_end", async () = const loader = new DefaultResourceLoader({ cwd: dir, agentDir: dir, settingsManager: settings, extensionFactories: [inline], builtinPackagePaths: [], noSkills: true, noPromptTemplates: true, noThemes: true, noContextFiles: true }); await loader.reload(); const auth = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(auth); - const { session } = await createAgentSession({ cwd: dir, agentDir: dir, settingsManager: settings, resourceLoader: loader, authStorage: auth, modelRegistry: registry, sessionManager: SessionManager.inMemory(dir), model, noTools: "all" }); + const modelRuntime = await ModelRuntime.create({ credentials: auth, modelsPath: null }); + const { session } = await createAgentSession({ cwd: dir, agentDir: dir, settingsManager: settings, resourceLoader: loader, modelRuntime, sessionManager: SessionManager.inMemory(dir), model, noTools: "all" }); session.subscribe((event) => observed.push(event.type)); assert.ok(extensionApi); extensionApi.appendEntry("state", { ready: true }); diff --git a/test/unit/rpc-model-refresh.test.ts b/test/unit/rpc-model-refresh.test.ts index 429b80354..5eb532500 100644 --- a/test/unit/rpc-model-refresh.test.ts +++ b/test/unit/rpc-model-refresh.test.ts @@ -6,8 +6,8 @@ import { join } from "node:path"; import type { AgentSession } from "../../packages/coding-agent/src/core/agent-session.ts"; import type { AgentSessionRuntime } from "../../packages/coding-agent/src/core/agent-session-runtime.ts"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.ts"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.ts"; -import { getOAuthProviderMetadata } from "../../packages/coding-agent/src/core/oauth-provider-bridge.ts"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.ts"; +import { clearApiKeyCache } from "../../packages/coding-agent/src/core/provider-composer.ts"; import { RemoteModelCatalog } from "../../packages/coding-agent/src/modes/interactive-engine/remote-model-catalog.ts"; import type { RpcClient } from "../../packages/coding-agent/src/modes/rpc/rpc-client.ts"; import { createRpcCommandHandler } from "../../packages/coding-agent/src/modes/rpc/rpc-command-handler.ts"; @@ -41,8 +41,8 @@ for (const scenario of [ try { const authPath = join(tempDir, "auth.json"); const childAuth = AuthStorage.create(authPath); - const childRegistry = ModelRegistry.create(childAuth, join(tempDir, "models.json")); - const session = { modelRegistry: childRegistry, scopedModels: [] } as unknown as AgentSession; + const childRuntime = await ModelRuntime.create({ credentials: childAuth, modelsPath: null }); + const session = { modelRuntime: childRuntime, scopedModels: [] } as unknown as AgentSession; const handle = createRpcCommandHandler({ runtimeHost: { services: { agentDir: tempDir } } as unknown as AgentSessionRuntime, getSession: () => session, @@ -50,66 +50,108 @@ for (const scenario of [ output: () => {}, }); - assert.equal(childAuth.hasAuth(scenario.provider), false); - assert.equal(childRegistry.getAvailable().some((model) => model.provider === scenario.provider), false); + assert.equal(await childAuth.read(scenario.provider), undefined); + assert.equal(childRuntime.getAvailableSnapshot().some((model) => model.provider === scenario.provider), false); const hostAuth = AuthStorage.create(authPath); - hostAuth.set(scenario.provider, scenario.credential); - assert.equal(childAuth.hasAuth(scenario.provider), false, "child keeps its startup auth snapshot before RPC refresh"); + await hostAuth.modify(scenario.provider, async () => scenario.credential); + assert.equal(await childAuth.read(scenario.provider), undefined); const response = await handle({ id: scenario.provider, type: "refresh_models", allowNetwork: false }); - assert.equal(childAuth.hasAuth(scenario.provider), true); + assert.equal((await childAuth.read(scenario.provider))?.type, scenario.credential.type); assert.equal(availableProviders(response).has(scenario.provider), true); - assert.deepEqual(response && "data" in response ? response.data : undefined, { - aborted: false, - errors: [], - models: childRegistry.getAvailable(), - scopedModels: [], - customAuthProviders: [], - oauthProviders: getOAuthProviderMetadata(), - }); } finally { rmSync(tempDir, { recursive: true, force: true }); } }); } -test("refresh_models uses a newly persisted credential for dynamic model discovery", async () => { +test("refresh_models uses newly persisted credentials for forced dynamic discovery", async () => { const tempDir = mkdtempSync(join(tmpdir(), "atomic-rpc-dynamic-model-refresh-")); try { const authPath = join(tempDir, "auth.json"); const childAuth = AuthStorage.create(authPath); - const childRegistry = ModelRegistry.create(childAuth, join(tempDir, "models.json")); - const template = ModelRegistry.inMemory(AuthStorage.inMemory({ - "kimi-coding": { type: "api_key", key: "template-key" }, - })).getAvailable().find((model) => model.provider === "kimi-coding"); - assert.ok(template); + const childRuntime = await ModelRuntime.create({ credentials: childAuth, modelsPath: null }); let observedKey: string | undefined; - childRegistry.registerProvider("dynamic-login", { - refreshModels: async ({ credential, allowNetwork, force }) => { + let observedForce: boolean | undefined; + childRuntime.registerProvider("dynamic-login", { + baseUrl: "https://dynamic.test/v1", + api: "openai-completions", + apiKey: "$DYNAMIC_LOGIN_KEY", + models: [], + refreshModels: async ({ credential, force }) => { observedKey = credential?.type === "api_key" ? credential.key : undefined; - assert.equal(allowNetwork, false); - assert.equal(force, true); - return observedKey ? [{ ...template, provider: "dynamic-login", id: "discovered-after-login" }] : []; + observedForce = force; + return observedKey ? [{ + id: "discovered-after-login", + name: "Discovered", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128_000, + maxTokens: 4_096, + }] : []; }, }); - const session = { modelRegistry: childRegistry, scopedModels: [] } as unknown as AgentSession; + const session = { modelRuntime: childRuntime, scopedModels: [] } as unknown as AgentSession; const handle = createRpcCommandHandler({ runtimeHost: { services: { agentDir: tempDir } } as unknown as AgentSessionRuntime, getSession: () => session, rebindSession: async () => {}, output: () => {}, }); - AuthStorage.create(authPath).set("dynamic-login", { type: "api_key", key: "new-dynamic-key" }); + await AuthStorage.create(authPath).modify("dynamic-login", async () => ({ type: "api_key", key: "new-dynamic-key" })); + assert.equal(await childAuth.read("dynamic-login"), undefined); + const response = await handle({ type: "refresh_models", allowNetwork: false, force: true }); + + assert.deepEqual(await childAuth.read("dynamic-login"), { type: "api_key", key: "new-dynamic-key" }); assert.equal(observedKey, "new-dynamic-key"); + assert.equal(observedForce, true); assert.equal(availableProviders(response).has("dynamic-login"), true); } finally { rmSync(tempDir, { recursive: true, force: true }); } }); +test("refresh_models is unbounded by default and honors an explicit caller timeout", async () => { + const signals: Array = []; + let reloads = 0; + const modelRuntime = { + reloadCredentials: async () => { reloads += 1; }, + refresh: async ({ signal }: { signal?: AbortSignal }) => { + signals.push(signal); + if (signal && !signal.aborted) { + await new Promise((resolve) => signal.addEventListener("abort", () => resolve(), { once: true })); + } + return { aborted: signal?.aborted ?? false, errors: new Map() }; + }, + getAvailableSnapshot: () => [], + getProviders: () => [], + getOAuthProviderMetadata: () => [], + }; + const session = { modelRuntime, scopedModels: [] } as unknown as AgentSession; + const handle = createRpcCommandHandler({ + runtimeHost: {} as AgentSessionRuntime, + getSession: () => session, + rebindSession: async () => {}, + output: () => {}, + }); + + const unbounded = await handle({ type: "refresh_models", allowNetwork: false }); + const bounded = await handle({ type: "refresh_models", allowNetwork: false, timeoutMs: 5 }); + + assert.equal(unbounded?.success, true); + assert.equal(signals[0], undefined); + assert.ok(signals[1] instanceof AbortSignal); + assert.equal(signals[1].aborted, true); + assert.ok(bounded?.success); + assert.equal(bounded.command, "refresh_models"); + assert.ok("data" in bounded); + assert.equal(bounded.data.aborted, true); + assert.equal(reloads, 2); +}); test("bearer-only Anthropic models survive RPC and isolated catalog transport", async () => { const originalAuthToken = process.env.ANTHROPIC_AUTH_TOKEN; const originalApiKey = process.env.ANTHROPIC_API_KEY; @@ -118,8 +160,8 @@ test("bearer-only Anthropic models survive RPC and isolated catalog transport", process.env.ANTHROPIC_AUTH_TOKEN = "gateway-token"; delete process.env.ANTHROPIC_API_KEY; delete process.env.ANTHROPIC_OAUTH_TOKEN; - const hostRegistry = ModelRegistry.inMemory(AuthStorage.inMemory()); - const hostSession = { modelRegistry: hostRegistry, scopedModels: [] } as unknown as AgentSession; + const hostRuntime = await ModelRuntime.create({ modelsPath: null }); + const hostSession = { modelRuntime: hostRuntime, scopedModels: [] } as unknown as AgentSession; const handle = createRpcCommandHandler({ runtimeHost: { services: { agentDir: "/tmp/atomic-anthropic-bearer-rpc" } } as unknown as AgentSessionRuntime, getSession: () => hostSession, @@ -134,13 +176,13 @@ test("bearer-only Anthropic models survive RPC and isolated catalog transport", assert.ok(response.data.models.some((model) => model.provider === "anthropic")); delete process.env.ANTHROPIC_AUTH_TOKEN; - const isolatedRegistry = ModelRegistry.inMemory(AuthStorage.inMemory()); - const isolatedSession = { modelRegistry: isolatedRegistry, scopedModels: [] } as unknown as AgentSession; - assert.equal(isolatedRegistry.getAvailable().some((model) => model.provider === "anthropic"), false); + clearApiKeyCache(); + const isolatedRuntime = await ModelRuntime.create({ modelsPath: null }); + const isolatedSession = { modelRuntime: isolatedRuntime, scopedModels: [] } as unknown as AgentSession; const remoteCatalog = new RemoteModelCatalog({} as RpcClient); remoteCatalog.apply(response.data); remoteCatalog.patch(isolatedSession); - assert.equal(isolatedRegistry.getAvailable().some((model) => model.provider === "anthropic"), true); + assert.deepEqual(isolatedRuntime.getAvailableSnapshot(), response.data.models); } finally { if (originalAuthToken === undefined) delete process.env.ANTHROPIC_AUTH_TOKEN; else process.env.ANTHROPIC_AUTH_TOKEN = originalAuthToken; diff --git a/test/unit/rpc-provider-login.test.ts b/test/unit/rpc-provider-login.test.ts index 93e3564b1..1e0702cb8 100644 --- a/test/unit/rpc-provider-login.test.ts +++ b/test/unit/rpc-provider-login.test.ts @@ -3,33 +3,22 @@ import assert from "node:assert/strict"; import type { AgentSession } from "../../packages/coding-agent/src/core/agent-session.ts"; import type { AgentSessionRuntime } from "../../packages/coding-agent/src/core/agent-session-runtime.ts"; import { AuthStorage } from "../../packages/coding-agent/src/core/auth-storage.ts"; -import { ModelRegistry } from "../../packages/coding-agent/src/core/model-registry.ts"; +import { ModelRuntime } from "../../packages/coding-agent/src/core/model-runtime.ts"; import { createRpcCommandHandler } from "../../packages/coding-agent/src/modes/rpc/rpc-command-handler.ts"; -function createSessionRegistry(login: (prompt: (message: string) => Promise) => Promise) { +async function createSessionRuntime() { const authStorage = AuthStorage.inMemory(); - const registry = ModelRegistry.inMemory(authStorage); - const template = ModelRegistry.inMemory(AuthStorage.inMemory({ - "kimi-coding": { type: "api_key", key: "template" }, - })).getAvailable().find((model) => model.provider === "kimi-coding"); + const runtime = await ModelRuntime.create({ credentials: authStorage, modelsPath: null }); + const template = runtime.getModels("kimi-coding")[0]; assert.ok(template); let refreshCount = 0; - registry.registerProvider("extension-login", { - auth: { - apiKey: { - name: "Extension Login", - login: async ({ prompt }) => ({ - type: "api_key", - key: await login((message) => prompt({ type: "secret", message, placeholder: "token" })), - }), - }, - }, + runtime.registerProvider("extension-login", { refreshModels: async ({ credential }) => { refreshCount += 1; - return credential ? [{ ...template, id: "extension-model" }] : []; + return credential ? [{ ...template, provider: "extension-login", id: "extension-model" }] : []; }, }); - return { authStorage, registry, refreshCount: () => refreshCount }; + return { authStorage, runtime, refreshCount: () => refreshCount }; } function runtimeHost(): AgentSessionRuntime { @@ -37,8 +26,8 @@ function runtimeHost(): AgentSessionRuntime { } test("login_provider prompts in the host, persists the credential, refreshes, and returns provider metadata", async () => { - const state = createSessionRegistry(async (prompt) => prompt("Enter extension token")); - const session = { modelRegistry: state.registry, scopedModels: [] } as unknown as AgentSession; + const state = await createSessionRuntime(); + const session = { modelRuntime: state.runtime, scopedModels: [] } as unknown as AgentSession; const handle = createRpcCommandHandler({ runtimeHost: runtimeHost(), getSession: () => session, @@ -46,8 +35,8 @@ test("login_provider prompts in the host, persists the credential, refreshes, an output: () => {}, inputForm: { open: async (request) => { - assert.equal(request.title, "Enter extension token"); - assert.equal(request.fields[0]?.placeholder, "token"); + assert.equal(request.title, "Enter API key"); + assert.equal(request.fields[0]?.placeholder, undefined); return { value: "child-secret" }; }, }, @@ -58,17 +47,17 @@ test("login_provider prompts in the host, persists the credential, refreshes, an assert.ok(response?.success && "data" in response); assert.equal(response.command, "login_provider"); assert.equal(response.data.cancelled, false); - assert.deepEqual(state.authStorage.get("extension-login"), { type: "api_key", key: "child-secret" }); + assert.deepEqual(await state.authStorage.read("extension-login"), { type: "api_key", key: "child-secret" }); assert.ok(state.refreshCount() > 0); if (!response.data.cancelled) { - assert.deepEqual(response.data.customAuthProviders, [{ id: "extension-login", name: "Extension Login" }]); + assert.deepEqual(response.data.customAuthProviders, []); assert.equal(response.data.models.some((model) => model.provider === "extension-login"), true); } }); test("cancel_login_provider aborts an active child login without storing credentials", async () => { - const state = createSessionRegistry(async (prompt) => prompt("Enter extension token")); - const session = { modelRegistry: state.registry, scopedModels: [] } as unknown as AgentSession; + const state = await createSessionRuntime(); + const session = { modelRuntime: state.runtime, scopedModels: [] } as unknown as AgentSession; const handle = createRpcCommandHandler({ runtimeHost: runtimeHost(), getSession: () => session, @@ -89,20 +78,18 @@ test("cancel_login_provider aborts an active child login without storing credent assert.ok(cancelled?.success); assert.ok(response?.success && "data" in response); assert.deepEqual(response.data, { provider: "extension-login", cancelled: true }); - assert.equal(state.authStorage.get("extension-login"), undefined); + assert.equal(await state.authStorage.read("extension-login"), undefined); }); -test("non-isolated registries retain the local custom authentication contract", async () => { - const state = createSessionRegistry(async (prompt) => prompt("Local prompt")); - const auth = state.registry.getCustomApiKeyAuth("extension-login"); - assert.ok(auth); - assert.equal(auth.name, "Extension Login"); - const credential = await auth.login({ +test("non-isolated runtimes retain provider-owned API-key authentication", async () => { + const state = await createSessionRuntime(); + const credential = await state.runtime.login("extension-login", "api_key", { signal: new AbortController().signal, prompt: async ({ message }) => { - assert.equal(message, "Local prompt"); + assert.equal(message, "Enter API key"); return "local-secret"; }, + notify: () => {}, }); assert.deepEqual(credential, { type: "api_key", key: "local-secret" }); }); diff --git a/test/unit/stage-chat-view-helpers.ts b/test/unit/stage-chat-view-helpers.ts index 404a487e8..765f4ca38 100644 --- a/test/unit/stage-chat-view-helpers.ts +++ b/test/unit/stage-chat-view-helpers.ts @@ -239,7 +239,7 @@ export function fakeFooterAgentSession(isStreaming = false): AgentSession { }, ], }, - modelRegistry: { + modelRuntime: { isUsingOAuth: () => false, }, settingsManager: { diff --git a/test/unit/stage-runner-fallback-shared-registry.test.ts b/test/unit/stage-runner-fallback-shared-registry.test.ts index 44912d9df..6d2bcbed9 100644 --- a/test/unit/stage-runner-fallback-shared-registry.test.ts +++ b/test/unit/stage-runner-fallback-shared-registry.test.ts @@ -1,5 +1,5 @@ /** - * Regression: within a single stage, model fallback must REUSE one ModelRegistry + * Regression: within a single stage, model fallback must REUSE one ModelRuntime * (and its already-loaded AuthStorage) across every candidate instead of letting * each fallback candidate build a fresh one. * @@ -8,7 +8,7 @@ * model 404'd: each fresh AuthStorage re-read auth.json under lock contention and * could silently fall back to an empty credential set (issue #1431). * - * This pins the stage-runner behavior: the registry captured from the first + * This pins the stage-runner behavior: the runtime captured from the first * session is threaded into the options of every subsequent fallback candidate. * * cross-ref: packages/workflows/src/runs/foreground/stage-runner.ts @@ -23,8 +23,8 @@ import type { } from "../../packages/workflows/src/runs/foreground/stage-runner.js"; interface FakeSessionConfig { - /** Marker object used to identify which ModelRegistry the session carries. */ - modelRegistry: unknown; + /** Marker object used to identify which ModelRuntime the session carries. */ + modelRuntime: unknown; /** When set, prompt() throws this error (to drive fallback). */ promptError?: Error; onTransfer?: (target: object) => void; @@ -87,10 +87,10 @@ function makeFakeStageSession(config: FakeSessionConfig): StageSessionRuntime { return undefined; }, }; - // The real SDK AgentSession exposes `.modelRegistry`; the stage runner reads it - // (via a structural cast) to capture/reuse the registry across candidates. + // The real SDK AgentSession exposes `.modelRuntime`; the stage runner reads it + // (via a structural cast) to capture/reuse the runtime across candidates. return Object.assign(base, { - modelRegistry: config.modelRegistry, + modelRuntime: config.modelRuntime, state: { messages: [] }, sessionManager: {}, getContextUsage: () => undefined, @@ -98,12 +98,12 @@ function makeFakeStageSession(config: FakeSessionConfig): StageSessionRuntime { }); } -describe("stage model fallback reuses one ModelRegistry across candidates (#1431)", () => { - test("threads the first session's registry into every later fallback candidate", async () => { - const registryA = { id: "registry-A" }; - const registryB = { id: "registry-B" }; +describe("stage model fallback reuses one ModelRuntime across candidates (#1431)", () => { + test("threads the first session's runtime into every later fallback candidate", async () => { + const runtimeA = { id: "runtime-A" }; + const runtimeB = { id: "runtime-B" }; - const createdWith: Array<{ model: unknown; modelRegistry: unknown }> = []; + const createdWith: Array<{ model: unknown; modelRuntime: unknown }> = []; let createCount = 0; let fallbackDeliveryTarget: object | undefined; let sharedOrchestrationContext: StageSessionCreateOptions["orchestrationContext"]; @@ -120,7 +120,7 @@ describe("stage model fallback reuses one ModelRegistry across candidates (#1431 agentSession: { async create(options: StageSessionCreateOptions, meta) { createCount += 1; - createdWith.push({ model: options?.model, modelRegistry: options?.modelRegistry }); + createdWith.push({ model: options?.model, modelRuntime: options?.modelRuntime }); if (createCount === 1) { sharedOrchestrationContext = { kind: "workflow-stage", @@ -130,7 +130,7 @@ describe("stage model fallback reuses one ModelRegistry across candidates (#1431 constraints: { disableWorkflowTool: true, maxSubagentDepth: 5 }, }; return makeFakeStageSession({ - modelRegistry: registryA, + modelRuntime: runtimeA, orchestrationContext: sharedOrchestrationContext, promptError: new Error("rate limit exceeded"), onTransfer: (target) => { fallbackDeliveryTarget = target; }, @@ -138,7 +138,7 @@ describe("stage model fallback reuses one ModelRegistry across candidates (#1431 } assert.equal(meta?.orchestrationContext, sharedOrchestrationContext); return makeFakeStageSession({ - modelRegistry: registryB, + modelRuntime: runtimeB, orchestrationContext: sharedOrchestrationContext, }); } @@ -155,13 +155,13 @@ describe("stage model fallback reuses one ModelRegistry across candidates (#1431 assert.ok(fallbackDeliveryTarget, "failed candidate must transfer detached deliveries to its replacement"); await (fallbackDeliveryTarget as StageSessionRuntime).followUp("completion received while stage remains open"); - // The primary candidate builds its own registry (none injected). + // The primary candidate builds its own runtime (none injected). assert.equal(createdWith[0]?.model, "anthropic/model-a"); - assert.equal(createdWith[0]?.modelRegistry, undefined); + assert.equal(createdWith[0]?.modelRuntime, undefined); - // The fallback candidate is created WITH the registry captured from the + // The fallback candidate is created WITH the runtime captured from the // first session — not a fresh one. assert.equal(createdWith[1]?.model, "anthropic/model-b"); - assert.equal(createdWith[1]?.modelRegistry, registryA); + assert.equal(createdWith[1]?.modelRuntime, runtimeA); }); }); diff --git a/test/unit/workflow-idle-prompt-start-race.test.ts b/test/unit/workflow-idle-prompt-start-race.test.ts index e0d1b6c3e..185818c48 100644 --- a/test/unit/workflow-idle-prompt-start-race.test.ts +++ b/test/unit/workflow-idle-prompt-start-race.test.ts @@ -58,7 +58,7 @@ test("production prompt wiring holds idle admission until the first agent turn s }, _flushPendingBashMessages() {}, model: { provider: "test", id: "test" }, - _modelRegistry: { hasConfiguredAuth: () => true }, + _modelRuntime: { hasConfiguredAuth: () => true }, _findLastAssistantMessage: () => undefined, _pendingNextTurnMessages: [] as AgentMessage[], _baseSystemPrompt: "base", diff --git a/test/unit/workflow-model-catalog-context.test.ts b/test/unit/workflow-model-catalog-context.test.ts new file mode 100644 index 000000000..71ec90101 --- /dev/null +++ b/test/unit/workflow-model-catalog-context.test.ts @@ -0,0 +1,67 @@ +import { test } from "bun:test"; +import assert from "node:assert/strict"; +import type { Api, Model } from "@earendil-works/pi-ai/compat"; +import { WORKFLOW_STAGE_SUBAGENT_GUARD_ENV } from "@bastani/atomic"; +import type { ExtensionRuntime } from "../../packages/workflows/src/extension/runtime.js"; +import { workflowModelCatalogFromContext } from "../../packages/workflows/src/extension/workflow-model-catalog.js"; +import type { PiExecuteContext } from "../../packages/workflows/src/extension/public-types.js"; +import { makeExecuteWorkflowTool } from "../../packages/workflows/src/extension/workflow-tool.js"; + +function model(provider: string, id: string): Model { + return { + provider, + id, + name: id, + api: "anthropic-messages", + baseUrl: "https://example.invalid", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 8_192, + maxTokens: 1_024, + }; +} + +test("workflow models action lists the host model registry catalog", async () => { + const previousGuard = process.env[WORKFLOW_STAGE_SUBAGENT_GUARD_ENV]; + delete process.env[WORKFLOW_STAGE_SUBAGENT_GUARD_ENV]; + try { + const current = model("provider-a", "current"); + const alternate = model("provider-b", "alternate"); + const execute = makeExecuteWorkflowTool({} as ExtensionRuntime, () => undefined); + const context = { + model: current, + modelRegistry: { getAvailable: () => [current, alternate] }, + } as unknown as PiExecuteContext; + + const result = await execute({ action: "models" }, context); + + assert.equal(result.action, "models"); + if (result.action !== "models") return; + assert.deepEqual(result.models.map(({ fullId }) => fullId), [ + "provider-a/current", + "provider-b/alternate", + ]); + assert.deepEqual(result.models.map(({ isCurrent }) => isCurrent), [true, false]); + } finally { + if (previousGuard === undefined) delete process.env[WORKFLOW_STAGE_SUBAGENT_GUARD_ENV]; + else process.env[WORKFLOW_STAGE_SUBAGENT_GUARD_ENV] = previousGuard; + } +}); + +test("workflow stage model catalog includes alternatives beyond the current model", async () => { + const current = model("provider-a", "current"); + const alternate = model("provider-b", "alternate"); + const catalog = workflowModelCatalogFromContext({ + model: current, + modelRegistry: { getAvailable: () => [current, alternate] }, + }); + + assert.ok(catalog); + const available = await catalog.listModels(); + assert.deepEqual(available.map(({ fullId }) => fullId), [ + "provider-a/current", + "provider-b/alternate", + ]); + assert.equal(catalog.currentModel, current); +});