Skip to content
Draft
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
32 changes: 23 additions & 9 deletions packages/appkit/src/connectors/lakebase/endpoint-host.ts
Original file line number Diff line number Diff line change
Expand Up @@ -74,16 +74,30 @@ async function compareHosts(
reason,
);
try {
const path = `/api/2.0/postgres/${endpoint}`;
const lookup =
"request" in client
? // AppKit's modular client: throws ApiError (with statusCode) on non-2xx.
// `signal` isn't in lakebase's public request type but AppKit honors it.
client
.request({
method: "GET",
path,
headers: { Accept: "application/json" },
signal: controller.signal,
} as Parameters<typeof client.request>[0])
.then((res) => res.json())
: client.apiClient.request(
{
path,
method: "GET",
headers: new Headers({ Accept: "application/json" }),
raw: false,
},
contextFromAbortSignal(controller.signal),
);
const response = await Promise.race([
client.apiClient.request(
{
path: `/api/2.0/postgres/${endpoint}`,
method: "GET",
headers: new Headers({ Accept: "application/json" }),
raw: false,
},
contextFromAbortSignal(controller.signal),
),
lookup,
new Promise<undefined>((resolve) => {
timer = setTimeout(() => {
controller.abort();
Expand Down
2 changes: 1 addition & 1 deletion packages/appkit/src/connectors/lakebase/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ export async function initializeLakebasePool(
const client = ServiceContext.isInitialized()
? ServiceContext.get().client
: createWorkspaceClient({ clientOptions: getClientOptions() });
resolved.workspaceClient = client.toLegacyWorkspaceClient();
resolved.workspaceClient = client;
}
const [user] = await Promise.all([
getUsernameWithApiLookup(resolved),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -316,3 +316,69 @@ describe("assertEndpointHostMatches", () => {
]);
});
});

// AppKit passes its modular workspace client (with `request()`) to lakebase,
// so in apps the lookup goes through `request`, not the legacy `apiClient`.
describe("assertEndpointHostMatches with the modular request() client", () => {
function modularClient(
impl: (req: { signal?: AbortSignal }) => Promise<Response>,
) {
const request = vi.fn(impl);
return { request, client: { request } as unknown as Client };
}

test("rejects a mismatch read through request()", async () => {
const { endpoint, host } = names();
const { request, client } = modularClient(async () =>
Response.json(endpointWith({ host: "ep-expected.database.example.com" })),
);
vi.spyOn(console, "error").mockImplementation(() => undefined);

await expect(
assertEndpointHostMatches({ endpoint, host, workspaceClient: client }),
).rejects.toThrow(host);
expect(request).toHaveBeenCalledWith(
expect.objectContaining({
method: "GET",
path: `/api/2.0/postgres/${endpoint}`,
signal: expect.any(AbortSignal),
}),
);
});

test("warns when request() reports the endpoint no longer exists (404)", async () => {
const { endpoint, host } = names();
const { client } = modularClient(async () => {
throw Object.assign(new Error("not found"), { statusCode: 404 });
});

await assertEndpointHostMatches({
endpoint,
host,
workspaceClient: client,
});
expect(warnings()).toContain("was not found");
});

test("aborts the request() lookup when the deadline expires", async () => {
vi.useFakeTimers();
const { endpoint, host } = names();
let seen: AbortSignal | undefined;
const { client } = modularClient(
(req) =>
new Promise<Response>(() => {
seen = req.signal;
}),
);

const check = assertEndpointHostMatches({
endpoint,
host,
workspaceClient: client,
});
await vi.advanceTimersByTimeAsync(3_000);
await check;
expect(seen?.aborted).toBe(true);
expect(warnings()).toContain("timed out");
});
});
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,7 @@ const mocks = vi.hoisted(() => {
request,
client,
createPool: vi.fn(),
createWorkspaceClient: vi.fn(() => ({
toLegacyWorkspaceClient: () => client,
})),
createWorkspaceClient: vi.fn(() => client),
};
});

Expand Down Expand Up @@ -90,7 +88,7 @@ describe("AppKit Lakebase connector initialization", () => {
};
vi.mocked(ServiceContext.isInitialized).mockReturnValue(true);
vi.spyOn(ServiceContext, "get").mockReturnValue({
client: { toLegacyWorkspaceClient: () => legacy },
client: legacy,
} as unknown as ReturnType<typeof ServiceContext.get>);
await initializeLakebasePool();
expect(poolConfig()).toMatchObject({
Expand Down Expand Up @@ -125,7 +123,7 @@ describe("AppKit Lakebase connector initialization", () => {
const requestClient = { currentUser: { me: requestLookup } };
await runInUserContext(
{
client: { toLegacyWorkspaceClient: () => requestClient },
client: requestClient,
userId: "request-user-id",
userEmail: "request-user@example.test",
workspaceId: Promise.resolve("workspace"),
Expand Down
6 changes: 3 additions & 3 deletions packages/appkit/src/plugins/lakebase/lakebase.ts
Original file line number Diff line number Diff line change
Expand Up @@ -84,7 +84,7 @@ export class LakebasePlugin extends Plugin implements ToolProvider {
this.config.pool?.workspaceClient ??
createWorkspaceClient({
clientOptions: getClientOptions(),
}).toLegacyWorkspaceClient(),
}),
};
const [user] = await Promise.all([
getUsernameWithApiLookup(poolConfig),
Expand All @@ -109,7 +109,7 @@ export class LakebasePlugin extends Plugin implements ToolProvider {
const pool = oboManager.getPool(
userKey,
{
workspaceClient: ctx.client.toLegacyWorkspaceClient(),
workspaceClient: ctx.client,
user: userKey,
},
ctx.tokenFingerprint,
Expand Down Expand Up @@ -305,7 +305,7 @@ export class LakebasePlugin extends Plugin implements ToolProvider {
const user = ctx.principal.userEmail ?? ctx.principal.userId;
return {
...this.config.pool,
workspaceClient: ctx.client.toLegacyWorkspaceClient(),
workspaceClient: ctx.client,
user,
};
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -373,7 +373,7 @@ describe("LakebasePlugin - OBO via RoutingPool", () => {
await plugin.setup();

const userCtx = {
client: { toLegacyWorkspaceClient: () => ({}) } as any,
client: {} as any,
userId: "user-123",
userEmail: "alice@example.com",
workspaceId: Promise.resolve("ws-1"),
Expand Down Expand Up @@ -402,7 +402,7 @@ describe("LakebasePlugin - OBO via RoutingPool", () => {
await plugin.setup();

const userCtx = {
client: { toLegacyWorkspaceClient: () => ({}) } as any,
client: {} as any,
userId: "user-123",
userEmail: "alice@example.com",
workspaceId: Promise.resolve("ws-1"),
Expand Down Expand Up @@ -431,7 +431,7 @@ describe("LakebasePlugin - OBO via RoutingPool", () => {
await plugin.setup();

const userCtx = {
client: { toLegacyWorkspaceClient: () => ({}) } as any,
client: {} as any,
userId: "user-123",
workspaceId: Promise.resolve("ws-1"),
isUserContext: true as const,
Expand Down
3 changes: 2 additions & 1 deletion packages/lakebase/package.json
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,8 @@
"release:sbom": "pnpm exec cdxgen -t js --no-recurse --required-only -o tmp/sbom.cdx.json ."
},
"dependencies": {
"@databricks/sdk-experimental": "0.17.0",
"@databricks/sdk-auth": "0.51.0",
"@databricks/sdk-core": "0.51.0",
"pg": "8.18.0",
"@opentelemetry/api": "1.9.0"
},
Expand Down
140 changes: 80 additions & 60 deletions packages/lakebase/src/__tests__/credentials.test.ts
Original file line number Diff line number Diff line change
@@ -1,47 +1,24 @@
import type { WorkspaceClient } from "@databricks/sdk-experimental";
import { ApiClient, Config } from "@databricks/sdk-experimental";
import { beforeEach, describe, expect, it, vi } from "vitest";

import { getUsernameWithApiLookup } from "../config";
import { generateDatabaseCredential } from "../credentials";
import {
type DatabaseCredential,
type LegacyWorkspaceClientLike,
RequestedClaimsPermissionSet,
} from "../types";

// Mock the @databricks/sdk-experimental module
vi.mock("@databricks/sdk-experimental", () => {
const mockRequest = vi.fn();

return {
Config: vi.fn(),
ApiClient: vi.fn().mockImplementation(() => ({
request: mockRequest,
})),
};
});

describe("Lakebase Authentication", () => {
let mockWorkspaceClient: WorkspaceClient;
let mockApiClient: ApiClient;
let mockWorkspaceClient: LegacyWorkspaceClientLike;
let mockApiClient: LegacyWorkspaceClientLike["apiClient"];

beforeEach(() => {
vi.clearAllMocks();

// Get the mocked ApiClient constructor
const ApiClientConstructor = ApiClient as unknown as ReturnType<
typeof vi.fn
>;
mockApiClient = new ApiClientConstructor(
new Config({ host: "https://test.databricks.com" }),
);

// Setup mock workspace client with apiClient
mockApiClient = { request: vi.fn() };
mockWorkspaceClient = {
config: {
host: "https://test.databricks.com",
},
currentUser: { me: vi.fn() },
apiClient: mockApiClient,
} as WorkspaceClient;
};
});

describe("generateDatabaseCredential", () => {
Expand Down Expand Up @@ -145,44 +122,87 @@ describe("Lakebase Authentication", () => {
}),
).rejects.toThrow("API request failed");
});
});

it("should use correct workspace host for API calls", async () => {
const customHost = "https://custom-workspace.databricks.com";
describe("request-capable (AppKit modular) client", () => {
const credential: DatabaseCredential = {
token: "modular-token",
expire_time: "2026-02-06T18:00:00Z",
};

it("posts the snake_case request body through request()", async () => {
const request = vi.fn(async () => Response.json(credential));
const claims = [
{
permission_set: RequestedClaimsPermissionSet.READ_ONLY,
resources: [{ table_name: "catalog.schema.users" }],
},
];

// Create a new mock API client for the custom workspace
const ApiClientConstructor = ApiClient as unknown as ReturnType<
typeof vi.fn
>;
const customApiClient = new ApiClientConstructor(
new Config({ host: customHost }),
const result = await generateDatabaseCredential(
{ request },
{ endpoint: "projects/p/branches/main/endpoints/primary", claims },
);

const customWorkspaceClient = {
config: { host: customHost },
apiClient: customApiClient,
} as WorkspaceClient;
expect(result).toEqual(credential);
expect(request).toHaveBeenCalledWith({
method: "POST",
path: "/api/2.0/postgres/credentials",
headers: {
Accept: "application/json",
"Content-Type": "application/json",
},
body: JSON.stringify({
endpoint: "projects/p/branches/main/endpoints/primary",
claims,
}),
});
});

const mockCredential: DatabaseCredential = {
token: "mock-token",
expire_time: "2026-02-06T18:00:00Z",
};
it("prefers request() over a legacy apiClient on the same object", async () => {
// The AppKit facade exposes both; the modular path must win.
const request = vi.fn(async () => Response.json(credential));
const client = { ...mockWorkspaceClient, request };

vi.mocked(customApiClient.request).mockResolvedValue(mockCredential);
await generateDatabaseCredential(client, { endpoint: "e" });

await generateDatabaseCredential(customWorkspaceClient, {
endpoint: "projects/test/branches/main/endpoints/primary",
});
expect(request).toHaveBeenCalledOnce();
expect(mockApiClient.request).not.toHaveBeenCalled();
});

// Verify the request was made with the correct workspace client
expect(customApiClient.request).toHaveBeenCalledWith({
path: "/api/2.0/postgres/credentials",
method: "POST",
headers: expect.any(Headers),
raw: false,
payload: {
endpoint: "projects/test/branches/main/endpoints/primary",
},
});
it("rejects a malformed credential response", async () => {
const request = vi.fn(async () => Response.json({ token: "t" }));
await expect(
generateDatabaseCredential({ request }, { endpoint: "e" }),
).rejects.toThrow();
});

it("resolves the username via the SCIM Me endpoint", async () => {
const prev = {
PGUSER: process.env.PGUSER,
DATABRICKS_CLIENT_ID: process.env.DATABRICKS_CLIENT_ID,
};
delete process.env.PGUSER;
delete process.env.DATABRICKS_CLIENT_ID;
try {
const request = vi.fn(async () =>
Response.json({ userName: "someone@example.com" }),
);
await expect(
getUsernameWithApiLookup({ workspaceClient: { request } }),
).resolves.toBe("someone@example.com");
expect(request).toHaveBeenCalledWith(
expect.objectContaining({
method: "GET",
path: "/api/2.0/preview/scim/v2/Me",
}),
);
} finally {
Object.assign(process.env, prev);
for (const [k, v] of Object.entries(prev)) {
if (v === undefined) delete process.env[k];
}
}
});
});
});
Loading