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 +5 -2
- package/package.json +2 -2
- package/src/api/client.ts +121 -0
- package/src/constants.ts +1 -1
- package/src/interfaces/endpoints/models.ts +3 -0
- package/src/interfaces/endpoints/props.ts +33 -5
- package/src/managers/command.ts +15 -1
- package/src/managers/server.ts +10 -0
- package/src/models/baseModel.ts +17 -4
- package/src/server.ts +29 -84
- package/src/sse/client.ts +158 -0
- package/src/sse/manager.ts +208 -0
- package/src/sse/types.ts +65 -0
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
|
|
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
|
-
|
|
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.
|
|
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.
|
|
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
|
|
47
|
+
* Timeout (ms) for server verification and SSE support probe
|
|
48
48
|
*/
|
|
49
49
|
export const SERVER_TIMEOUT = 1000;
|
|
50
50
|
|
|
@@ -2,28 +2,56 @@
|
|
|
2
2
|
* The structure of llama-server's /props endpoint
|
|
3
3
|
*/
|
|
4
4
|
export interface PropsEndpoint {
|
|
5
|
-
role
|
|
6
|
-
|
|
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
|
-
|
|
20
|
-
|
|
37
|
+
ui: boolean;
|
|
38
|
+
ui_settings: Record<string, any>;
|
|
21
39
|
chat_template: string;
|
|
22
|
-
chat_template_caps:
|
|
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 {
|
package/src/managers/command.ts
CHANGED
|
@@ -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(
|
|
185
|
+
.finally(() => {
|
|
186
|
+
cleanupProgress();
|
|
187
|
+
EventManager.resetInflightModel();
|
|
188
|
+
});
|
|
175
189
|
}
|
|
176
190
|
}
|
|
177
191
|
|
package/src/managers/server.ts
CHANGED
|
@@ -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
|
*
|
package/src/models/baseModel.ts
CHANGED
|
@@ -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
|
|
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
|
-
|
|
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 {
|
|
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 {
|
|
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 {
|
|
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
|
|
19
|
-
private
|
|
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.
|
|
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.
|
|
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.
|
|
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.
|
|
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<
|
|
144
|
-
return await this.
|
|
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.
|
|
160
|
-
return await this.
|
|
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
|
+
}
|
package/src/sse/types.ts
ADDED
|
@@ -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;
|