Archived
Merge pull request #64 from jmfederico/fix/issue-62-authstorage
fix: migrate auth/model plumbing to ModelRuntime (fixes #62)
This commit is contained in:
@@ -1,9 +1,12 @@
|
||||
import { createAssistantMessageEventStream, type AssistantMessage } from "@earendil-works/pi-ai";
|
||||
import { mkdtemp, rm, writeFile } from "node:fs/promises";
|
||||
import { tmpdir } from "node:os";
|
||||
import { join } from "node:path";
|
||||
import { createAssistantMessageEventStream, InMemoryCredentialStore, type AssistantMessage } from "@earendil-works/pi-ai";
|
||||
import type { StreamFn } from "@earendil-works/pi-agent-core";
|
||||
import { AuthStorage, ModelRegistry } from "@earendil-works/pi-coding-agent";
|
||||
import { ModelRuntime } from "@earendil-works/pi-coding-agent";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { PiSessionService } from "./piSessionService.js";
|
||||
import { CapturingSessionEventHub, fakeRuntime, runtimeCreator, sessionGateway, sessionRecord, sessionRef, TEST_MODEL_ID, TEST_MODEL_PROVIDER, testModel, type RuntimeCreator } from "./piSessionService.testSupport.js";
|
||||
import { CapturingSessionEventHub, createTestModelRuntime, fakeRuntime, runtimeCreator, seedCredential, sessionGateway, sessionRecord, sessionRef, TEST_MODEL_ID, TEST_MODEL_PROVIDER, testModel, testModelRuntime, type RuntimeCreator } from "./piSessionService.testSupport.js";
|
||||
|
||||
const TEST_AGENT_DIR = "/tmp/pi-web-test-agent";
|
||||
|
||||
@@ -12,6 +15,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const fake = fakeRuntime("prompt-session");
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("prompt-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -30,6 +34,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const hub = new CapturingSessionEventHub();
|
||||
const service = new PiSessionService(hub, {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("echo-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -60,6 +65,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
};
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime,
|
||||
sessionManager: sessionGateway([sessionRecord("prompt-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -96,6 +102,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const fake = fakeRuntime("name-session", { model, agent: { streamFn } });
|
||||
const service = new PiSessionService(hub, {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("name-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -118,6 +125,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
});
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("status-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -139,6 +147,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
});
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("dedupe-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -155,6 +164,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const fake = fakeRuntime("queued-session", { isStreaming: true });
|
||||
const service = new PiSessionService(hub, {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("queued-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -181,6 +191,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
};
|
||||
const service = new PiSessionService(hub, {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("compacting-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -250,6 +261,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
fake.session.clearQueue = clearRuntimeQueue;
|
||||
const service = new PiSessionService(hub, {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("clear-queue-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -290,6 +302,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const fake = fakeRuntime("clear-empty-queue-session");
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("clear-empty-queue-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -309,6 +322,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const fake = fakeRuntime("abort-session");
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("abort-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -326,6 +340,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const fake = fakeRuntime("abort-compaction-session", { isCompacting: true });
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("abort-compaction-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -341,17 +356,66 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
await service.dispose();
|
||||
});
|
||||
|
||||
it("reloads models.json before listing and selecting models", async () => {
|
||||
const agentDir = await mkdtemp(join(tmpdir(), "pi-web-model-runtime-"));
|
||||
try {
|
||||
const modelsPath = join(agentDir, "models.json");
|
||||
await writeLocalModelsConfig(modelsPath, "initial-model");
|
||||
const modelRuntime = await ModelRuntime.create({
|
||||
credentials: new InMemoryCredentialStore(),
|
||||
modelsPath,
|
||||
allowModelNetwork: false,
|
||||
});
|
||||
const setSessionModel = vi.fn(() => Promise.resolve());
|
||||
const fake = fakeRuntime("models-session", { modelRuntime, setModel: setSessionModel });
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir,
|
||||
modelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("models-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
});
|
||||
|
||||
try {
|
||||
await writeLocalModelsConfig(modelsPath, "listed-model");
|
||||
const listed = await service.availableModels(sessionRef("models-session"));
|
||||
expect(listed).toEqual(expect.arrayContaining([
|
||||
expect.objectContaining({ provider: "test-local", id: "listed-model" }),
|
||||
]));
|
||||
expect(listed).not.toEqual(expect.arrayContaining([
|
||||
expect.objectContaining({ provider: "test-local", id: "initial-model" }),
|
||||
]));
|
||||
|
||||
await writeLocalModelsConfig(modelsPath, "selected-model");
|
||||
await expect(service.setModel(sessionRef("models-session"), "test-local", "selected-model")).resolves.toBeDefined();
|
||||
expect(setSessionModel).toHaveBeenCalledWith(expect.objectContaining({
|
||||
provider: "test-local",
|
||||
id: "selected-model",
|
||||
}));
|
||||
} finally {
|
||||
await service.dispose();
|
||||
}
|
||||
} finally {
|
||||
await rm(agentDir, { recursive: true, force: true });
|
||||
}
|
||||
});
|
||||
|
||||
it("refreshes auth state and dedupes warnings when logout removes the current model's credentials", async () => {
|
||||
const hub = new CapturingSessionEventHub();
|
||||
const authStorage = AuthStorage.inMemory({ anthropic: { type: "api_key", key: "sk-test" } });
|
||||
const modelRegistry = ModelRegistry.inMemory(authStorage);
|
||||
const model = modelRegistry.find(TEST_MODEL_PROVIDER, TEST_MODEL_ID);
|
||||
// The shared model runtime reads a live credential store. Mutating the store
|
||||
// and refreshing here simulates the committed snapshot that
|
||||
// ModelRuntime.login()/logout() establishes before AuthService emits.
|
||||
// applyAuthChange then only needs to notify active sessions.
|
||||
const credentials = new InMemoryCredentialStore();
|
||||
await seedCredential(credentials, "anthropic", { type: "api_key", key: "sk-test" });
|
||||
const modelRuntime = await createTestModelRuntime(credentials);
|
||||
const model = modelRuntime.getModel(TEST_MODEL_PROVIDER, TEST_MODEL_ID);
|
||||
if (model === undefined) throw new Error("Expected Anthropic model fixture");
|
||||
const fake = fakeRuntime("auth-session", { model, modelRegistry });
|
||||
const fake = fakeRuntime("auth-session", { model, modelRuntime });
|
||||
|
||||
const service = new PiSessionService(hub, {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRegistry,
|
||||
modelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("auth-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -361,7 +425,8 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
hub.sessionEvents.length = 0;
|
||||
hub.globalEvents.length = 0;
|
||||
|
||||
authStorage.logout("anthropic");
|
||||
await credentials.delete("anthropic");
|
||||
await modelRuntime.refresh();
|
||||
service.applyAuthChange({ removedProviderId: "anthropic" });
|
||||
service.applyAuthChange({ removedProviderId: "anthropic" });
|
||||
|
||||
@@ -369,9 +434,11 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
expect(warningCount()).toBe(1);
|
||||
expect(hub.globalEvents.some((event) => event.type === "status.update" && event.status.sessionId === "auth-session")).toBe(true);
|
||||
|
||||
authStorage.set("anthropic", { type: "api_key", key: "sk-new" });
|
||||
await seedCredential(credentials, "anthropic", { type: "api_key", key: "sk-new" });
|
||||
await modelRuntime.refresh();
|
||||
service.applyAuthChange();
|
||||
authStorage.logout("anthropic");
|
||||
await credentials.delete("anthropic");
|
||||
await modelRuntime.refresh();
|
||||
service.applyAuthChange({ removedProviderId: "anthropic" });
|
||||
expect(warningCount()).toBe(2);
|
||||
|
||||
@@ -382,6 +449,7 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
const fake = fakeRuntime("stop-session");
|
||||
const service = new PiSessionService(new CapturingSessionEventHub(), {
|
||||
agentDir: TEST_AGENT_DIR,
|
||||
modelRuntime: testModelRuntime,
|
||||
createAgentRuntime: runtimeCreator(fake.runtime),
|
||||
sessionManager: sessionGateway([sessionRecord("stop-session")]),
|
||||
heartbeatIntervalMs: 60_000,
|
||||
@@ -394,3 +462,25 @@ describe("PiSessionService prompt, queue, and auth warnings", () => {
|
||||
await service.dispose();
|
||||
});
|
||||
});
|
||||
|
||||
async function writeLocalModelsConfig(path: string, modelId: string): Promise<void> {
|
||||
await writeFile(path, JSON.stringify({
|
||||
providers: {
|
||||
"test-local": {
|
||||
name: "Test Local",
|
||||
baseUrl: "http://127.0.0.1:1234/v1",
|
||||
apiKey: "offline-test-key",
|
||||
api: "openai-completions",
|
||||
models: [{
|
||||
id: modelId,
|
||||
name: modelId,
|
||||
reasoning: false,
|
||||
input: ["text"],
|
||||
cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
|
||||
contextWindow: 1_000,
|
||||
maxTokens: 100,
|
||||
}],
|
||||
},
|
||||
},
|
||||
}));
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user