Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions db/services/settings.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import { and, eq } from "drizzle-orm";
import type { AccessScope } from "@shared/identity/access-scope";
import { db, settings } from "@db";
import { ensureScope } from "@db/services/scope";

const gatewayModelKey = "gateway_model";
const defaultGatewayModel = "openai/gpt-6.1-sol-fast";
Expand All @@ -24,6 +25,7 @@ export async function getGatewayModel(scope: AccessScope) {
}

export async function selectGatewayModel(scope: AccessScope, modelId: string) {
await ensureScope(scope);
await db
.insert(settings)
.values({
Expand Down
64 changes: 64 additions & 0 deletions web/trpc/model-settings.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
import { PGlite } from "@electric-sql/pglite";
import { drizzle } from "drizzle-orm/pglite";
import { migrate } from "drizzle-orm/pglite/migrator";
import {
afterAll,
beforeAll,
beforeEach,
describe,
expect,
it,
vi,
} from "vitest";
import * as Database from "@db";
import * as schema from "@db/schema";
import { getGatewayModel } from "@db/services/settings";
import { ensureScope } from "@db/services/scope";
import { accessScopeForUser } from "@shared/identity/access-scope";
import { appRouter } from "./router";

const client = new PGlite();
const database = drizzle(client, { schema });
const scope = accessScopeForUser("first-model-selection");
const caller = appRouter.createCaller({ origin: "https://example.com", scope });
const defaultModel = "openai/gpt-6.1-sol-fast";

beforeAll(async () => {
await migrate(database, { migrationsFolder: "db/migrations" });
// SAFETY: PGlite implements the real Drizzle query-builder contract; only the database driver is replaced.
// oxlint-disable-next-line typescript/no-unsafe-type-assertion -- Real committed migrations and real tRPC/service caller use an owned PostgreSQL-compatible database.
vi.spyOn(Database, "db", "get").mockReturnValue(database as never);
}, 20_000);

beforeEach(async () => {
await database.delete(schema.workspaces);
});
afterAll(async () => {
vi.restoreAllMocks();
await client.close();
});

describe("workspace model selection", () => {
it("saves the first model selection before a chat or vault initializes the workspace", async () => {
expect(await getGatewayModel(scope)).toBe(defaultModel);
expect(await database.select().from(schema.workspaces)).toHaveLength(0);
await caller.settings.selectModel({ modelId: defaultModel });
expect(await database.select().from(schema.settings)).toMatchObject([
{ value: defaultModel, workspaceId: scope.workspaceId },
]);
expect(
await database.select().from(schema.workspaceMemberships)
).toMatchObject([{ userId: scope.userId, workspaceId: scope.workspaceId }]);
});

it("still updates an already initialized workspace without duplicate rows", async () => {
await ensureScope(scope);
await caller.settings.selectModel({ modelId: defaultModel });
await caller.settings.selectModel({ modelId: "meta/muse-spark-1.3" });
expect(await getGatewayModel(scope)).toBe("meta/muse-spark-1.3");
expect(await database.select().from(schema.settings)).toHaveLength(1);
expect(
await database.select().from(schema.workspaceMemberships)
).toHaveLength(1);
});
});