pi-llama-cpp 0.7.2 → 0.8.1

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.
package/README.md CHANGED
@@ -13,6 +13,7 @@ A [Pi Coding Agent](https://pi.dev/) extension that integrates with running [lla
13
13
  - **Auth support** — allows to login into a llama.cpp server that was secured with an API key
14
14
  - **Multiple server support** — connect to multiple llama.cpp servers simultaneously by separating URLs with semicolons
15
15
  - **Thinking budget support** — configurable token budgets for model reasoning/thinking, mapped to Pi's thinking levels
16
+ - **Real-time progress tracking** — live loading progress via SSE (falls back to polling)
16
17
 
17
18
  ### Status Indicators
18
19
 
@@ -200,9 +201,11 @@ This keeps the server in sync with the active model in Pi, regardless of how the
200
201
 
201
202
  ### Loading Models
202
203
 
203
- When you trigger a load, switch, or retry action, the extension polls the server to track progress. If a model takes longer than **60 seconds** to load, the polling times out with an error.
204
+ When you trigger a load, switch, or retry action, the extension uses SSE (Server-Sent Events) to receive real-time progress updates from the server. If SSE is not available, it falls back to polling.
204
205
 
205
- > **Note:** The timeout is only for the polling. The model might still be loading.
206
+ If loading takes longer than **60 seconds**, the operation times out with an error.
207
+
208
+ > **Note:** The timeout only applies to the progress detection. The model might still be loading in the background.
206
209
 
207
210
  ### Model Configuration
208
211
 
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "pi-llama-cpp",
3
- "version": "0.7.2",
3
+ "version": "0.8.1",
4
4
  "description": "Pi extension for llama.cpp integration. Supports router, single and legacy models. Supports multiple servers.",
5
5
  "keywords": [
6
6
  "pi",
@@ -36,7 +36,7 @@
36
36
  "@earendil-works/pi-tui": "*"
37
37
  },
38
38
  "devDependencies": {
39
- "@types/node": "^26.0.0",
39
+ "@types/node": "^26.0.1",
40
40
  "prettier-plugin-organize-imports": "^4.3.0",
41
41
  "vitest": "^4.1.9"
42
42
  }
@@ -0,0 +1,121 @@
1
+ import { POLLING_INTERVAL } from "../constants";
2
+ import { Cache } from "../utils/cache";
3
+ import { Mutex } from "../utils/mutex";
4
+
5
+ /**
6
+ * HTTP client for llama-server with caching and deduplication.
7
+ */
8
+ export class ApiClient {
9
+ private cache = new Cache(POLLING_INTERVAL / 2);
10
+ private mutex = new Mutex();
11
+
12
+ /**
13
+ * Creates a new ApiClient.
14
+ *
15
+ * @param baseUrl The base URL of the llama-server
16
+ * @param apiKey The API key for authentication
17
+ */
18
+ constructor(
19
+ private readonly baseUrl: string,
20
+ private readonly apiKey: string,
21
+ ) {}
22
+
23
+ /**
24
+ * Makes a cached, deduplicated GET request to the llama-server.
25
+ * Results are cached for half the polling interval and in-flight requests are deduplicated.
26
+ *
27
+ * @param endpoint The endpoint path to fetch (e.g. "/health")
28
+ * @returns The parsed JSON response from the server
29
+ */
30
+ async get<T>(endpoint: string): Promise<T> {
31
+ const cached = this.cache.get<T>(endpoint);
32
+ if (cached !== undefined) return cached;
33
+
34
+ return this.mutex.getOrCreate(endpoint, async () => {
35
+ const data = (await this.do_get<T>(endpoint)) as T;
36
+ this.cache.set(endpoint, data);
37
+ return data;
38
+ });
39
+ }
40
+
41
+ /**
42
+ * Makes a cached, deduplicated POST request to the llama-server.
43
+ * Results are cached for half the polling interval and in-flight requests are deduplicated.
44
+ *
45
+ * @param endpoint The endpoint path to post to
46
+ * @param body The optional request body
47
+ * @returns The parsed JSON response from the server
48
+ */
49
+ async post<T>(endpoint: string, body?: Record<string, unknown>): Promise<T> {
50
+ const key = this.cacheKey(endpoint, body);
51
+ const cached = this.cache.get<T>(key);
52
+ if (cached !== undefined) return cached;
53
+
54
+ return this.mutex.getOrCreate(key, async () => {
55
+ const data = (await this.do_post<T>(endpoint, body)) as T;
56
+ this.cache.set(key, data);
57
+ return data;
58
+ });
59
+ }
60
+
61
+ /**
62
+ * Clears the entire cache.
63
+ */
64
+ clearCache(): void {
65
+ this.cache.clear();
66
+ }
67
+
68
+ /**
69
+ * Makes a raw GET request to the llama-server.
70
+ * This bypasses caching and deduplication.
71
+ *
72
+ * @param endpoint The endpoint path to fetch (e.g. "/health")
73
+ * @returns The parsed JSON response from the server
74
+ */
75
+ private async do_get<T>(endpoint: string): Promise<T> {
76
+ const url = `${this.baseUrl}${endpoint}`;
77
+
78
+ const res = await fetch(url, {
79
+ headers: { Authorization: `Bearer ${this.apiKey}` },
80
+ });
81
+
82
+ return res.json();
83
+ }
84
+
85
+ /**
86
+ * Makes a raw POST request to the llama-server.
87
+ * This bypasses caching and deduplication.
88
+ *
89
+ * @param endpoint The endpoint path to post to
90
+ * @param body The optional request body
91
+ * @returns The parsed JSON response from the server
92
+ */
93
+ private async do_post<T>(
94
+ endpoint: string,
95
+ body?: Record<string, unknown>,
96
+ ): Promise<T> {
97
+ const url = `${this.baseUrl}${endpoint}`;
98
+
99
+ const res = await fetch(url, {
100
+ method: "POST",
101
+ headers: {
102
+ "Content-Type": "application/json",
103
+ Authorization: `Bearer ${this.apiKey}`,
104
+ },
105
+ body: body ? JSON.stringify(body) : undefined,
106
+ });
107
+
108
+ return res.json();
109
+ }
110
+
111
+ /**
112
+ * Sets a cache key
113
+ *
114
+ * @param endpoint The endpoint path to post to
115
+ * @param body The optional request body
116
+ * @returns The cache key
117
+ */
118
+ private cacheKey(endpoint: string, body?: Record<string, unknown>): string {
119
+ return body ? `${endpoint}:${JSON.stringify(body)}` : endpoint;
120
+ }
121
+ }
package/src/constants.ts CHANGED
@@ -44,7 +44,7 @@ export const POLLING_TIMEOUT = 60000;
44
44
  export const READABLE_TIMEOUT = 15000;
45
45
 
46
46
  /**
47
- * Timeout (ms) for server verification before assuming failure
47
+ * Timeout (ms) for server verification and SSE support probe
48
48
  */
49
49
  export const SERVER_TIMEOUT = 1000;
50
50
 
@@ -40,6 +40,9 @@ export interface DataProperty {
40
40
  created: number;
41
41
  status?: StatusProperty;
42
42
  architecture?: ArchitectureProperty;
43
+ source?: string;
44
+ can_remove?: boolean;
45
+ need_download?: boolean;
43
46
  meta?: MetaProperty;
44
47
  }
45
48
 
@@ -2,28 +2,56 @@
2
2
  * The structure of llama-server's /props endpoint
3
3
  */
4
4
  export interface PropsEndpoint {
5
- role?: "router";
6
- error?: PropsError;
5
+ role: "router";
6
+ max_instances: number;
7
+ models_autoload: boolean;
8
+ model_alias: string;
9
+ model_path: string;
7
10
  default_generation_settings: Record<string, any>;
11
+ ui_settings: Record<string, any>;
12
+ build_info: string;
13
+ cors_proxy_enabled: boolean;
14
+ }
15
+
16
+ /**
17
+ * The structure of llama-server's /props?model=<id> endpoint
18
+ */
19
+ export interface PropsModelEndpoint {
20
+ error?: PropsError;
21
+ default_generation_settings: {
22
+ params: Record<string, any>;
23
+ n_ctx: number;
24
+ };
8
25
  total_slots: number;
9
26
  model_alias: string;
10
27
  model_path: string;
11
28
  modalities: {
12
29
  vision: boolean;
30
+ video: boolean;
13
31
  audio: boolean;
14
32
  };
15
33
  media_marker: string;
16
34
  endpoint_slots: boolean;
17
35
  endpoint_props: boolean;
18
36
  endpoint_metrics: boolean;
19
- webui: boolean;
20
- webui_settings: Record<string, any>;
37
+ ui: boolean;
38
+ ui_settings: Record<string, any>;
21
39
  chat_template: string;
22
- chat_template_caps: Record<string, boolean>;
40
+ chat_template_caps: {
41
+ supports_object_arguments: boolean;
42
+ supports_parallel_tool_calls: boolean;
43
+ supports_preserve_reasoning: boolean;
44
+ supports_string_content: boolean;
45
+ supports_system_role: boolean;
46
+ supports_tool_calls: boolean;
47
+ supports_tools: boolean;
48
+ supports_typed_content: boolean;
49
+ };
23
50
  bos_token: string;
24
51
  eos_token: string;
25
52
  build_info: string;
26
53
  is_sleeping: boolean;
54
+ cors_proxy_enabled: boolean;
27
55
  }
28
56
 
29
57
  export interface PropsError {
@@ -134,6 +134,17 @@ export class CommandManager {
134
134
  ctx.ui.notify(`Loading ${model.name}...`, "info");
135
135
  EventManager.inflightModel = model;
136
136
 
137
+ // Subscribe to progress events
138
+ const cleanupProgress = this.serverManager
139
+ .getServer(model)
140
+ .sseManager.subscribeToProgress(model.id, (percentage, stage) => {
141
+ const stageText = stage ? ` (${stage})` : "";
142
+ ctx.ui.notify(
143
+ `Loading ${model.name}... [${percentage}%${stageText}]`,
144
+ "info",
145
+ );
146
+ });
147
+
137
148
  const onSuccess = async () => {
138
149
  const { serverId } = model;
139
150
  const piModel = ctx.modelRegistry.find(serverId, model.id);
@@ -171,7 +182,10 @@ export class CommandManager {
171
182
  .load()
172
183
  .then(onSuccess)
173
184
  .catch(onFailure)
174
- .finally(EventManager.resetInflightModel);
185
+ .finally(() => {
186
+ cleanupProgress();
187
+ EventManager.resetInflightModel();
188
+ });
175
189
  }
176
190
  }
177
191
 
@@ -117,6 +117,16 @@ export class ServerManager {
117
117
  return warnings;
118
118
  }
119
119
 
120
+ /**
121
+ * Returns the server for a given model.
122
+ *
123
+ * @param model - The model to find the server for
124
+ * @returns The server containing the model
125
+ */
126
+ getServer(model: BaseModel): Server {
127
+ return this.servers.find((s) => s.baseUrl === model.serverUrl)!;
128
+ }
129
+
120
130
  /**
121
131
  * Returns all models from all servers.
122
132
  *
@@ -87,9 +87,11 @@ export abstract class BaseModel {
87
87
  const model = data.find((d) => d.id === this.id);
88
88
  if (!model) return ["text"];
89
89
 
90
- const { input_modalities } = model.architecture!;
90
+ const input_modalities: ("text" | "image" | "audio")[] = model
91
+ .architecture?.input_modalities ?? ["text"];
92
+
91
93
  const response = input_modalities.filter(
92
- (mod) => mod === "text" || mod === "image",
94
+ (mod): mod is "text" | "image" => mod === "text" || mod === "image",
93
95
  );
94
96
 
95
97
  return response;
@@ -189,14 +191,25 @@ export abstract class BaseModel {
189
191
  }
190
192
 
191
193
  /**
192
- * Loads the model in llama-server
194
+ * Loads the model in llama-server.
195
+ * Uses SSE status events when available, falling back to polling.
193
196
  */
194
197
  async load(): Promise<void> {
195
198
  const status = await this.getStatus();
196
199
  if (status === Status.LOADED || status === Status.SLEEPING) return;
197
200
 
198
201
  await this.server.postRequest("load", this.id);
199
- await this.pollStatus();
202
+
203
+ if (await this.server.sseManager.probeSSE()) {
204
+ const { status, exit_code } =
205
+ await this.server.sseManager.subscribeToStatus(this.id);
206
+
207
+ if (status === "failed" || (status === "unloaded" && exit_code !== 0)) {
208
+ throw new Error(`Model loading failed: ${this.id}`);
209
+ }
210
+ } else {
211
+ await this.pollStatus();
212
+ }
200
213
  }
201
214
 
202
215
  /**
package/src/server.ts CHANGED
@@ -1,25 +1,35 @@
1
- import { POLLING_INTERVAL, PROVIDER_NAME, PROVIDER_PREFIX } from "./constants";
1
+ import { ApiClient } from "./api/client";
2
+ import { PROVIDER_NAME, PROVIDER_PREFIX } from "./constants";
2
3
  import { Mode } from "./enums/mode";
3
4
  import { ServerStatus } from "./enums/serverStatus";
4
5
  import { HealthEndpoint } from "./interfaces/endpoints/health";
5
6
  import { ModelsEndpoint } from "./interfaces/endpoints/models";
6
- import { PropsEndpoint } from "./interfaces/endpoints/props";
7
+ import {
8
+ PropsEndpoint,
9
+ PropsModelEndpoint,
10
+ } from "./interfaces/endpoints/props";
7
11
  import { BaseModel } from "./models/baseModel";
8
12
  import { LegacyModel } from "./models/legacyModel";
9
13
  import { RouterModel } from "./models/routerModel";
10
14
  import { SingleModel } from "./models/singleModel";
11
15
  import { ConfigResolver } from "./resolver";
12
- import { Cache } from "./utils/cache";
13
- import { Mutex } from "./utils/mutex";
16
+ import { SSEManager } from "./sse/manager";
14
17
 
15
18
  export class Server {
16
19
  public readonly models: BaseModel[] = [];
17
20
  private configResolver = new ConfigResolver();
18
- private cache = new Cache(POLLING_INTERVAL / 2);
19
- private mutex = new Mutex();
21
+ private apiClient!: ApiClient;
22
+ private sse!: SSEManager;
20
23
 
21
24
  constructor(readonly baseUrl: string) {}
22
25
 
26
+ /**
27
+ * Provides access to the SSE manager for direct subscriptions.
28
+ */
29
+ get sseManager(): SSEManager {
30
+ return this.sse;
31
+ }
32
+
23
33
  /**
24
34
  * Generates a unique provider ID from a server URL.
25
35
  */
@@ -47,7 +57,9 @@ export class Server {
47
57
  * Clears the cache first so we always fetch fresh data.
48
58
  */
49
59
  async initialize() {
50
- this.cache.clear();
60
+ const apiKey = await this.getApiKey();
61
+ this.apiClient = new ApiClient(this.baseUrl, apiKey);
62
+ this.sse = new SSEManager(this.baseUrl, apiKey);
51
63
  const { data } = await this.fetchModels();
52
64
  const mode = await this.detectServerMode();
53
65
 
@@ -87,6 +99,8 @@ export class Server {
87
99
  * @returns The server status
88
100
  */
89
101
  async isReady(timeout: number): Promise<ServerStatus> {
102
+ this.apiClient ??= new ApiClient(this.baseUrl, await this.getApiKey());
103
+
90
104
  try {
91
105
  const timeoutPromise = new Promise<never>((_, reject) =>
92
106
  setTimeout(() => reject(new Error("timeout")), timeout),
@@ -113,7 +127,7 @@ export class Server {
113
127
  * @returns The health status
114
128
  */
115
129
  async fetchServerHealth(): Promise<HealthEndpoint> {
116
- return await this.rpc<HealthEndpoint>("/health");
130
+ return await this.apiClient.get<HealthEndpoint>("/health");
117
131
  }
118
132
 
119
133
  /**
@@ -122,7 +136,7 @@ export class Server {
122
136
  * @return The models from the server
123
137
  */
124
138
  async fetchModels(): Promise<ModelsEndpoint> {
125
- return await this.rpc<ModelsEndpoint>("/v1/models");
139
+ return await this.apiClient.get<ModelsEndpoint>("/v1/models");
126
140
  }
127
141
 
128
142
  /**
@@ -131,7 +145,7 @@ export class Server {
131
145
  * @return The properties of the server
132
146
  */
133
147
  async fetchServerProps(): Promise<PropsEndpoint> {
134
- return await this.rpc<PropsEndpoint>("/props?autoload=false");
148
+ return await this.apiClient.get<PropsEndpoint>("/props?autoload=false");
135
149
  }
136
150
 
137
151
  /**
@@ -140,8 +154,8 @@ export class Server {
140
154
  * @param modelId The ID of the model
141
155
  * @return The properties of the specified model
142
156
  */
143
- async fetchModelProps(modelId: string): Promise<PropsEndpoint> {
144
- return await this.rpc<PropsEndpoint>(
157
+ async fetchModelProps(modelId: string): Promise<PropsModelEndpoint> {
158
+ return await this.apiClient.get<PropsModelEndpoint>(
145
159
  `/props?model=${modelId}&autoload=false`,
146
160
  );
147
161
  }
@@ -156,78 +170,9 @@ export class Server {
156
170
  resource: "load" | "unload",
157
171
  model: string,
158
172
  ): Promise<ModelsEndpoint> {
159
- this.cache.clear();
160
- return await this.rpc<ModelsEndpoint>(`/models/${resource}`, { model });
161
- }
162
-
163
- /**
164
- * Makes a cached, deduplicated request to the llama-server.
165
- * Results are cached for half the {@link POLLING_INTERVAL} and in-flight requests are deduplicated.
166
- *
167
- * @param endpoint The endpoint path to fetch (e.g. "/health")
168
- * @param body The optional request body for POST requests
169
- * @returns The parsed JSON response from the server
170
- */
171
- private async rpc<T>(
172
- endpoint: string,
173
- body?: Record<string, unknown>,
174
- ): Promise<T> {
175
- const key = this.cacheKey(endpoint, body);
176
-
177
- // Check cache
178
- const cached = this.cache.get<T>(key);
179
- if (cached !== undefined) {
180
- return cached;
181
- }
182
-
183
- // Deduplicate in-flight requests
184
- return this.mutex.getOrCreate(key, async () => {
185
- const data = await this.fetch<T>(endpoint, body);
186
- this.cache.set(key, data);
187
- return data;
173
+ this.apiClient.clearCache();
174
+ return await this.apiClient.post<ModelsEndpoint>(`/models/${resource}`, {
175
+ model,
188
176
  });
189
177
  }
190
-
191
- /**
192
- * Makes an HTTP request to the llama-server and returns the parsed JSON response.
193
- *
194
- * @param endpoint The endpoint path to fetch (e.g. "/health")
195
- * @param body The optional request body for POST requests
196
- * @returns The parsed JSON response from the server
197
- */
198
- private async fetch<T>(
199
- endpoint: string,
200
- body?: Record<string, unknown>,
201
- ): Promise<T> {
202
- const url = `${this.baseUrl}${endpoint}`;
203
- const apiKey = await this.getApiKey();
204
-
205
- const data = {
206
- method: body ? "POST" : "GET",
207
- headers: body ? { "Content-Type": "application/json" } : undefined,
208
- body: body ? JSON.stringify(body) : undefined,
209
- };
210
-
211
- const res = await fetch(url, {
212
- ...data,
213
- headers: {
214
- ...data.headers,
215
- ...(apiKey ? { Authorization: `Bearer ${apiKey}` } : {}),
216
- },
217
- });
218
-
219
- const response: T = await res.json();
220
- return response;
221
- }
222
-
223
- /**
224
- * Generates a cache key from the endpoint and body.
225
- *
226
- * @param endpoint The endpoint path (e.g. "/health")
227
- * @param body The optional request body for POST requests
228
- * @returns A key used for caching
229
- */
230
- private cacheKey(endpoint: string, body?: Record<string, unknown>): string {
231
- return body ? `${endpoint}:${JSON.stringify(body)}` : endpoint;
232
- }
233
178
  }
@@ -0,0 +1,158 @@
1
+ import { POLLING_INTERVAL } from "../constants";
2
+ import type { SSECallback, SSECleanup, SSEEvent } from "./types";
3
+
4
+ /**
5
+ * SSE client for llama-server's /models/sse endpoint.
6
+ *
7
+ * Uses a single shared EventSource per server instance.
8
+ * Supports multiple model subscriptions with automatic event routing.
9
+ * Handles reconnection by re-subscribing all callbacks.
10
+ */
11
+ export class SSEClient {
12
+ private eventSource: EventSource | null = null;
13
+ private subscribers: Map<string, SSECallback> = new Map();
14
+ private connected: boolean = false;
15
+ private reconnecting: boolean = false; // tracks if EventSource auto-reconnect is in progress
16
+ private _onConnectFailed: (() => void) | null = null;
17
+ private _hasReceivedEvents: boolean = false;
18
+
19
+ /**
20
+ * @param sseEndpoint - The full SSE endpoint URL (e.g., "http://127.0.0.1:8080/models/sse")
21
+ * @param apiKey - Optional API key for authenticated servers
22
+ */
23
+ constructor(
24
+ private readonly sseEndpoint: string,
25
+ private readonly apiKey?: string,
26
+ ) {}
27
+
28
+ /**
29
+ * Connects to the SSE endpoint.
30
+ *
31
+ * @returns true if the connection was established successfully
32
+ */
33
+ async connect(): Promise<boolean> {
34
+ if (this.connected) return true;
35
+
36
+ const url = this.buildUrl();
37
+
38
+ try {
39
+ this.eventSource = new EventSource(url);
40
+ } catch {
41
+ this.connected = false;
42
+ return false;
43
+ }
44
+
45
+ this.eventSource.onopen = () => {
46
+ this.connected = true;
47
+ this.reconnecting = false;
48
+ };
49
+
50
+ this.eventSource.onerror = () => {
51
+ // EventSource will auto-reconnect; we just track state
52
+ this.connected = false;
53
+ this.reconnecting = true;
54
+
55
+ // Notify subscriber if connection fails before any event is received
56
+ if (!this._hasReceivedEvents && this._onConnectFailed) {
57
+ this._onConnectFailed();
58
+ }
59
+ };
60
+
61
+ this.eventSource.onmessage = (event: MessageEvent) => {
62
+ try {
63
+ const data = JSON.parse(event.data);
64
+ const sseEvent: SSEEvent = {
65
+ event: data.event ?? "unknown",
66
+ model: data.model ?? "*",
67
+ data: data.data,
68
+ };
69
+ this._hasReceivedEvents = true;
70
+ this.dispatch(sseEvent);
71
+ } catch {
72
+ // Invalid JSON, ignore
73
+ }
74
+ };
75
+
76
+ // Wait a bit for the connection to establish
77
+ await new Promise<void>((resolve) => {
78
+ const timeout = setTimeout(() => resolve(), POLLING_INTERVAL);
79
+ this.eventSource!.onopen = () => {
80
+ clearTimeout(timeout);
81
+ this.connected = true;
82
+ this.reconnecting = false;
83
+ resolve();
84
+ };
85
+ });
86
+
87
+ return this.connected;
88
+ }
89
+
90
+ /**
91
+ * Sets a callback to be called when the connection fails before
92
+ * any event is received. Useful for rejecting promises early.
93
+ *
94
+ * @param callback - Called once when connection fails
95
+ */
96
+ setOnConnectFailed(callback: () => void): void {
97
+ this._onConnectFailed = callback;
98
+ }
99
+
100
+ /**
101
+ * Subscribes to SSE events for a specific model.
102
+ * Auto-connects if not already connected.
103
+ *
104
+ * @param modelId - The model ID to subscribe to
105
+ * @param callback - Callback to receive SSE events
106
+ * @returns A cleanup function to unsubscribe
107
+ */
108
+ subscribe(modelId: string, callback: SSECallback): SSECleanup {
109
+ this.subscribers.set(modelId, callback);
110
+
111
+ if (!this.connected && !this.reconnecting) {
112
+ this.connect();
113
+ }
114
+
115
+ return () => {
116
+ this.subscribers.delete(modelId);
117
+ };
118
+ }
119
+
120
+ /**
121
+ * Disconnects from the SSE endpoint and clears all subscriptions.
122
+ */
123
+ disconnect(): void {
124
+ if (this.eventSource) {
125
+ this.eventSource.close();
126
+ this.eventSource = null;
127
+ }
128
+ this.connected = false;
129
+ this.subscribers.clear();
130
+ }
131
+
132
+ /**
133
+ * Builds the full URL with optional API key query param.
134
+ */
135
+ private buildUrl(): string {
136
+ if (this.apiKey) {
137
+ return `${this.sseEndpoint}?api_key=${encodeURIComponent(this.apiKey)}`;
138
+ }
139
+ return this.sseEndpoint;
140
+ }
141
+
142
+ /**
143
+ * Dispatches an SSE event to all matching subscribers.
144
+ */
145
+ private dispatch(event: SSEEvent): void {
146
+ // Dispatch to model-specific subscriber
147
+ const modelCallback = this.subscribers.get(event.model);
148
+ if (modelCallback) {
149
+ modelCallback(event);
150
+ }
151
+
152
+ // Also dispatch to wildcard subscriber if present
153
+ const wildcardCallback = this.subscribers.get("*");
154
+ if (wildcardCallback && event.model !== "*") {
155
+ wildcardCallback(event);
156
+ }
157
+ }
158
+ }
@@ -0,0 +1,208 @@
1
+ import { POLLING_TIMEOUT, SERVER_TIMEOUT } from "../constants";
2
+ import { SSEClient } from "./client";
3
+ import {
4
+ DownloadProgressData,
5
+ ProgressData,
6
+ SSECallback,
7
+ SSECleanup,
8
+ SSEEvent,
9
+ SSEEventType,
10
+ StatusChangeData,
11
+ } from "./types";
12
+
13
+ /**
14
+ * Manages SSE connections and event routing for a single llama-server instance.
15
+ *
16
+ * Handles:
17
+ * - Shared EventSource connection
18
+ * - Model-based event subscription with callback aggregation
19
+ * - Progress parsing and callback dispatch
20
+ */
21
+ export class SSEManager {
22
+ private sseClient: SSEClient | null = null;
23
+ private sseSubscribers: Map<string, SSECleanup> = new Map();
24
+ private modelCallbacks: Map<string, SSECallback[]> = new Map();
25
+ private sseSupported: boolean | null = null;
26
+
27
+ constructor(
28
+ private readonly baseUrl: string,
29
+ private readonly apiKey: string,
30
+ ) {}
31
+
32
+ /**
33
+ * The SSE endpoint URL.
34
+ */
35
+ private get sseEndpoint(): string {
36
+ return `${this.baseUrl}/models/sse`;
37
+ }
38
+
39
+ /**
40
+ * Probes the SSE endpoint to check if it's supported.
41
+ * Result is cached for the lifetime of the manager.
42
+ *
43
+ * @returns true if SSE is supported
44
+ */
45
+ async probeSSE(): Promise<boolean> {
46
+ if (this.sseSupported !== null) return this.sseSupported;
47
+
48
+ try {
49
+ let url = this.sseEndpoint;
50
+ if (this.apiKey) {
51
+ url = `${url}?api_key=${encodeURIComponent(this.apiKey)}`;
52
+ }
53
+ const response = await fetch(url, {
54
+ method: "GET",
55
+ signal: AbortSignal.timeout(SERVER_TIMEOUT),
56
+ });
57
+ this.sseSupported =
58
+ response.ok &&
59
+ !!response.headers.get("content-type")?.includes("text/event-stream");
60
+ } catch {
61
+ this.sseSupported = false;
62
+ }
63
+
64
+ return this.sseSupported;
65
+ }
66
+
67
+ /**
68
+ * Subscribes to SSE events for a specific model.
69
+ * Uses a shared SSE connection per server.
70
+ * Aggregates multiple callbacks into one SSEClient subscription.
71
+ *
72
+ * @param modelId - The model ID to subscribe to
73
+ * @param callback - Callback to receive SSE events
74
+ * @returns A cleanup function to unsubscribe
75
+ */
76
+ private subscribeToSSE(
77
+ modelId: string,
78
+ callback: (event: SSEEvent) => void,
79
+ ): SSECleanup {
80
+ // Aggregate callbacks for this model
81
+ const callbacks = this.modelCallbacks.get(modelId) ?? [];
82
+ callbacks.push(callback);
83
+ this.modelCallbacks.set(modelId, callbacks);
84
+
85
+ // Create SSE client if not already created
86
+ this.sseClient ??= new SSEClient(this.sseEndpoint, this.apiKey);
87
+
88
+ // Subscribe a single dispatching callback to the SSE client
89
+ if (!this.sseSubscribers.has(modelId)) {
90
+ const dispatch = (event: SSEEvent) => {
91
+ for (const cb of callbacks) cb(event);
92
+ };
93
+ const cleanup = this.sseClient!.subscribe(modelId, dispatch);
94
+ this.sseSubscribers.set(modelId, cleanup);
95
+ }
96
+
97
+ return () => {
98
+ const list = this.modelCallbacks.get(modelId);
99
+ if (list) {
100
+ const idx = list.indexOf(callback);
101
+ if (idx !== -1) list.splice(idx, 1);
102
+ if (list.length === 0) {
103
+ this.modelCallbacks.delete(modelId);
104
+ const cleanup = this.sseSubscribers.get(modelId);
105
+ if (cleanup) cleanup();
106
+ this.sseSubscribers.delete(modelId);
107
+ }
108
+ }
109
+ };
110
+ }
111
+
112
+ /**
113
+ * Subscribes to SSE progress events for a specific model.
114
+ * Parses SSE events and calls the progress callback with percentage and stage.
115
+ *
116
+ * @param modelId - The model ID to subscribe to
117
+ * @param onProgress - Callback to receive progress updates (percentage 0-100, stage name)
118
+ * @returns A cleanup function to unsubscribe
119
+ */
120
+ subscribeToProgress(
121
+ modelId: string,
122
+ onProgress: (percentage: number, stage?: string) => void,
123
+ ): SSECleanup {
124
+ // Track download progress across multiple URLs
125
+ let totalDownloaded = 0;
126
+ let totalToDownload = 0;
127
+
128
+ return this.subscribeToSSE(modelId, (event: SSEEvent) => {
129
+ if (event.event === SSEEventType.status_change && event.data) {
130
+ const data = event.data as unknown as StatusChangeData;
131
+
132
+ if (data.status === "loading" && data.progress) {
133
+ const progress = data.progress as ProgressData;
134
+ const percentage = Math.round(progress.value * 100);
135
+ onProgress(percentage, progress.current);
136
+ } else if (data.status === "loaded" || data.status === "failed") {
137
+ // Reset download tracking on final state
138
+ totalDownloaded = 0;
139
+ totalToDownload = 0;
140
+ }
141
+ } else if (event.event === SSEEventType.download_progress && event.data) {
142
+ const downloadData = event.data as DownloadProgressData;
143
+ totalDownloaded = 0;
144
+ totalToDownload = 0;
145
+
146
+ for (const urlData of Object.values(downloadData)) {
147
+ totalDownloaded += urlData.done;
148
+ totalToDownload += urlData.total;
149
+ }
150
+
151
+ if (totalToDownload > 0) {
152
+ const percentage = Math.round(
153
+ (totalDownloaded / totalToDownload) * 100,
154
+ );
155
+ onProgress(percentage, "downloading");
156
+ }
157
+ }
158
+ });
159
+ }
160
+
161
+ /**
162
+ * Subscribes to SSE status change events for a specific model.
163
+ * Resolves with the final status string once the model reaches a terminal state.
164
+ * Rejects immediately if the connection fails before any event is received.
165
+ *
166
+ * @param modelId - The model ID to subscribe to
167
+ * @returns Promise that resolves with the final status string
168
+ */
169
+ subscribeToStatus(modelId: string): Promise<StatusChangeData> {
170
+ return new Promise((resolve, reject) => {
171
+ const timeout = setTimeout(
172
+ () => reject(new Error(`SSE status timeout for model: ${modelId}`)),
173
+ POLLING_TIMEOUT,
174
+ );
175
+
176
+ this.subscribeToSSE(modelId, (event: SSEEvent) => {
177
+ if (event.event === SSEEventType.status_change && event.data) {
178
+ const data = event.data as unknown as StatusChangeData;
179
+ if (["loaded", "unloaded", "failed"].includes(data.status)) {
180
+ clearTimeout(timeout);
181
+ resolve(data);
182
+ }
183
+ }
184
+ });
185
+
186
+ // Reject immediately if the connection fails before any event is received
187
+ this.sseClient?.setOnConnectFailed(() => {
188
+ clearTimeout(timeout);
189
+ reject(new Error(`SSE connection failed for model: ${modelId}`));
190
+ });
191
+ });
192
+ }
193
+
194
+ /**
195
+ * Disconnects the SSE client and cleans up all subscriptions.
196
+ */
197
+ disconnect(): void {
198
+ for (const cleanup of this.sseSubscribers.values()) {
199
+ cleanup();
200
+ }
201
+ this.sseSubscribers.clear();
202
+ this.modelCallbacks.clear();
203
+ if (this.sseClient) {
204
+ this.sseClient.disconnect();
205
+ this.sseClient = null;
206
+ }
207
+ }
208
+ }
@@ -0,0 +1,65 @@
1
+ /**
2
+ * SSE event types from llama-server's /models/sse endpoint
3
+ */
4
+
5
+ /**
6
+ * Possible event types from the SSE stream
7
+ */
8
+ export const SSEEventType = {
9
+ status_change: "status_change",
10
+ download_progress: "download_progress",
11
+ download_finished: "download_finished",
12
+ download_failed: "download_failed",
13
+ models_reload: "models_reload",
14
+ model_remove: "model_remove",
15
+ } as const;
16
+
17
+ export type SSEEventType = (typeof SSEEventType)[keyof typeof SSEEventType];
18
+
19
+ /**
20
+ * A parsed SSE event from the /models/sse endpoint
21
+ */
22
+ export interface SSEEvent {
23
+ event: SSEEventType;
24
+ model: string; // model ID or "*" for global events
25
+ data?: Record<string, unknown>;
26
+ }
27
+
28
+ /**
29
+ * Progress data sent during model loading
30
+ */
31
+ export interface ProgressData {
32
+ stages: string[];
33
+ current: string;
34
+ value: number; // 0.0 to 1.0
35
+ }
36
+
37
+ /**
38
+ * Data payload for status_change events
39
+ */
40
+ export interface StatusChangeData {
41
+ status: string;
42
+ exit_code?: number;
43
+ info?: Record<string, unknown>;
44
+ progress?: ProgressData;
45
+ }
46
+
47
+ /**
48
+ * Data payload for download_progress events
49
+ */
50
+ export interface DownloadProgressData {
51
+ [url: string]: {
52
+ done: number;
53
+ total: number;
54
+ };
55
+ }
56
+
57
+ /**
58
+ * Subscriber callback type for SSE events
59
+ */
60
+ export type SSECallback = (event: SSEEvent) => void;
61
+
62
+ /**
63
+ * Cleanup function to unsubscribe from SSE events
64
+ */
65
+ export type SSECleanup = () => void;