pi-llama-cpp 0.10.0 → 0.12.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (41) hide show
  1. package/README.md +187 -10
  2. package/package.json +6 -6
  3. package/src/api/client.ts +76 -32
  4. package/src/constants.ts +15 -0
  5. package/src/index.ts +3 -5
  6. package/src/interfaces/events.ts +14 -3
  7. package/src/interfaces/server.ts +27 -0
  8. package/src/interfaces/settings.ts +76 -2
  9. package/src/managers/command.ts +369 -50
  10. package/src/managers/events.ts +28 -10
  11. package/src/managers/server.ts +90 -16
  12. package/src/managers/settings.ts +181 -44
  13. package/src/models/baseModel.ts +37 -15
  14. package/src/models/routerModel.ts +2 -1
  15. package/src/server.ts +103 -28
  16. package/src/sse/client.ts +28 -16
  17. package/src/sse/manager.ts +26 -13
  18. package/src/ui/dialog.ts +287 -0
  19. package/src/ui/overrideEntryEditor.ts +119 -0
  20. package/src/ui/overrideSettingsList.ts +682 -0
  21. package/src/ui/serverListEditor.ts +32 -0
  22. package/src/ui/serverSettingsList.ts +466 -0
  23. package/src/ui/strings.ts +127 -0
  24. package/src/utils/errors.ts +5 -0
  25. package/src/utils/settingsStore.ts +56 -0
  26. package/src/utils/urls.ts +16 -0
  27. package/tests/commandManager.test.ts +346 -11
  28. package/tests/dialog.test.ts +186 -0
  29. package/tests/events.test.ts +120 -88
  30. package/tests/legacyModel.test.ts +4 -19
  31. package/tests/mocks.ts +149 -32
  32. package/tests/overrides.test.ts +352 -0
  33. package/tests/server.test.ts +42 -40
  34. package/tests/serverManager.test.ts +264 -55
  35. package/tests/settings.test.ts +654 -51
  36. package/tests/settingsStore.test.ts +190 -0
  37. package/tests/singleModel.test.ts +32 -0
  38. package/tests/sseManager.test.ts +88 -11
  39. package/src/interfaces/auth.ts +0 -6
  40. package/src/utils/cache.ts +0 -39
  41. package/src/utils/mutex.ts +0 -24
@@ -1,15 +1,27 @@
1
1
  import type { ExtensionAPI } from "@earendil-works/pi-coding-agent";
2
- import { API_TYPE, PROVIDER_NAME } from "../constants";
2
+ import { API_TYPE, PROVIDER_NAME, type SortBy } from "../constants";
3
3
  import { ServerStatus } from "../enums/serverStatus";
4
4
  import { BaseModel } from "../models/baseModel";
5
5
  import { Server } from "../server";
6
- import { settings } from "./settings";
6
+ import type { LlamaSettingsManager } from "./settings";
7
+
8
+ /** Model-list comparator: negative if a sorts first, positive if b does. */
9
+ type ModelComparator = (a: BaseModel, b: BaseModel) => number;
7
10
 
