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"];