From f90c6b84d0b693b7bd1125dd07ea5453f7355400 Mon Sep 17 00:00:00 2001 From: MarioCadenas Date: Wed, 7 Oct 2026 16:22:24 +0200 Subject: [PATCH] feat(shared): add a modular auth + raw-request seam to the workspace client Add getHost(), authenticate(headers) and request(req) to the WorkspaceClient facade, built only from mapToClientOptions and resolved the same way the modular SDK's resolveClientConfig does (profile resolve + defaultCredentials). request() sends through the modular transport, so it carries the AppKit User-Agent, returns the raw Response, and throws the wrapper ApiError on non-2xx. Migrate the drop-in callers off the legacy config/apiClient: SCIM workspace id probe, dev warehouse discovery (still skip_cannot_use=true), internal telemetry, files upload, mlflow eval auth, and agents hosted-tool auth. Also cache env-M2M OAuth tokens until 40s before expiry. sdk-auth 0.51.0's newM2mCredentials mints a new token per request, unlike the legacy SDK; this affected the already-migrated warehouses/statementExecution clients too. Co-authored-by: Isaac Signed-off-by: MarioCadenas --- .../appkit/src/connectors/files/client.ts | 60 ++-- .../src/connectors/files/tests/client.test.ts | 113 +++---- packages/appkit/src/connectors/mlflow/auth.ts | 6 +- .../src/connectors/mlflow/tests/auth.test.ts | 38 +++ .../appkit/src/context/service-context.ts | 14 +- .../src/context/tests/service-context.test.ts | 11 +- .../core/tests/appkit-as-user-exports.test.ts | 8 +- .../appkit/src/internal-telemetry/reporter.ts | 8 +- .../internal-telemetry/tests/reporter.test.ts | 13 +- packages/appkit/src/plugins/agents/agents.ts | 5 +- packages/appkit/src/plugins/files/plugin.ts | 6 +- .../src/plugins/files/tests/plugin.test.ts | 110 ++----- .../src/resources/tests/warehouse.test.ts | 12 +- packages/appkit/src/resources/warehouse.ts | 14 +- .../src/testing/mock-workspace-client.ts | 19 ++ .../shared/src/workspace-client/client.ts | 24 ++ .../shared/src/workspace-client/modular.ts | 189 ++++++++++-- .../workspace-client/tests/modular.test.ts | 290 ++++++++++++++---- packages/shared/src/workspace-client/types.ts | 20 +- 19 files changed, 620 insertions(+), 340 deletions(-) create mode 100644 packages/appkit/src/connectors/mlflow/tests/auth.test.ts diff --git a/packages/appkit/src/connectors/files/client.ts b/packages/appkit/src/connectors/files/client.ts index ffdc06dac..006fdb63e 100644 --- a/packages/appkit/src/connectors/files/client.ts +++ b/packages/appkit/src/connectors/files/client.ts @@ -294,49 +294,37 @@ export class FilesConnector { const body = contents; const overwrite = options?.overwrite ?? true; - // Workaround: The SDK's files.upload() has two bugs: - // 1. It ignores the `contents` field (sets body to undefined) - // 2. apiClient.request() checks `instanceof` against its own ReadableStream - // subclass, so standard ReadableStream instances get JSON.stringified to "{}" - // Bypass both by calling the REST API directly with SDK-provided auth. - const hostValue = client.config.host; - if (!hostValue) { - throw new Error( - "Databricks host is not configured. Set DATABRICKS_HOST or configure client.config.host.", - ); - } - const host = hostValue.startsWith("http") - ? hostValue - : `https://${hostValue}`; - const url = new URL(`/api/2.0/fs/files${resolvedPath}`, host); - url.searchParams.set("overwrite", String(overwrite)); - - const headers = new Headers({ + // Workaround: the legacy SDK's files.upload() ignores `contents` and + // JSON-stringifies standard ReadableStreams, so PUT the REST API directly + // through the modular transport (auth + AppKit User-Agent; it sets + // `duplex: "half"` for stream bodies). + const headers: Record = { "Content-Type": "application/octet-stream", - "User-Agent": client.apiClient.userAgent(), - }); - const fetchOptions: RequestInit = { method: "PUT", headers, body }; - - if (body instanceof ReadableStream) { - fetchOptions.duplex = "half"; - } else if (body instanceof Buffer) { - headers.set("Content-Length", String(body.length)); + }; + if (body instanceof Buffer) { + headers["Content-Length"] = String(body.length); } else if (typeof body === "string") { - headers.set("Content-Length", String(Buffer.byteLength(body))); + headers["Content-Length"] = String(Buffer.byteLength(body)); } - await client.config.authenticate(headers); - - const res = await fetch(url.toString(), fetchOptions); - - if (!res.ok) { - const text = await res.text(); - logger.error(`Upload failed (${res.status}): ${text}`); - const safeMessage = text.length > 200 ? `${text.slice(0, 200)}…` : text; + try { + const res = await client.request({ + method: "PUT", + path: `/api/2.0/fs/files${resolvedPath}`, + query: { overwrite: String(overwrite) }, + headers, + body, + }); + await res.body?.cancel(); + } catch (e) { + if (!(e instanceof ApiError)) throw e; + logger.error(`Upload failed (${e.statusCode}): ${e.message}`); + const safeMessage = + e.message.length > 200 ? `${e.message.slice(0, 200)}…` : e.message; throw new ApiError( `Upload failed: ${safeMessage}`, "UPLOAD_FAILED", - res.status, + e.statusCode, undefined, [], ); diff --git a/packages/appkit/src/connectors/files/tests/client.test.ts b/packages/appkit/src/connectors/files/tests/client.test.ts index 09be16e7c..a04a010c4 100644 --- a/packages/appkit/src/connectors/files/tests/client.test.ts +++ b/packages/appkit/src/connectors/files/tests/client.test.ts @@ -7,7 +7,7 @@ import { ApiError } from "../../../workspace-client"; import { FilesConnector } from "../client"; import { streamFromChunks, streamFromString } from "./utils"; -const { mockFilesApi, mockConfig, mockClient } = vi.hoisted(() => { +const { mockFilesApi, mockRequest, mockClient } = vi.hoisted(() => { const mockFilesApi = { listDirectoryContents: vi.fn(), download: vi.fn(), @@ -17,21 +17,13 @@ const { mockFilesApi, mockConfig, mockClient } = vi.hoisted(() => { delete: vi.fn(), }; - const mockConfig = { - host: "https://test.databricks.com", - authenticate: vi.fn(), - }; - - const mockApiClient = { - userAgent: vi.fn(() => "@databricks/appkit/9.9.9"), - }; + const mockRequest = vi.fn(); const mockClient = { files: mockFilesApi, - config: mockConfig, - apiClient: mockApiClient, + request: mockRequest, } as unknown as WorkspaceClient; - return { mockFilesApi, mockConfig, mockClient }; + return { mockFilesApi, mockRequest, mockClient }; }); vi.mock("../../../workspace-client", async (importOriginal) => { @@ -459,32 +451,29 @@ describe("FilesConnector", () => { describe("upload()", () => { let connector: FilesConnector; - let fetchSpy: ReturnType; beforeEach(() => { vi.clearAllMocks(); connector = new FilesConnector({ defaultVolume: "/Volumes/catalog/schema/vol", }); - mockConfig.authenticate.mockResolvedValue(undefined); - fetchSpy = vi.fn().mockResolvedValue({ ok: true }); - vi.stubGlobal("fetch", fetchSpy); - }); - - afterEach(() => { - vi.unstubAllGlobals(); + mockRequest.mockImplementation(async () => new Response(null)); }); + // URL, auth, User-Agent, and stream `duplex` are the modular transport's + // job (covered in shared's modular.test.ts); here we assert what we send. test("handles string input", async () => { await connector.upload(mockClient, "file.txt", "hello world"); - expect(fetchSpy).toHaveBeenCalledWith( - expect.stringContaining( - "/api/2.0/fs/files/Volumes/catalog/schema/vol/file.txt", - ), + expect(mockRequest).toHaveBeenCalledWith( expect.objectContaining({ method: "PUT", + path: "/api/2.0/fs/files/Volumes/catalog/schema/vol/file.txt", body: "hello world", + headers: { + "Content-Type": "application/octet-stream", + "Content-Length": "11", + }, }), ); }); @@ -493,11 +482,11 @@ describe("FilesConnector", () => { const buf = Buffer.from("buffer data"); await connector.upload(mockClient, "file.bin", buf); - expect(fetchSpy).toHaveBeenCalledWith( - expect.any(String), + expect(mockRequest).toHaveBeenCalledWith( expect.objectContaining({ method: "PUT", body: buf, + headers: expect.objectContaining({ "Content-Length": "11" }), }), ); }); @@ -506,21 +495,15 @@ describe("FilesConnector", () => { const stream = streamFromString("stream data"); await connector.upload(mockClient, "file.txt", stream); - expect(fetchSpy).toHaveBeenCalledWith( - expect.any(String), - expect.objectContaining({ - method: "PUT", - body: expect.any(ReadableStream), - duplex: "half", - }), - ); + const req = mockRequest.mock.calls[0][0]; + expect(req.body).toBe(stream); + expect(req.headers["Content-Length"]).toBeUndefined(); }); test("defaults overwrite to true", async () => { await connector.upload(mockClient, "file.txt", "data"); - const url = fetchSpy.mock.calls[0][0] as string; - expect(url).toContain("overwrite=true"); + expect(mockRequest.mock.calls[0][0].query).toEqual({ overwrite: "true" }); }); test("sets overwrite=false when specified", async () => { @@ -528,57 +511,31 @@ describe("FilesConnector", () => { overwrite: false, }); - const url = fetchSpy.mock.calls[0][0] as string; - expect(url).toContain("overwrite=false"); - }); - - test("calls config.authenticate on the headers", async () => { - await connector.upload(mockClient, "file.txt", "data"); - - expect(mockConfig.authenticate).toHaveBeenCalledWith(expect.any(Headers)); - }); - - test("stamps the AppKit User-Agent from the SDK apiClient", async () => { - await connector.upload(mockClient, "file.txt", "data"); - - const init = fetchSpy.mock.calls[0][1] as RequestInit; - const headers = init.headers as Headers; - expect(headers.get("User-Agent")).toBe("@databricks/appkit/9.9.9"); - }); - - test("builds URL from client.config.host", async () => { - await connector.upload(mockClient, "file.txt", "data"); - - const url = fetchSpy.mock.calls[0][0] as string; - expect(url).toMatch( - /^https:\/\/test\.databricks\.com\/api\/2\.0\/fs\/files/, - ); + expect(mockRequest.mock.calls[0][0].query).toEqual({ + overwrite: "false", + }); }); test("throws ApiError on non-ok response", async () => { - fetchSpy.mockResolvedValue({ - ok: false, - status: 403, - text: () => Promise.resolve("Forbidden"), - }); - - await expect( - connector.upload(mockClient, "file.txt", "data"), - ).rejects.toThrow("Upload failed: Forbidden"); + mockRequest.mockRejectedValue( + new ApiError("Forbidden", "PERMISSION_DENIED", 403, undefined, []), + ); - try { - await connector.upload(mockClient, "file.txt", "data"); - } catch (error) { - expect(error).toBeInstanceOf(ApiError); - expect((error as any).statusCode).toBe(403); - } + const error = await connector + .upload(mockClient, "file.txt", "data") + .catch((e: unknown) => e); + expect(error).toBeInstanceOf(ApiError); + expect((error as ApiError).message).toBe("Upload failed: Forbidden"); + expect((error as ApiError).errorCode).toBe("UPLOAD_FAILED"); + expect((error as ApiError).statusCode).toBe(403); }); test("resolves absolute paths directly", async () => { await connector.upload(mockClient, "/Volumes/other/vol/file.txt", "data"); - const url = fetchSpy.mock.calls[0][0] as string; - expect(url).toContain("/api/2.0/fs/files/Volumes/other/vol/file.txt"); + expect(mockRequest.mock.calls[0][0].path).toBe( + "/api/2.0/fs/files/Volumes/other/vol/file.txt", + ); }); }); diff --git a/packages/appkit/src/connectors/mlflow/auth.ts b/packages/appkit/src/connectors/mlflow/auth.ts index 9502aebfd..983419d95 100644 --- a/packages/appkit/src/connectors/mlflow/auth.ts +++ b/packages/appkit/src/connectors/mlflow/auth.ts @@ -49,11 +49,9 @@ async function resolveViaSdk( // Mints the OAuth access token (or reuses a PAT from the profile) and adds // an `Authorization: Bearer ` header — the same call the connectors // use before each request. - await client.config.authenticate(headers); + await client.authenticate(headers); const token = options.token ?? extractBearer(headers); - const host = - options.host ?? - (await client.config.getHost()).toString().replace(/\/+$/, ""); + const host = options.host ?? (await client.getHost()); if (!token || !host) return undefined; return { host, token }; } catch { diff --git a/packages/appkit/src/connectors/mlflow/tests/auth.test.ts b/packages/appkit/src/connectors/mlflow/tests/auth.test.ts new file mode 100644 index 000000000..1123037e2 --- /dev/null +++ b/packages/appkit/src/connectors/mlflow/tests/auth.test.ts @@ -0,0 +1,38 @@ +import { describe, expect, test, vi } from "vitest"; + +const { createWorkspaceClient } = vi.hoisted(() => ({ + createWorkspaceClient: vi.fn(), +})); +vi.mock("../../../workspace-client", () => ({ createWorkspaceClient })); + +import { resolveDatabricksAuth } from "../auth"; + +describe("resolveDatabricksAuth", () => { + test("profile path: takes the bearer + host from the client's modular auth seam", async () => { + createWorkspaceClient.mockReturnValue({ + authenticate: async (headers: Headers) => { + headers.set("Authorization", "Bearer minted"); + }, + getHost: async () => "https://ws.cloud.databricks.com", + }); + + await expect( + resolveDatabricksAuth({ profile: "dogfood" }), + ).resolves.toEqual({ + host: "https://ws.cloud.databricks.com", + token: "minted", + }); + expect(createWorkspaceClient).toHaveBeenCalledWith({ profile: "dogfood" }); + }); + + test("returns undefined when credentials can't be resolved", async () => { + createWorkspaceClient.mockReturnValue({ + authenticate: async () => { + throw new Error("no auth configured"); + }, + getHost: async () => "https://x", + }); + + await expect(resolveDatabricksAuth({})).resolves.toBeUndefined(); + }); +}); diff --git a/packages/appkit/src/context/service-context.ts b/packages/appkit/src/context/service-context.ts index 1d463a06b..789457c93 100644 --- a/packages/appkit/src/context/service-context.ts +++ b/packages/appkit/src/context/service-context.ts @@ -226,20 +226,18 @@ export class ServiceContext { return process.env.DATABRICKS_WORKSPACE_ID; } - const response = (await client.apiClient.request({ + const response = await client.request({ path: "/api/2.0/preview/scim/v2/Me", method: "GET", - headers: new Headers(), - raw: false, - query: {}, - responseHeaders: ["x-databricks-org-id"], - })) as { "x-databricks-org-id": string }; + }); + await response.body?.cancel(); + const workspaceId = response.headers.get("x-databricks-org-id"); - if (!response["x-databricks-org-id"]) { + if (!workspaceId) { throw ConfigurationError.resourceNotFound("Workspace ID"); } - return response["x-databricks-org-id"]; + return workspaceId; } /** diff --git a/packages/appkit/src/context/tests/service-context.test.ts b/packages/appkit/src/context/tests/service-context.test.ts index 4bb5f27f6..3251ea867 100644 --- a/packages/appkit/src/context/tests/service-context.test.ts +++ b/packages/appkit/src/context/tests/service-context.test.ts @@ -19,7 +19,15 @@ const { mockMe, mockApiRequest, MockWorkspaceClient, MockConfigError } = const MockWorkspaceClient = vi.fn().mockImplementation(() => ({ currentUser: { me: mockMe }, - apiClient: { request: mockApiRequest }, + // Tests script legacy-style results; adapt them to the raw `Response` + // `client.request` returns (org id → response header, else JSON body). + request: async (req: unknown) => { + const result = await mockApiRequest(req); + const orgId = result?.["x-databricks-org-id"]; + return orgId !== undefined + ? new Response(null, { headers: { "x-databricks-org-id": orgId } }) + : new Response(JSON.stringify(result ?? {})); + }, })); class MockConfigError extends Error { @@ -411,7 +419,6 @@ describe("ServiceContext", () => { expect.objectContaining({ path: "/api/2.0/preview/scim/v2/Me", method: "GET", - responseHeaders: ["x-databricks-org-id"], }), ); }); diff --git a/packages/appkit/src/core/tests/appkit-as-user-exports.test.ts b/packages/appkit/src/core/tests/appkit-as-user-exports.test.ts index bd6655dbf..cf57c548e 100644 --- a/packages/appkit/src/core/tests/appkit-as-user-exports.test.ts +++ b/packages/appkit/src/core/tests/appkit-as-user-exports.test.ts @@ -98,9 +98,11 @@ import { createApp } from "../appkit"; const { MockWorkspaceClient } = vi.hoisted(() => { const MockWorkspaceClient = vi.fn().mockImplementation(() => ({ currentUser: { me: vi.fn().mockResolvedValue({ id: "sp-user-123" }) }, - apiClient: { - request: vi.fn().mockResolvedValue({ "x-databricks-org-id": "ws-456" }), - }, + request: vi + .fn() + .mockResolvedValue( + new Response(null, { headers: { "x-databricks-org-id": "ws-456" } }), + ), })); return { MockWorkspaceClient }; }); diff --git a/packages/appkit/src/internal-telemetry/reporter.ts b/packages/appkit/src/internal-telemetry/reporter.ts index fefb788f2..cf99c9456 100644 --- a/packages/appkit/src/internal-telemetry/reporter.ts +++ b/packages/appkit/src/internal-telemetry/reporter.ts @@ -174,13 +174,13 @@ export class TelemetryReporter { async #send(logs: AppkitLog[]): Promise { if (logs.length === 0) return; const workspaceId = await this.#workspaceIdPromise; - await this.#client.apiClient.request({ + const response = await this.#client.request({ path: "/telemetry-ext", method: "POST", query: { o: workspaceId }, - headers: new Headers(), - payload: buildAppkitPayload(logs), - raw: false, + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(buildAppkitPayload(logs)), }); + await response.body?.cancel(); } } diff --git a/packages/appkit/src/internal-telemetry/tests/reporter.test.ts b/packages/appkit/src/internal-telemetry/tests/reporter.test.ts index a571e4d3e..41eb5cf2b 100644 --- a/packages/appkit/src/internal-telemetry/tests/reporter.test.ts +++ b/packages/appkit/src/internal-telemetry/tests/reporter.test.ts @@ -9,8 +9,8 @@ function createMockClient(): { client: WorkspaceClient; request: RequestSpy; } { - const request = vi.fn().mockResolvedValue({}); - const client = { apiClient: { request } } as unknown as WorkspaceClient; + const request = vi.fn(async () => new Response("{}")); + const client = { request } as unknown as WorkspaceClient; return { client, request }; } @@ -43,7 +43,7 @@ afterEach(() => { function lastProtoLog(spy: RequestSpy, callIndex = -1) { const calls = spy.mock.calls; const idx = callIndex < 0 ? calls.length + callIndex : callIndex; - const payload = calls[idx][0].payload as { protoLogs: string[] }; + const payload = JSON.parse(calls[idx][0].body) as { protoLogs: string[] }; return JSON.parse(payload.protoLogs[0]); } @@ -52,7 +52,7 @@ describe("TelemetryReporter", () => { expect(TelemetryReporter.getInstance()).toBeNull(); }); - test("sendStartup emits an APP_STARTUP appkit_log via apiClient.request", async () => { + test("sendStartup emits an APP_STARTUP appkit_log via client.request", async () => { const opts = baseOpts(); const reporter = TelemetryReporter.initialize(opts); await reporter.sendStartup(); @@ -63,7 +63,7 @@ describe("TelemetryReporter", () => { path: "/telemetry-ext", method: "POST", query: { o: "1234567890" }, - raw: false, + headers: { "Content-Type": "application/json" }, }); expect(lastProtoLog(opts.__spy).entry.appkit_log).toMatchObject({ event_name: "APP_STARTUP", @@ -95,7 +95,8 @@ describe("TelemetryReporter", () => { await reporter.flushRequestMetrics(); expect(opts.__spy).toHaveBeenCalledOnce(); - const protoLogs = opts.__spy.mock.calls[0][0].payload.protoLogs as string[]; + const protoLogs = JSON.parse(opts.__spy.mock.calls[0][0].body) + .protoLogs as string[]; expect(protoLogs).toHaveLength(2); const events = protoLogs diff --git a/packages/appkit/src/plugins/agents/agents.ts b/packages/appkit/src/plugins/agents/agents.ts index 298c0cf0b..f3a4508ee 100644 --- a/packages/appkit/src/plugins/agents/agents.ts +++ b/packages/appkit/src/plugins/agents/agents.ts @@ -864,11 +864,10 @@ export class AgentsPlugin extends Plugin implements ToolProvider { try { const { getWorkspaceClient } = await import("../../context"); const wsClient = getWorkspaceClient(); - await wsClient.config.ensureResolved(); - host = wsClient.config.host; + host = await wsClient.getHost(); authenticate = async () => { const headers = new Headers(); - await wsClient.config.authenticate(headers); + await wsClient.authenticate(headers); return Object.fromEntries(headers.entries()); }; } catch { diff --git a/packages/appkit/src/plugins/files/plugin.ts b/packages/appkit/src/plugins/files/plugin.ts index 1800b6265..52bf8b913 100644 --- a/packages/appkit/src/plugins/files/plugin.ts +++ b/packages/appkit/src/plugins/files/plugin.ts @@ -1233,10 +1233,10 @@ export class FilesPlugin extends Plugin implements ToolProvider { ); const settings = this._writeSettings(mode); // The connector's `upload` resolves `getWorkspaceClient()` and - // `client.config.authenticate(headers)` synchronously inside this + // sends `client.request(...)` (user-token auth) inside this // callback. When `_runWithAuth` wraps us in `runInCallerContext`, that - // chain produces user-token Authorization headers on the outgoing - // `fetch PUT`. The OBO upload-headers test pins this contract. + // chain sends the PUT on the user-token client. The OBO upload + // identity test pins this contract. const result = await this.trackWrite(() => this.execute(async () => { await connector.upload(getWorkspaceClient(), path, webStream); diff --git a/packages/appkit/src/plugins/files/tests/plugin.test.ts b/packages/appkit/src/plugins/files/tests/plugin.test.ts index f9ae1ad01..a0a3b84b1 100644 --- a/packages/appkit/src/plugins/files/tests/plugin.test.ts +++ b/packages/appkit/src/plugins/files/tests/plugin.test.ts @@ -31,10 +31,7 @@ const { mockClient, MockApiError } = await vi.hoisted(async () => { const mockClient = { files: mockFilesApi, - config: { - host: "https://test.databricks.com", - authenticate: vi.fn(), - }, + request: vi.fn(async () => new Response(null)), }; class MockApiError extends Error { @@ -2658,7 +2655,7 @@ describe("FilesPlugin", () => { * while calls outside the user context fall back to the * service-principal client from `ServiceContext.get()`. Required to * exercise the real `_runWithAuth → runInUserContext → getWorkspaceClient - * → client.config.authenticate → fetch headers` chain. + * → client.request` chain. */ async function useRealGetWorkspaceClient() { const actual = @@ -2689,49 +2686,26 @@ describe("FilesPlugin", () => { }); /** - * NON-NEGOTIABLE upload-headers contract. + * NON-NEGOTIABLE upload identity contract. * - * `_handleUpload` does a hand-rolled `fetch PUT` (not a typed SDK call) - * via the connector's `upload()`. Inside `_runWithAuth` on an OBO volume, - * the chain is: - * - * getWorkspaceClient() → user-token WorkspaceClient - * client.config.authenticate(h) → injects "Bearer " - * fetch(url, { headers }) → outgoing request as the user - * - * This test pins that chain end-to-end. If any future SDK upgrade or - * refactor changes `client.config.authenticate`'s signature, removes the - * `runInUserContext` wrap from `_handleUpload`, or rewires - * `getWorkspaceClient()` so it returns the SP client inside the OBO - * scope, the user-token Authorization header will not reach `fetch` and - * this assertion fails. SP-token would silently leak to UC otherwise. + * `_handleUpload` does a raw PUT (not a typed SDK call) via the connector's + * `upload()`, which sends `client.request(...)` on whatever + * `getWorkspaceClient()` returns. Inside `_runWithAuth` on an OBO volume + * that must be the user-token client, whose modular transport stamps + * "Bearer " (PAT stamping is pinned in shared's + * modular.test.ts). If a refactor removes the `runInUserContext` wrap or + * rewires `getWorkspaceClient()` to the SP client inside the OBO scope, the + * PUT goes out on the SP client and this test fails. SP-token would + * silently leak to UC otherwise. */ - test("OBO upload: outgoing fetch PUT carries user-token Authorization header (not SP)", async () => { + test("OBO upload: outgoing PUT goes through the user-token client (not SP)", async () => { await useRealGetCurrentPrincipalId(); await useRealGetWorkspaceClient(); - // SP-token marker — what the existing mockClient would inject if the - // OBO wrap leaked. We assert this NEVER reaches the outgoing fetch. - mockClient.config.authenticate.mockImplementation( - async (headers: Headers) => { - headers.set("Authorization", "Bearer SP-TOKEN"); - }, - ); - - // User-token marker — what the OBO scope MUST inject. const userClient = { - config: { - host: "https://test.databricks.com", - authenticate: vi.fn(async (headers: Headers) => { - headers.set("Authorization", "Bearer USER-TOKEN-FOO"); - }), - }, // `_handleUpload` only routes through the connector's `upload()`, - // which uses host + authenticate + apiClient.userAgent() + fetch. No - // `files.*` accessor is touched on the user client during this path. - apiClient: { - userAgent: vi.fn(() => "@databricks/appkit/9.9.9"), - }, + // which sends one raw `request`. No `files.*` accessor is touched. + request: vi.fn(async () => new Response(null)), }; // Wire `_buildUserContextOrNull → ServiceContext.createUserContext` to @@ -2746,12 +2720,6 @@ describe("FilesPlugin", () => { }), ); - // Capture the outgoing PUT. - const fetchSpy = vi - .fn() - .mockResolvedValue({ ok: true, status: 200, text: async () => "" }); - vi.stubGlobal("fetch", fetchSpy); - const plugin = new FilesPlugin({ volumes: { obo_vol: { @@ -2778,31 +2746,19 @@ describe("FilesPlugin", () => { res, ); - // The user-token authenticator was consulted exactly when upload ran. - expect(userClient.config.authenticate).toHaveBeenCalledTimes(1); - - // The hand-rolled fetch PUT happened exactly once. - expect(fetchSpy).toHaveBeenCalledTimes(1); - const fetchArgs = fetchSpy.mock.calls[0]; - const init = fetchArgs[1] as RequestInit & { headers: Headers }; - expect(init.method).toBe("PUT"); - - // The contract — proves the user-token Authorization header reached - // fetch. Toggling the `_runWithAuth` wrap off in `_handleUpload` - // breaks this assertion (fetch would carry "Bearer SP-TOKEN" instead). - expect(init.headers.get("Authorization")).toBe("Bearer USER-TOKEN-FOO"); - expect(init.headers.get("Authorization")).not.toBe("Bearer SP-TOKEN"); + // The PUT happened exactly once, on the user-token client. + expect(userClient.request).toHaveBeenCalledTimes(1); + expect(userClient.request).toHaveBeenCalledWith( + expect.objectContaining({ method: "PUT" }), + ); - // Defense-in-depth: SP authenticator was NOT called along the OBO path. - expect(mockClient.config.authenticate).not.toHaveBeenCalled(); + // The contract: the SP client never sent it. + expect(mockClient.request).not.toHaveBeenCalled(); }); test("OBO upload + missing token + NODE_ENV=production → 401 before any SDK or fetch call", async () => { process.env.NODE_ENV = "production"; - const fetchSpy = vi.fn(); - vi.stubGlobal("fetch", fetchSpy); - const plugin = new FilesPlugin({ volumes: { obo_vol: { @@ -2838,9 +2794,9 @@ describe("FilesPlugin", () => { plugin: "files", }); - // Neither the SDK upload nor the hand-rolled fetch ran. + // Neither the SDK upload nor the raw PUT ran. expect(mockClient.files.upload).not.toHaveBeenCalled(); - expect(fetchSpy).not.toHaveBeenCalled(); + expect(mockClient.request).not.toHaveBeenCalled(); }); test("OBO mkdir + policy denies → 403 PolicyDeniedError; SDK not invoked", async () => { @@ -2855,9 +2811,6 @@ describe("FilesPlugin", () => { const handler = getRouteHandler(plugin, "post", "/mkdir"); const res = mockRes(); - const fetchSpy = vi.fn(); - vi.stubGlobal("fetch", fetchSpy); - await handler( mockReq( "obo_vol", @@ -2887,9 +2840,9 @@ describe("FilesPlugin", () => { }), ); - // SDK + fetch not invoked. + // SDK + raw request not invoked. expect(mockClient.files.createDirectory).not.toHaveBeenCalled(); - expect(fetchSpy).not.toHaveBeenCalled(); + expect(mockClient.request).not.toHaveBeenCalled(); }); test("OBO delete + valid token + UC denies → user-token client invoked, error propagated", async () => { @@ -3219,13 +3172,6 @@ describe("FilesPlugin", () => { } }, ); - mockClient.config.authenticate.mockImplementation(async (h: Headers) => { - h.set("Authorization", "Bearer ALICE"); - }); - const fetchSpy = vi - .fn() - .mockResolvedValue({ ok: true, status: 200, text: async () => "" }); - vi.stubGlobal("fetch", fetchSpy); // Bob's first list — empty. const bobRes1 = mockRes(); @@ -3335,10 +3281,6 @@ describe("FilesPlugin", () => { yield { name: "user.txt", path: "/user.txt", is_directory: false }; }); const userClient = { - config: { - host: "https://test.databricks.com", - authenticate: vi.fn(), - }, files: { listDirectoryContents: userListSpy }, }; diff --git a/packages/appkit/src/resources/tests/warehouse.test.ts b/packages/appkit/src/resources/tests/warehouse.test.ts index 7bc134519..f0b0195f5 100644 --- a/packages/appkit/src/resources/tests/warehouse.test.ts +++ b/packages/appkit/src/resources/tests/warehouse.test.ts @@ -67,7 +67,7 @@ describe("warehouse resource bindings", () => { const client = createMockWorkspaceClient(); const bindings = await WarehouseResource.resolve(client, true); expect(await bindings.warehouseId).toBe("configured-warehouse"); - expect(client.apiClient.request).not.toHaveBeenCalled(); + expect(client.request).not.toHaveBeenCalled(); expect(getWarehouseId).toThrow(InitializationError); WarehouseResource.bind(bindings); expect(getWarehouseId()).toBe(bindings.warehouseId); @@ -78,7 +78,7 @@ describe("warehouse resource bindings", () => { const client = createMockWorkspaceClient(); const bindings = await WarehouseResource.resolve(client); expect(bindings.warehouseId).toBeUndefined(); - expect(client.apiClient.request).not.toHaveBeenCalled(); + expect(client.request).not.toHaveBeenCalled(); WarehouseResource.bind(bindings); expect(getWarehouseId).toThrow(ConfigurationError); expect(getWarehouseId).toThrow("No plugin requires a SQL Warehouse"); @@ -95,7 +95,7 @@ describe("warehouse resource bindings", () => { ); WarehouseResource.bind(binding); expect(await getWarehouseId()).toBe("plugin-warehouse"); - expect(client.apiClient.request).not.toHaveBeenCalled(); + expect(client.request).not.toHaveBeenCalled(); }); test("shares one ID across declarations with different environment variables", () => { @@ -174,13 +174,15 @@ describe("warehouse resource bindings", () => { vi.stubEnv("DATABRICKS_WAREHOUSE_ID", ""); vi.stubEnv("DATABRICKS_APPS_AGENTIC_MODE", ""); const client = createMockWorkspaceClient(); - vi.mocked(client.apiClient.request).mockResolvedValue(response); + vi.mocked(client.request).mockResolvedValue( + new Response(JSON.stringify(response)), + ); const error = await WarehouseResource.resolve(client, true).catch( (e: unknown) => e, ); expect(error).toBeInstanceOf(ConfigurationError); expect(error).not.toBeInstanceOf(TypeError); - expect(client.apiClient.request).toHaveBeenCalledTimes(1); + expect(client.request).toHaveBeenCalledTimes(1); }, ); }); diff --git a/packages/appkit/src/resources/warehouse.ts b/packages/appkit/src/resources/warehouse.ts index 461f24699..73971deee 100644 --- a/packages/appkit/src/resources/warehouse.ts +++ b/packages/appkit/src/resources/warehouse.ts @@ -78,13 +78,13 @@ async function discoverWarehouseId(client: WorkspaceClient): Promise { process.env.DATABRICKS_APPS_AGENTIC_MODE === "1"; if (process.env.NODE_ENV === "development" && !agenticMode) { - const response = (await client.apiClient.request({ - path: "/api/2.0/sql/warehouses", - method: "GET", - headers: new Headers(), - raw: false, - query: { skip_cannot_use: "true" }, - })) as { warehouses: sql.EndpointInfo[] }; + const response = (await ( + await client.request({ + path: "/api/2.0/sql/warehouses", + method: "GET", + query: { skip_cannot_use: "true" }, + }) + ).json()) as { warehouses?: sql.EndpointInfo[] }; const priorities: Record = { RUNNING: 0, diff --git a/packages/appkit/src/testing/mock-workspace-client.ts b/packages/appkit/src/testing/mock-workspace-client.ts index 59ec97635..bd65eff3d 100644 --- a/packages/appkit/src/testing/mock-workspace-client.ts +++ b/packages/appkit/src/testing/mock-workspace-client.ts @@ -271,7 +271,26 @@ export function createMockWorkspaceClient( return legacy; } + // Modular auth seam. Seed via `responses: { request: ... }` (a `Response`, or + // a function returning one); `request` defaults to a fresh `{}` JSON body. + const seamDefaults: Record Any> = { + getHost: async () => configTarget.host, + authenticate: async (headers: Headers) => { + headers.set("Authorization", "Bearer test-token"); + }, + request: async () => new Response("{}"), + }; + const seam: Pick = + Object.fromEntries( + Object.entries(seamDefaults).map(([key, impl]) => { + const fn = key in merged ? mint(key) : vi.fn(impl); + fns.set(key, fn); + return [key, fn]; + }), + ) as Any; + const client: WorkspaceClient = { + ...seam, ...(Object.fromEntries( FACADE_SERVICES.map((name) => [name, service(name)]), ) as Pick), diff --git a/packages/shared/src/workspace-client/client.ts b/packages/shared/src/workspace-client/client.ts index adf30c8db..db0fab36e 100644 --- a/packages/shared/src/workspace-client/client.ts +++ b/packages/shared/src/workspace-client/client.ts @@ -15,8 +15,11 @@ import { import { buildStatementExecutionClient, buildWarehousesClient, + buildWorkspaceAuth, type StatementExecutionClient, type WarehousesClient, + type WorkspaceAuth, + type WorkspaceRequest, } from "./modular"; import type { WorkspaceClient } from "./types"; @@ -25,6 +28,7 @@ export class AppKitWorkspaceClient implements WorkspaceClient { #legacy?: LegacyWorkspaceClient; #warehouses?: WarehousesClient; #statementExecution?: StatementExecutionClient; + #auth?: WorkspaceAuth; constructor(opts: WorkspaceClientOptions) { this.#opts = opts; @@ -66,6 +70,19 @@ export class AppKitWorkspaceClient implements WorkspaceClient { return this.#getLegacy().currentUser; } + // Modular auth + raw-request seam — built lazily, independent of the legacy client. + getHost(): Promise { + return this.#getAuth().getHost(); + } + + authenticate(headers: Headers): Promise { + return this.#getAuth().authenticate(headers); + } + + request(req: WorkspaceRequest): Promise { + return this.#getAuth().request(req); + } + get config() { return this.#getLegacy().config; } @@ -78,6 +95,13 @@ export class AppKitWorkspaceClient implements WorkspaceClient { return this.#getLegacy(); } + #getAuth(): WorkspaceAuth { + if (!this.#auth) { + this.#auth = buildWorkspaceAuth(this.#opts); + } + return this.#auth; + } + #getLegacy(): LegacyWorkspaceClient { if (!this.#legacy) { this.#legacy = buildLegacyWorkspaceClient(this.#opts); diff --git a/packages/shared/src/workspace-client/modular.ts b/packages/shared/src/workspace-client/modular.ts index c22590de0..1b1c78754 100644 --- a/packages/shared/src/workspace-client/modular.ts +++ b/packages/shared/src/workspace-client/modular.ts @@ -8,8 +8,9 @@ * * Migrated services are built here as per-service clients; the facade delegates * their accessors to these instead of the legacy monolithic client. Currently - * `warehouses` and `statementExecution` are migrated; every other service still - * routes through `legacy.ts`. + * `warehouses` and `statementExecution` are migrated, plus the auth + raw-request + * seam ({@link buildWorkspaceAuth}); every other service still routes through + * `legacy.ts`. * * NOTE: statementExecution relies on a pinned pnpm patch * (`patches/@databricks__sdk-statementexecution@0.46.0.patch`) that restores the @@ -17,20 +18,28 @@ * unmarshal transform would otherwise strip. */ import { + type Credentials, newTokenCredentials, type Token, type TokenCredentials, tokenProviderFn, } from "@databricks/sdk-auth"; import { + defaultCredentials, newM2mCredentials, newPatCredentials, } from "@databricks/sdk-auth/credentials"; -import { type HttpClient, newFetchHttpClient } from "@databricks/sdk-core/http"; +import { + type HttpClient, + type HttpRequest, + newFetchHttpClient, +} from "@databricks/sdk-core/http"; +import { resolve } from "@databricks/sdk-core/profiles"; import type { ClientOptions } from "@databricks/sdk-options/client"; import { StatementExecutionClient } from "@databricks/sdk-statementexecution/v1"; import { WarehousesClient } from "@databricks/sdk-warehouses/v1"; +import { ApiError } from "./errors"; import type { WorkspaceClientOptions } from "./legacy"; /** @@ -45,37 +54,6 @@ function normalizeHost(host: string | undefined): string | undefined { return /^https?:\/\//i.test(trimmed) ? trimmed : `https://${trimmed}`; } -// Same margin as the legacy SDK: refresh 40s early, since Azure Databricks -// rejects tokens that expire in 30s or less. -const TOKEN_REFRESH_MARGIN_MS = 40_000; - -/** - * Cache a token until shortly before it expires. The modular `newM2mCredentials` - * caches only the token endpoint and mints a fresh OAuth token on EVERY request; - * the legacy SDK reused it until expiry. Concurrent callers share one in-flight - * fetch. Like the legacy SDK, a token without an expiry is reused indefinitely. - */ -function withTokenCache(credentials: TokenCredentials): TokenCredentials { - let current: Token | undefined; - let inflight: Promise | undefined; - const isFresh = (t: Token) => - t.expiry === undefined || - t.expiry.getTime() - TOKEN_REFRESH_MARGIN_MS > Date.now(); - return newTokenCredentials( - credentials.name(), - tokenProviderFn(async () => { - if (current && isFresh(current)) return current; - inflight ??= credentials - .token() - .then((t) => (current = t)) - .finally(() => { - inflight = undefined; - }); - return inflight; - }), - ); -} - /** * Map wrapper options onto the modular SDK's `ClientOptions`, reproducing the * legacy SDK's auth resolution: explicit token → PAT (the OBO path); profile → @@ -177,6 +155,149 @@ function buildHttpClient(opts: WorkspaceClientOptions): HttpClient | undefined { }; } +// Same margin as the legacy SDK: refresh 40s early, since Azure Databricks +// rejects tokens that expire in 30s or less. +const TOKEN_REFRESH_MARGIN_MS = 40_000; + +/** + * Cache a token until shortly before it expires. The modular `newM2mCredentials` + * caches only the token endpoint and mints a fresh OAuth token on EVERY request; + * the legacy SDK reused it until expiry. Concurrent callers share one in-flight + * fetch. Like the legacy SDK, a token without an expiry is reused indefinitely. + */ +function withTokenCache(credentials: TokenCredentials): TokenCredentials { + let current: Token | undefined; + let inflight: Promise | undefined; + const isFresh = (t: Token) => + t.expiry === undefined || + t.expiry.getTime() - TOKEN_REFRESH_MARGIN_MS > Date.now(); + return newTokenCredentials( + credentials.name(), + tokenProviderFn(async () => { + if (current && isFresh(current)) return current; + inflight ??= credentials + .token() + .then((t) => (current = t)) + .finally(() => { + inflight = undefined; + }); + return inflight; + }), + ); +} + +/** A raw REST call against the workspace host. */ +export interface WorkspaceRequest { + method: string; + /** Path on the workspace host, e.g. `/api/2.0/preview/scim/v2/Me`. */ + path: string; + query?: Record; + headers?: Record; + body?: HttpRequest["body"]; + signal?: AbortSignal; +} + +/** Host, auth headers, and raw requests resolved exactly like the modular clients. */ +export interface WorkspaceAuth { + /** Scheme-normalized workspace host, without a trailing slash. */ + getHost(): Promise; + /** Set the auth header(s) (e.g. `Authorization`) on `headers`. */ + authenticate(headers: Headers): Promise; + /** + * Send a request through the modular transport (AppKit User-Agent + auth). + * Returns the raw `Response` (body unread, so it can stream); throws + * {@link ApiError} on a non-2xx status. + */ + request(req: WorkspaceRequest): Promise; +} + +/** + * Build host/credential resolution from {@link mapToClientOptions}, mirroring + * the modular SDK's `resolveClientConfig` (sdk-warehouses `dist/v1/transport.js`): + * resolve the profile from config file + env, explicit options win, and fall + * back to `defaultCredentials` over the resolved profile. Resolved once, lazily; + * a failed resolution is retried on the next call. + */ +export function buildWorkspaceAuth( + opts: WorkspaceClientOptions, +): WorkspaceAuth { + const options = mapToClientOptions(opts); + const transport = options.httpClient ?? newFetchHttpClient(); + let resolved: Promise<{ host: string; credentials: Credentials }> | undefined; + const resolveOnce = () => { + resolved ??= (async () => { + const profile = await resolve(options.profileOptions); + const host = normalizeHost(options.host ?? profile.host)?.replace( + /\/+$/, + "", + ); + if (!host) throw new Error("Host is required."); + const credentials = + options.credentials ?? + defaultCredentials({ profile: { ...profile, host } }); + return { host, credentials }; + })().catch((e) => { + resolved = undefined; + throw e; + }); + return resolved; + }; + const authenticate = async (headers: Headers) => { + const { credentials } = await resolveOnce(); + for (const h of await credentials.authHeaders()) { + headers.set(h.key, h.value); + } + }; + return { + getHost: async () => (await resolveOnce()).host, + authenticate, + async request(req) { + const url = new URL(req.path, (await resolveOnce()).host); + for (const [k, v] of Object.entries(req.query ?? {})) { + url.searchParams.set(k, v); + } + const headers = new Headers(req.headers); + await authenticate(headers); + const res = await transport.send({ + url: url.toString(), + method: req.method, + headers, + body: req.body, + signal: req.signal, + }); + // `Response` rejects any body (even empty) on null-body statuses. + const nullBody = [204, 205, 304].includes(res.statusCode); + const response = new Response(nullBody ? null : res.body, { + status: res.statusCode, + headers: res.headers, + }); + if (!response.ok) throw await toApiError(response); + return response; + }, + }; +} + +/** + * Same error class the legacy `apiClient.request` threw, so existing catch + * sites (`instanceof ApiError`, `.statusCode`, `.errorCode`) keep working. + */ +async function toApiError(response: Response): Promise { + const text = await response.text(); + let parsed: { error_code?: string; message?: string } = {}; + try { + parsed = JSON.parse(text); + } catch { + // Non-JSON error body; fall back to the raw text. + } + return new ApiError( + parsed.message || text || response.statusText, + parsed.error_code ?? "UNKNOWN", + response.status, + undefined, + [], + ); +} + /** Build a modular Warehouses client from wrapper options. */ export function buildWarehousesClient( opts: WorkspaceClientOptions, diff --git a/packages/shared/src/workspace-client/tests/modular.test.ts b/packages/shared/src/workspace-client/tests/modular.test.ts index ae0bda9d2..868059ec1 100644 --- a/packages/shared/src/workspace-client/tests/modular.test.ts +++ b/packages/shared/src/workspace-client/tests/modular.test.ts @@ -3,12 +3,25 @@ import { afterEach, beforeEach, describe, expect, test, vi } from "vitest"; // The wrapper's own tests are the one place allowed to mock the SDK directly. // Capture the `ClientOptions` the modular `WarehousesClient` constructor receives // so we can assert how wrapper options map onto the modular SDK's config. -const { ctorOpts, patTokens, m2mOpts, m2mMint } = vi.hoisted(() => ({ +const { + ctorOpts, + patTokens, + m2mOpts, + m2mToken, + resolveProfile, + defaultCreds, + sent, + nextResponse, +} = vi.hoisted(() => ({ ctorOpts: [] as Array>, patTokens: [] as string[], m2mOpts: [] as Array>, - // Each M2M `token()` call mints a new token, like the real (uncached) SDK. - m2mMint: { count: 0, ttlMs: 60 * 60 * 1000 }, + // Each M2M `token()` call mints a new token, like the real SDK's (uncached). + m2mToken: vi.fn(), + resolveProfile: vi.fn(), + defaultCreds: vi.fn(), + sent: [] as Array<{ url: string; method: string; headers: Headers }>, + nextResponse: { statusCode: 200, body: "{}" }, })); vi.mock("@databricks/sdk-warehouses/v1", () => ({ @@ -23,34 +36,40 @@ vi.mock("@databricks/sdk-statementexecution/v1", () => ({ vi.mock("@databricks/sdk-auth/credentials", () => ({ newPatCredentials: vi.fn((token: string) => { patTokens.push(token); - return { kind: "pat", token }; + return { + kind: "pat", + token, + authHeaders: async () => [ + { key: "Authorization", value: `Bearer ${token}` }, + ], + }; }), newM2mCredentials: vi.fn((opts: Record) => { m2mOpts.push(opts); - return { - name: () => "oauth-m2m", - token: async () => ({ - value: `m2m-token-${++m2mMint.count}`, - expiry: new Date(Date.now() + m2mMint.ttlMs), - }), - }; + return { name: () => "oauth-m2m", token: m2mToken }; }), + defaultCredentials: defaultCreds, })); -// The default transport: its `send` echoes the final request headers so tests -// can assert the User-Agent the wrapper set before delegating. +vi.mock("@databricks/sdk-core/profiles", () => ({ resolve: resolveProfile })); +// The default transport: records each request and echoes its final headers so +// tests can assert the User-Agent / auth the wrapper set before delegating. vi.mock("@databricks/sdk-core/http", () => ({ newFetchHttpClient: vi.fn(() => ({ - send: vi.fn((request: { headers: Headers }) => - Promise.resolve({ - statusCode: 200, - headers: request.headers, - body: null, - }), + send: vi.fn( + (request: { url: string; method: string; headers: Headers }) => { + sent.push(request); + return Promise.resolve({ + statusCode: nextResponse.statusCode, + headers: request.headers, + body: new Response(nextResponse.body).body, + }); + }, ), })), })); -import { buildWarehousesClient } from "../modular"; +import { ApiError } from "../errors"; +import { buildWarehousesClient, buildWorkspaceAuth } from "../modular"; /** Drive the wrapped httpClient with one request and return the UA it set. */ async function sentUserAgent( @@ -87,8 +106,6 @@ describe("modular mapToClientOptions (via buildWarehousesClient)", () => { ctorOpts.length = 0; patTokens.length = 0; m2mOpts.length = 0; - m2mMint.count = 0; - m2mMint.ttlMs = 60 * 60 * 1000; for (const key of AUTH_ENV) { originalEnv[key] = process.env[key]; delete process.env[key]; @@ -122,13 +139,16 @@ describe("modular mapToClientOptions (via buildWarehousesClient)", () => { buildWarehousesClient({ token: "abc", host: "https://x" }); expect(patTokens).toEqual(["abc"]); expect(ctorOpts[0].host).toBe("https://x"); - expect(ctorOpts[0].credentials).toEqual({ kind: "pat", token: "abc" }); + expect(ctorOpts[0].credentials).toMatchObject({ + kind: "pat", + token: "abc", + }); }); test("an empty-string token still uses PAT (no silent fall-through to default auth)", () => { buildWarehousesClient({ token: "", host: "https://x" }); expect(patTokens).toEqual([""]); - expect(ctorOpts[0].credentials).toEqual({ kind: "pat", token: "" }); + expect(ctorOpts[0].credentials).toMatchObject({ kind: "pat", token: "" }); }); test("a profile sets profileOptions and defers host to the SDK (ignores env)", () => { @@ -163,46 +183,21 @@ describe("modular mapToClientOptions (via buildWarehousesClient)", () => { }, ]); // Wrapped in the token cache, so assert the strategy, not identity. - expect((ctorOpts[0].credentials as { name(): string }).name()).toBe( + expect((ctorOpts[0].credentials as { name: () => string }).name()).toBe( "oauth-m2m", ); expect(patTokens).toEqual([]); }); - test("env M2M: caches the OAuth token until 40s before expiry, then refreshes", async () => { - // sdk-auth's newM2mCredentials mints a new OAuth token on every request; - // the legacy SDK reused it until expiry. - process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; - process.env.DATABRICKS_CLIENT_ID = "sp-client-id"; - process.env.DATABRICKS_CLIENT_SECRET = "sp-secret"; - buildWarehousesClient({}); - const creds = ctorOpts[0].credentials as { - authHeaders(): Promise>; - }; - const bearer = async () => - (await creds.authHeaders()).find((h) => h.key === "Authorization")?.value; - expect(await bearer()).toBe("Bearer m2m-token-1"); - expect(await bearer()).toBe("Bearer m2m-token-1"); - expect(m2mMint.count).toBe(1); - - // A token inside the 40s refresh margin is replaced on the next request. - m2mMint.ttlMs = 30_000; - vi.useFakeTimers({ toFake: ["Date"] }); - try { - vi.setSystemTime(Date.now() + 60 * 60 * 1000); - expect(await bearer()).toBe("Bearer m2m-token-2"); - expect(await bearer()).toBe("Bearer m2m-token-3"); - } finally { - vi.useRealTimers(); - } - }); - test("falls back to DATABRICKS_TOKEN (PAT) when no client id/secret is set", () => { process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; process.env.DATABRICKS_TOKEN = "env-pat"; buildWarehousesClient({}); expect(patTokens).toEqual(["env-pat"]); - expect(ctorOpts[0].credentials).toEqual({ kind: "pat", token: "env-pat" }); + expect(ctorOpts[0].credentials).toMatchObject({ + kind: "pat", + token: "env-pat", + }); expect(m2mOpts).toEqual([]); }); @@ -213,7 +208,7 @@ describe("modular mapToClientOptions (via buildWarehousesClient)", () => { process.env.DATABRICKS_CLIENT_SECRET = "sp-secret"; buildWarehousesClient({ token: "user-token", host: "https://x" }); expect(patTokens).toEqual(["user-token"]); - expect(ctorOpts[0].credentials).toEqual({ + expect(ctorOpts[0].credentials).toMatchObject({ kind: "pat", token: "user-token", }); @@ -260,3 +255,188 @@ describe("modular mapToClientOptions (via buildWarehousesClient)", () => { expect(ctorOpts[0].httpClient).toBeUndefined(); }); }); + +describe("buildWorkspaceAuth (auth + raw-request seam)", () => { + const AUTH_ENV = [ + "DATABRICKS_HOST", + "DATABRICKS_CLIENT_ID", + "DATABRICKS_CLIENT_SECRET", + "DATABRICKS_TOKEN", + ] as const; + const originalEnv: Record = {}; + const hour = 3_600_000; + + beforeEach(() => { + patTokens.length = 0; + m2mOpts.length = 0; + sent.length = 0; + nextResponse.statusCode = 200; + nextResponse.body = "{}"; + let n = 0; + m2mToken.mockReset().mockImplementation(async () => ({ + value: `m2m-${++n}`, + expiry: new Date(Date.now() + hour), + })); + // Never read the dev machine's ~/.databrickscfg. + resolveProfile.mockReset().mockResolvedValue({}); + defaultCreds.mockReset().mockReturnValue({ + authHeaders: async () => [{ key: "Authorization", value: "Bearer dflt" }], + }); + for (const key of AUTH_ENV) { + originalEnv[key] = process.env[key]; + delete process.env[key]; + } + }); + + afterEach(() => { + vi.useRealTimers(); + for (const key of AUTH_ENV) { + if (originalEnv[key] === undefined) delete process.env[key]; + else process.env[key] = originalEnv[key]; + } + }); + + async function authHeader(auth: ReturnType) { + const headers = new Headers(); + await auth.authenticate(headers); + return headers.get("Authorization"); + } + + test("an OBO token wins over env SP credentials — no escalation", async () => { + process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; + process.env.DATABRICKS_CLIENT_ID = "sp-client-id"; + process.env.DATABRICKS_CLIENT_SECRET = "sp-secret"; + const auth = buildWorkspaceAuth({ token: "user-token", host: "https://x" }); + expect(await authHeader(auth)).toBe("Bearer user-token"); + expect(await auth.getHost()).toBe("https://x"); + expect(m2mOpts).toEqual([]); + expect(defaultCreds).not.toHaveBeenCalled(); + }); + + test("an empty OBO token stays on PAT — never falls through to SP or the default chain", async () => { + process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; + process.env.DATABRICKS_CLIENT_ID = "sp-client-id"; + process.env.DATABRICKS_CLIENT_SECRET = "sp-secret"; + const auth = buildWorkspaceAuth({ token: "" }); + // `Headers` trims the trailing space of "Bearer ". + expect(await authHeader(auth)).toBe("Bearer"); + expect(patTokens).toEqual([""]); + expect(m2mOpts).toEqual([]); + expect(defaultCreds).not.toHaveBeenCalled(); + }); + + test("a profile resolves host + credentials through the SDK profile resolver", async () => { + process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; + resolveProfile.mockResolvedValue({ + host: "prof.cloud.databricks.com/", + token: "p", + }); + const auth = buildWorkspaceAuth({ profile: "myprofile" }); + // Profile host wins over env, scheme-normalized, trailing slash dropped. + expect(await auth.getHost()).toBe("https://prof.cloud.databricks.com"); + expect(await authHeader(auth)).toBe("Bearer dflt"); + expect(resolveProfile).toHaveBeenCalledWith({ profile: "myprofile" }); + expect(defaultCreds).toHaveBeenCalledWith({ + profile: { host: "https://prof.cloud.databricks.com", token: "p" }, + }); + // Resolved once and memoized. + await auth.getHost(); + expect(resolveProfile).toHaveBeenCalledTimes(1); + }); + + test("a failed resolution is retried instead of cached", async () => { + resolveProfile.mockRejectedValueOnce(new Error("bad cfg")); + const auth = buildWorkspaceAuth({ profile: "p", host: "https://x" }); + await expect(auth.getHost()).rejects.toThrow("bad cfg"); + expect(await auth.getHost()).toBe("https://x"); + }); + + test("no host anywhere → fails loudly", async () => { + await expect(buildWorkspaceAuth({}).getHost()).rejects.toThrow( + "Host is required.", + ); + }); + + test("env M2M: authenticates as the SP with the scheme-normalized host", async () => { + process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; + process.env.DATABRICKS_CLIENT_ID = "sp-client-id"; + process.env.DATABRICKS_CLIENT_SECRET = "sp-secret"; + const auth = buildWorkspaceAuth({}); + expect(await auth.getHost()).toBe("https://envhost.cloud.databricks.com"); + expect(await authHeader(auth)).toBe("Bearer m2m-1"); + expect(m2mOpts[0].host).toBe("https://envhost.cloud.databricks.com"); + expect(defaultCreds).not.toHaveBeenCalled(); + }); + + test("env M2M: caches the OAuth token until 40s before expiry, then refreshes", async () => { + // Regression: sdk-auth's newM2mCredentials mints a new token on EVERY + // call; the legacy SDK reused it until expiry. + vi.useFakeTimers(); + process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; + process.env.DATABRICKS_CLIENT_ID = "sp-client-id"; + process.env.DATABRICKS_CLIENT_SECRET = "sp-secret"; + const auth = buildWorkspaceAuth({}); + + // Concurrent first calls share one fetch. + const [a, b] = await Promise.all([authHeader(auth), authHeader(auth)]); + expect([a, b]).toEqual(["Bearer m2m-1", "Bearer m2m-1"]); + vi.advanceTimersByTime(hour - 41_000); + expect(await authHeader(auth)).toBe("Bearer m2m-1"); + expect(m2mToken).toHaveBeenCalledTimes(1); + + vi.advanceTimersByTime(2_000); // now inside the 40s refresh margin + expect(await authHeader(auth)).toBe("Bearer m2m-2"); + expect(m2mToken).toHaveBeenCalledTimes(2); + }); + + test("env M2M: a failed token fetch is not cached", async () => { + process.env.DATABRICKS_HOST = "envhost.cloud.databricks.com"; + process.env.DATABRICKS_CLIENT_ID = "sp-client-id"; + process.env.DATABRICKS_CLIENT_SECRET = "sp-secret"; + m2mToken.mockRejectedValueOnce(new Error("token endpoint down")); + const auth = buildWorkspaceAuth({}); + await expect(authHeader(auth)).rejects.toThrow("token endpoint down"); + expect(await authHeader(auth)).toBe("Bearer m2m-1"); + }); + + test("request: sends through the transport with the AppKit User-Agent, auth, and query", async () => { + const auth = buildWorkspaceAuth({ + host: "ws.cloud.databricks.com", + token: "t", + clientOptions: { + product: "@databricks/appkit", + productVersion: "0.64.0", + }, + } as never); + nextResponse.body = '{"warehouses":[]}'; + const res = await auth.request({ + method: "GET", + path: "/api/2.0/sql/warehouses", + query: { skip_cannot_use: "true" }, + headers: { "X-Extra": "1" }, + }); + expect(await res.json()).toEqual({ warehouses: [] }); + expect(sent[0].url).toBe( + "https://ws.cloud.databricks.com/api/2.0/sql/warehouses?skip_cannot_use=true", + ); + expect(sent[0].method).toBe("GET"); + expect(sent[0].headers.get("User-Agent")).toBe("@databricks/appkit/0.64.0"); + expect(sent[0].headers.get("Authorization")).toBe("Bearer t"); + expect(sent[0].headers.get("X-Extra")).toBe("1"); + }); + + test("request: a non-2xx status throws the wrapper ApiError with code + status", async () => { + nextResponse.statusCode = 403; + nextResponse.body = '{"error_code":"PERMISSION_DENIED","message":"nope"}'; + const auth = buildWorkspaceAuth({ host: "https://x", token: "t" }); + const error = await auth + .request({ method: "GET", path: "/api/x" }) + .catch((e: unknown) => e); + expect(error).toBeInstanceOf(ApiError); + expect(error).toMatchObject({ + message: "nope", + errorCode: "PERMISSION_DENIED", + statusCode: 403, + }); + }); +}); diff --git a/packages/shared/src/workspace-client/types.ts b/packages/shared/src/workspace-client/types.ts index 87071b004..b627afb23 100644 --- a/packages/shared/src/workspace-client/types.ts +++ b/packages/shared/src/workspace-client/types.ts @@ -14,7 +14,11 @@ * as each service migrates. */ import type { LegacyWorkspaceClient } from "./legacy"; -import type { StatementExecutionClient, WarehousesClient } from "./modular"; +import type { + StatementExecutionClient, + WarehousesClient, + WorkspaceAuth, +} from "./modular"; // Legacy SDK type namespaces for un-migrated services, re-exported so AppKit // modules import them from the wrapper rather than the SDK directly. `sql` @@ -36,7 +40,7 @@ export type * from "./modular"; * Accessors are legacy-typed for now (delegated to the underlying legacy SDK * client); see the module docblock. */ -export interface WorkspaceClient { +export interface WorkspaceClient extends WorkspaceAuth { /** UC Volumes / Files API. */ readonly files: LegacyWorkspaceClient["files"]; @@ -59,16 +63,16 @@ export interface WorkspaceClient { readonly currentUser: LegacyWorkspaceClient["currentUser"]; /** - * SDK `Config` — exposes `host` and `authenticate(headers)`. Used by the - * files-upload path and agents auth-header stamping, which bypass the typed - * services. + * Legacy SDK `Config`. Prefer `getHost()` / `authenticate(headers)` (modular, + * inherited from `WorkspaceAuth`); kept for structural `WorkspaceClientLike` + * callers (supervisor adapter) until they migrate. */ readonly config: LegacyWorkspaceClient["config"]; /** - * Low-level HTTP transport (`apiClient.request(...)`). Used for endpoints - * without a typed service method: SCIM header probe, warehouse listing, - * serving SSE streaming, vector search, internal telemetry. + * Legacy low-level HTTP transport. Prefer `request(...)` (modular, inherited + * from `WorkspaceAuth`); still used by serving SSE streaming, vector search, + * and the agents adapters, which take a structural `apiClient` shape. */ readonly apiClient: LegacyWorkspaceClient["apiClient"];