8
11
  export class ServerManager {
12
+ constructor(private readonly settings: LlamaSettingsManager) {}
9
13
  readonly failedUrls: string[] = [];
10
14
  private readonly warnings: string[] = [];
15
+ private readonly serverList: Server[] = [];
11
16
 
12
- constructor(private readonly servers: Server[]) {}
17
+ /**
18
+ * Live view of the server list. `update()` re-derives the list from
19
+ * settings on every scan (in place), so `/models servers` edits apply
20
+ * without a restart.
21
+ */
22
+ get servers(): readonly Server[] {
23
+ return this.serverList;
24
+ }
13
25
 
14
26
  /**
15
27
  * Verifies reachability of servers and registers the providers
@@ -18,7 +30,7 @@ export class ServerManager {
18
30
  */
19
31
  async initialize(pi: ExtensionAPI) {
20
32
  // Register the providers with the configured server timeout
21
- const { serverTimeout } = settings.resolveTimeouts();
33
+ const { serverTimeout } = await this.settings.resolveTimeouts();
22
34
  await this.update(pi, serverTimeout);
23
35
  }
24
36
 
@@ -32,6 +44,33 @@ export class ServerManager {
32
44
  async update(pi: ExtensionAPI, timeout?: number) {
33
45
  this.failedUrls.length = 0;
34
46
 
47
+ // Surface warnings from strict URL parsing (dropped invalid entries)
48
+ this.warnings.push(...this.settings.takeWarnings());
49
+
50
+ // Re-derive the server list from settings so `/models servers` edits
51
+ // (add / remove / URL / id / name) apply on the next scan
52
+ const fresh: Server[] = [];
53
+ const seen = new Set<string>(); // dedupe repeated URLs (same providerId)
54
+ for (const server of await this.settings.resolveServers()) {
55
+ if (seen.has(server.providerId)) continue;
56
+ seen.add(server.providerId);
57
+ fresh.push(server);
58
+ }
59
+
60
+ // Unregister providers that disappeared (removed or edited away);
61
+ // no-op for providers that were never registered
62
+ for (const old of this.servers) {
63
+ if (fresh.some((f) => f.providerId === old.providerId)) continue;
64
+ pi.unregisterProvider(old.providerId);
65
+ // Optional chain is intentional despite the non-optional type: `sse`
66
+ // is undefined until initialize() runs (async-constructor hack — see Server)
67
+ old.sseManager?.disconnect();
68
+ }
69
+
70
+ // Replace in place so the live `servers` view stays valid (D1)
71
+ this.serverList.length = 0;
72
+ this.serverList.push(...fresh);
73
+
35
74
  const registrableServers = timeout
36
75
  ? await this.findRegistrableServers(timeout)
37
76
  : this.servers;
@@ -95,7 +134,7 @@ export class ServerManager {
95
134
  */
96
135
  private async registerProvider(server: Server, pi: ExtensionAPI) {
97
136
  const { baseUrl, models, providerId, providerName } = server;
98
- const apiKey = await server.getApiKey();
137
+ const apiKey = server.getApiKey();
99
138
  const modelConfigs = await Promise.all(
100
139
  models.map((m) => m.toProviderConfig()),
101
140
  );
@@ -123,26 +162,61 @@ export class ServerManager {
123
162
  * Returns the server for a given model.
124
163
  *
125
164
  * @param model - The model to find the server for
126
- * @returns The server containing the model
165
+ * @returns The server containing the model, or `undefined` when no
166
+ * current server matches (e.g. removed while a model was loading)
127
167
  */
128
- getServer(model: BaseModel): Server {
129
- return this.servers.find((s) => s.baseUrl === model.serverUrl)!;
168
+ getServer(model: BaseModel): Server | undefined {
169
+ return this.servers.find((s) => s.baseUrl === model.serverUrl);
130
170
  }
131
171
 
132
172
  /**
133
- * Returns all models from all servers.
173
+ * Returns all models from all servers, sorted by the configured sort mode.
174
+ * Servers maintain their order from `llamaSettings`; sorting only applies
175
+ * to models within each server.
134
176
  *
135
177
  * @returns Flat array of all models across all servers
136
178
  */
137
- getAllModels(): BaseModel[] {
138
- const response = [];
179
+ async getAllModels(): Promise<BaseModel[]> {
180
+ const sortBy = await this.settings.resolveSortBy();
139
181
 
140
- for (const { models } of this.servers) {
141
- for (const model of models) {
142
- response.push(model);
143
- }
182
+ if (sortBy === "api") {
183
+ return this.servers.flatMap((s) => s.models);
144
184
  }
145
185
 
146
- return response;
186
+ const sorter = ServerManager.SORTERS[sortBy];
187
+ return this.servers.flatMap((s) => [...s.models].sort(sorter));
147
188
  }
189
+
190
+ private static sortByIdAsc(a: BaseModel, b: BaseModel): number {
191
+ return a.id.localeCompare(b.id);
192
+ }
193
+
194
+ private static sortByIdDesc(a: BaseModel, b: BaseModel): number {
195
+ return b.id.localeCompare(a.id);
196
+ }
197
+
198
+ /** Name ascending, with ID as tiebreaker. */
199
+ private static sortByNameAsc(a: BaseModel, b: BaseModel): number {
200
+ const cmp = a.name.localeCompare(b.name);
201
+ return cmp !== 0 ? cmp : a.id.localeCompare(b.id);
202
+ }
203
+
204
+ /** Name descending, with ID as tiebreaker. */
205
+ private static sortByNameDesc(a: BaseModel, b: BaseModel): number {
206
+ const cmp = b.name.localeCompare(a.name);
207
+ return cmp !== 0 ? cmp : a.id.localeCompare(b.id);
208
+ }
209
+
210
+ /**
211
+ * Comparators for sorting models within each server.
212
+ */
213
+ private static readonly SORTERS: Record<
214
+ Exclude<SortBy, "api">,
215
+ ModelComparator
216
+ > = {
217
+ asc: ServerManager.sortByIdAsc,
218
+ desc: ServerManager.sortByIdDesc,
219
+ "asc-name": ServerManager.sortByNameAsc,
220
+ "desc-name": ServerManager.sortByNameDesc,
221
+ };
148
222
  }
@@ -1,8 +1,11 @@
1
1
  import { ApiKeyCredential, ModelThinkingLevel } from "@earendil-works/pi-ai";
2
2
  import {
3
+ getAgentDir,
3
4
  readStoredCredential,
4
5
  SettingsManager,
5
6
  } from "@earendil-works/pi-coding-agent";
7
+ import { access } from "node:fs/promises";
8
+ import { join } from "node:path";
6
9
  import {
7
10
  API_KEY_PLACEHOLDER,
8
11
  AUTOLOAD_ON_MESSAGE,
@@ -10,32 +13,78 @@ import {
10
13
  POLLING_TIMEOUT,
11
14
  REACT_TO_MODEL_SELECT,
12
15
  SERVER_TIMEOUT,
16
+ SETTINGS_KEY,
17
+ SORT_BY,
13
18
  THINKING_BUDGETS,
19
+ type SortBy,
14
20
  } from "../constants";
15
- import { LlamaSettings } from "../interfaces/settings";
21
+ import {
22
+ LlamaServer,
23
+ LlamaSettings,
24
+ ModelOverride,
25
+ } from "../interfaces/settings";
16
26
  import { Server } from "../server";
17
-
18
- const SETTINGS_KEY = "llamaSettings";
27
+ import { SettingsStore } from "../utils/settingsStore";
28
+ import { isValidServerUrl, normalizeUrl } from "../utils/urls";
19
29
 
20
30
  export class LlamaSettingsManager {
21
31
  private settingsManager = SettingsManager.create(process.cwd());
22
32
 
33
+ private globalStore = new SettingsStore(join(getAgentDir(), "settings.json"));
34
+ private projectStore = new SettingsStore(
35
+ join(process.cwd(), ".pi", "settings.json"),
36
+ );
37
+
38
+ /**
39
+ * Check if project settings file exists in the current working directory.
40
+ */
41
+ private async hasProjectSettings(): Promise<boolean> {
42
+ try {
43
+ await access(join(process.cwd(), ".pi", "settings.json"));
44
+ return true;
45
+ } catch {
46
+ return false;
47
+ }
48
+ }
49
+
50
+ /** Warnings collected during URL resolution (dropped invalid entries). */
51
+ private warnings: string[] = [];
52
+
23
53
  /**
24
- * Convenience getter for merged project/global settings
54
+ * Returns and clears warnings collected during URL resolution.
25
55
  */
26
- private get mergedSettings(): Record<string, any> {
27
- const merged = {
56
+ takeWarnings(): string[] {
57
+ const warnings = [...this.warnings];
58
+ this.warnings.length = 0;
59
+ return warnings;
60
+ }
61
+
62
+ /**
63
+ * Reloads settings from disk and returns merged project/global settings.
64
+ * Project settings override global settings.
65
+ */
66
+ private async getMergedSettings(): Promise<Record<string, any>> {
67
+ await this.settingsManager.reload();
68
+ return {
28
69
  ...this.settingsManager.getGlobalSettings(),
29
70
  ...this.settingsManager.getProjectSettings(),
30
71
  } as Record<string, any>;
31
- return merged;
32
72
  }
33
73
 
34
74
  /**
35
- * Convenience getter for the `llamaSettings` key
75
+ * Convenience method for the `llamaSettings` key.
76
+ * Reloads settings from disk before reading.
36
77
  */
37
- private get llamaSettings(): LlamaSettings {
38
- return this.mergedSettings[SETTINGS_KEY] ?? {};
78
+ async getLlamaSettings(): Promise<LlamaSettings> {
79
+ return (await this.getMergedSettings())[SETTINGS_KEY] ?? {};
80
+ }
81
+
82
+ /**
83
+ * Convenience method for the merged `servers` list (project overrides
84
+ * global, per-key merge). Reloads settings from disk before reading.
85
+ */
86
+ async getLlamaServers(): Promise<LlamaServer[]> {
87
+ return (await this.getLlamaSettings()).servers ?? [];
39
88
  }
40
89
 
41
90
  /**
@@ -48,14 +97,14 @@ export class LlamaSettingsManager {
48
97
  *
49
98
  * @returns The list of URLs to use
50
99
  */
51
- resolveUrls(): string[] {
100
+ async resolveUrls(): Promise<string[]> {
52
101
  let response = this.resolveEnvUrls();
53
102
  if (response.length > 0) return response;
54
103
 
55
- response = this.resolveServerUrls();
104
+ response = await this.resolveServerUrls();
56
105
  if (response.length > 0) return response;
57
106
 
58
- response = this.resolveLegacyUrls();
107
+ response = await this.resolveLegacyUrls();
59
108
  if (response.length > 0) return response;
60
109
 
61
110
  return [LLAMA_SERVER_URL];
@@ -76,22 +125,24 @@ export class LlamaSettingsManager {
76
125
  /**
77
126
  * Resolves the llama-server URLs from `llamaSettings.servers`.
78
127
  * Settings are merged, prioritizing project over global settings.
128
+ * Reloads settings from disk before reading.
79
129
  *
80
130
  * @returns A list of detected URLs
81
131
  */
82
- private resolveServerUrls(): string[] {
83
- const { servers = [] } = this.llamaSettings;
132
+ private async resolveServerUrls(): Promise<string[]> {
133
+ const { servers = [] } = await this.getLlamaSettings();
84
134
  return servers.map((s) => this.parseUrls(s.url)).flat();
85
135
  }
86
136
 
87
137
  /**
88
- * Resolves the llama-server URLs from `llamaSettings.servers`.
138
+ * Resolves the llama-server URLs from `llamaServerUrl` legacy key.
89
139
  * Settings are merged, prioritizing project over global settings.
140
+ * Reloads settings from disk before reading.
90
141
  *
91
142
  * @returns A list of detected URLs
92
143
  */
93
- private resolveLegacyUrls(): string[] {
94
- const { llamaServerUrl = null } = this.mergedSettings;
144
+ private async resolveLegacyUrls(): Promise<string[]> {
145
+ const { llamaServerUrl = null } = await this.getMergedSettings();
95
146
  if (!llamaServerUrl) return [];
96
147
 
97
148
  return this.parseUrls(llamaServerUrl);
@@ -99,42 +150,76 @@ export class LlamaSettingsManager {
99
150
 
100
151
  /**
101
152
  * Parses a raw URL string into an array of cleaned URLs.
102
- * Splits on semicolons, trims whitespace, filters empty strings,
103
- * and strips trailing slashes.
153
+ * Splits on semicolons, trims whitespace, filters empty strings, strips
154
+ * trailing slashes, and drops entries without an http(s) scheme —
155
+ * collecting a warning for each dropped entry (same validation the
156
+ * `/models servers` editor applies).
104
157
  *
105
158
  * @returns A sanitized URL
106
159
  */
107
160
  private parseUrls(raw: string): string[] {
108
161
  return raw
109
162
  .split(";")
110
- .map((u) => u.trim())
111
- .filter((u) => u.length > 0)
112
- .map((u) => u.replace(/\/+$/, ""));
163
+ .map(normalizeUrl)
164
+ .filter((u) => {
165
+ if (u.length === 0) return false;
166
+ if (!isValidServerUrl(u)) {
167
+ this.warnings.push(
168
+ `Ignoring invalid server URL '${u}' (needs http(s)://)`,
169
+ );
170
+ return false;
171
+ }
172
+ return true;
173
+ });
174
+ }
175
+
176
+ /**
177
+ * Resolves the override map for a given server URL.
178
+ *
179
+ * Reads the `overrides` field from the matching server config and returns
180
+ * a map of model ID → override. Returns an empty object when the server
181
+ * has no `overrides` defined.
182
+ *
183
+ * @param serverUrl - The URL of the server to resolve overrides for
184
+ * @returns A map of model ID to override configuration (partial fields,
185
+ * fallbacks applied at consumption time)
186
+ */
187
+ async resolveServerOverrides(
188
+ serverUrl: string,
189
+ ): Promise<Record<string, ModelOverride>> {
190
+ const serverConfig = (await this.getLlamaSettings()).servers?.find(
191
+ (s: { url: string }) => s.url === serverUrl,
192
+ );
193
+ return serverConfig?.overrides ?? {};
113
194
  }
114
195
 
115
196
  /**
116
197
  * Resolves the servers that this extension will use.
117
198
  * Uses `resolveUrls()` as the source of truth for URLs (env > settings >
118
- * legacy > default), then applies `id`/`name` from `llamaSettings.servers`
119
- * as overrides when available.
199
+ * legacy > default), then applies `id`/`name`/`overrides` from
200
+ * `llamaSettings.servers` as overrides when available.
201
+ * Reloads settings from disk before reading.
120
202
  *
121
203
  * @returns A list of Server objects
122
204
  */
123
- resolveServers(): Server[] {
124
- const { pollingTimeout, serverTimeout } = this.resolveTimeouts();
125
- const urls = this.resolveUrls();
126
- const serverConfigs = this.llamaSettings.servers ?? [];
205
+ async resolveServers(): Promise<Server[]> {
206
+ const urls = await this.resolveUrls();
207
+ const serverConfigs = (await this.getLlamaSettings()).servers ?? [];
127
208
 
128
- return urls.map((url) => {
209
+ const servers: Server[] = [];
210
+ for (const url of urls) {
129
211
  const config = serverConfigs.find((s) => s.url === url);
130
- return new Server(
131
- url,
132
- config?.id,
133
- config?.name,
134
- serverTimeout,
135
- pollingTimeout,
212
+ const overrides = await this.resolveServerOverrides(url);
213
+ servers.push(
214
+ new Server(this, {
215
+ baseUrl: url,
216
+ customId: config?.id,
217
+ customName: config?.name,
218
+ overrides,
219
+ }),
136
220
  );
137
- });
221
+ }
222
+ return servers;
138
223
  }
139
224
 
140
225
  /**
@@ -174,8 +259,11 @@ export class LlamaSettingsManager {
174
259
  *
175
260
  * @returns `true` if the extension should load the model on model_select
176
261
  */
177
- resolveReactToModelSelect(): boolean {
178
- return this.llamaSettings.reactToModelSelect ?? REACT_TO_MODEL_SELECT;
262
+ async resolveReactToModelSelect(): Promise<boolean> {
263
+ return (
264
+ (await this.getLlamaSettings()).reactToModelSelect ??
265
+ REACT_TO_MODEL_SELECT
266
+ );
179
267
  }
180
268
 
181
269
  /**
@@ -183,8 +271,10 @@ export class LlamaSettingsManager {
183
271
  *
184
272
  * @returns `true` if the extension should auto-load models
185
273
  */
186
- resolveAutoloadOnMessage(): boolean {
187
- return this.llamaSettings.autoloadOnMessage ?? AUTOLOAD_ON_MESSAGE;
274
+ async resolveAutoloadOnMessage(): Promise<boolean> {
275
+ return (
276
+ (await this.getLlamaSettings()).autoloadOnMessage ?? AUTOLOAD_ON_MESSAGE
277
+ );
188
278
  }
189
279
 
190
280
  /**
@@ -192,12 +282,59 @@ export class LlamaSettingsManager {
192
282
  *
193
283
  * @returns Object with polling and server timeout values
194
284
  */
195
- resolveTimeouts(): { pollingTimeout: number; serverTimeout: number } {
285
+ async resolveTimeouts(): Promise<{
286
+ pollingTimeout: number;
287
+ serverTimeout: number;
288
+ }> {
289
+ const llamaSettings = await this.getLlamaSettings();
196
290
  return {
197
- pollingTimeout: this.llamaSettings.pollingTimeout ?? POLLING_TIMEOUT,
198
- serverTimeout: this.llamaSettings.serverTimeout ?? SERVER_TIMEOUT,
291
+ pollingTimeout: llamaSettings.pollingTimeout ?? POLLING_TIMEOUT,
292
+ serverTimeout: llamaSettings.serverTimeout ?? SERVER_TIMEOUT,
199
293
  };
200
294
  }
295
+
296
+ /**
297
+ * Resolves the sort order for model lists.
298
+ *
299
+ * @returns The sort order: "asc", "desc", "asc-name", "desc-name", or "api"
300
+ */
301
+ async resolveSortBy(): Promise<SortBy> {
302
+ return (await this.getLlamaSettings()).sortBy ?? SORT_BY;
303
+ }
304
+
305
+ /**
306
+ * Persists one llamaSettings field to settings and reloads the in-memory
307
+ * settings so resolvers see the change immediately.
308
+ *
309
+ * When `scope` is `"auto"` (default), writes to the project `.pi/settings.json`
310
+ * if it exists, otherwise to global `~/.pi/agent/settings.json`.
311
+ *
312
+ * Rejects if the file can't be read (e.g. invalid JSON) or written —
313
+ * in-memory state stays consistent (reload only on success).
314
+ */
315
+ async setLlamaSetting<K extends keyof LlamaSettings>(
316
+ key: K,
317
+ value: LlamaSettings[K],
318
+ scope: "auto" | "global" | "project" = "auto",
319
+ ): Promise<void> {
320
+ const store =
321
+ scope === "auto"
322
+ ? (await this.hasProjectSettings())
323
+ ? this.projectStore
324
+ : this.globalStore
325
+ : scope === "project"
326
+ ? this.projectStore
327
+ : this.globalStore;
328
+
329
+ await store.updateKey(SETTINGS_KEY, (current) => {
330
+ const merged =
331
+ typeof current === "object" && current !== null
332
+ ? (current as Record<string, unknown>)
333
+ : {};
334
+ return { ...merged, [key]: value };
335
+ });
336
+ await this.settingsManager.reload();
337
+ }
201
338
  }
202
339
 
203
340
  /**
@@ -1,3 +1,4 @@
1
+ import type { ModelCost } from "@earendil-works/pi-ai";
1
2
  import type { ProviderModelConfig } from "@earendil-works/pi-coding-agent";
2
3
  import { FALLBACK_CTX, POLLING_INTERVAL } from "../constants";
3
4
  import { Mode } from "../enums/mode";
@@ -16,14 +17,6 @@ export abstract class BaseModel {
16
17
  protected readonly server: Server,
17
18
  ) {}
18
19
 
19
- protected readonly statusMapper: Record<string, Status> = {
20
- loaded: Status.LOADED,
21
- loading: Status.LOADING,
22
- failed: Status.FAILED,
23
- sleeping: Status.SLEEPING,
24
- unloaded: Status.UNLOADED,
25
- };
26
-
27
20
  protected readonly labelIcons: Record<Status, string> = {
28
21
  [Status.LOADED]: "🟢",
29
22
  [Status.LOADING]: "🟡",
@@ -65,18 +58,24 @@ export abstract class BaseModel {
65
58
 
66
59
  /**
67
60
  * Whether the model is a reasoning model.
68
- * Currently always returns true since there's no way to detect this from llama-server.
61
+ * An override's `reasoning` wins; otherwise defaults to `true`, since
62
+ * there's no way to detect this from llama-server.
69
63
  */
70
64
  get reasoning(): boolean {
71
- return true;
65
+ return this.server.findOverrideForModel(this.id)?.reasoning ?? true;
72
66
  }
73
67
 
74
68
  /**
75
- * Detects the capabilities of the model
69
+ * Detects the capabilities of the model.
70
+ * An override's `capabilities` fully replaces detection; otherwise the
71
+ * model's modalities are probed from the server.
76
72
  *
77
73
  * @returns An array of capabilities, as expected by Pi
78
74
  */
79
75
  async getCapabilities(): Promise<("text" | "image")[]> {
76
+ const overridden = this.server.findOverrideForModel(this.id)?.capabilities;
77
+ if (overridden) return overridden;
78
+
80
79
  try {
81
80
  // When loaded, this works alright
82
81
  const { modalities } = await this.server.fetchModelProps(this.id);
@@ -123,9 +122,16 @@ export abstract class BaseModel {
123
122
  /**
124
123
  * Gets the context size of a particular model.
125
124
  *
125
+ * An override's `contextSize` (when set and `> 0`) replaces detection;
126
+ * otherwise the value is autodetected from the server, falling back to
127
+ * {@link FALLBACK_CTX}. A stored `0` behaves as if the key were absent.
128
+ *
126
129
  * @returns The context size in tokens
127
130
  */
128
131
  async getContextSize(): Promise<number> {
132
+ const overridden = this.server.findOverrideForModel(this.id)?.contextSize;
133
+ if (overridden && overridden > 0) return overridden;
134
+
129
135
  try {
130
136
  const { data } = await this.server.fetchModels();
131
137
  const { n_ctx } = data.find((m) => m.id === this.id)?.meta!;
@@ -170,7 +176,18 @@ export abstract class BaseModel {
170
176
  * @returns A Pi configuration object
171
177
  */
172
178
  async toProviderConfig(): Promise<ProviderModelConfig> {
173
- const response = {
179
+ const override = this.server.findOverrideForModel(this.id) ?? {};
180
+
181
+ // Merge the matched override's cost with zero defaults
182
+ const userCost = override.cost ?? {};
183
+ const cost: ModelCost = {
184
+ input: userCost.input ?? 0,
185
+ output: userCost.output ?? 0,
186
+ cacheRead: userCost.cacheRead ?? 0,
187
+ cacheWrite: userCost.cacheWrite ?? 0,
188
+ };
189
+
190
+ const response: ProviderModelConfig = {
174
191
  id: this.id,
175
192
  name: this.name,
176
193
  reasoning: this.reasoning,
@@ -184,10 +201,15 @@ export abstract class BaseModel {
184
201
  },
185
202
  input: await this.getCapabilities(),
186
203
  contextWindow: await this.getContextSize(),
187
- cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
188
- maxTokens: await this.getContextSize(),
204
+ cost,
205
+ maxTokens: override.maxTokens ?? (await this.getContextSize()),
189
206
  };
190
207
 
208
+ // Add compat if the override specifies it
209
+ if (override.compat) {
210
+ response.compat = override.compat;
211
+ }
212
+
191
213
  return response;
192
214
  }
193
215
 
@@ -253,7 +275,7 @@ export abstract class BaseModel {
253
275
  interval: number = POLLING_INTERVAL,
254
276
  ): Promise<void> {
255
277
  if (timeout === undefined) {
256
- timeout = this.server.pollingTimeout;
278
+ timeout = await this.server.getPollingTimeout();
257
279
  }
258
280
  while ((await this.getStatus()) === Status.LOADING) {
259
281
  // Force a timeout if we wasted too much time polling
@@ -41,7 +41,8 @@ export class RouterModel extends BaseModel {
41
41
  }
42
42
  }
43
43
 
44
- const timeout = this.server.pollingTimeout - elapsed;
44
+ const pollingTimeout = await this.server.getPollingTimeout();
45
+ const timeout = pollingTimeout - elapsed;
45
46
  return await super.pollStatus(startTime, timeout);
46
47
  }
47
48