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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 24 additions & 36 deletions packages/appkit/src/connectors/files/client.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<string, string> = {
"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,
[],
);
Expand Down
113 changes: 35 additions & 78 deletions packages/appkit/src/connectors/files/tests/client.test.ts
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import { createMockTelemetry } from "@tools/test-helpers";
import { afterEach, beforeEach, describe, expect, test, vi } from "vitest";

Check warning on line 2 in packages/appkit/src/connectors/files/tests/client.test.ts

View workflow job for this annotation

GitHub Actions / Lint & Type Check

eslint(no-unused-vars)

packages/appkit/src/connectors/files/tests/client.test.ts:2:10: Identifier 'afterEach' is imported but never used.

import { createApiError } from "../../../testing";
import type { WorkspaceClient } from "../../../workspace-client";
Expand All @@ -7,7 +7,7 @@
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(),
Expand All @@ -17,21 +17,13 @@
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) => {
Expand Down Expand Up @@ -459,32 +451,29 @@

describe("upload()", () => {
let connector: FilesConnector;
let fetchSpy: ReturnType<typeof vi.fn>;

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",
},
}),
);
});
Expand All @@ -493,11 +482,11 @@
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" }),
}),
);
});
Expand All @@ -506,79 +495,47 @@
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 () => {
await connector.upload(mockClient, "file.txt", "data", {
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",
);
});
});

Expand Down
6 changes: 2 additions & 4 deletions packages/appkit/src/connectors/mlflow/auth.ts
Original file line number Diff line number Diff line change
Expand Up @@ -49,11 +49,9 @@ async function resolveViaSdk(
// Mints the OAuth access token (or reuses a PAT from the profile) and adds
// an `Authorization: Bearer <token>` 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 {
Expand Down
38 changes: 38 additions & 0 deletions packages/appkit/src/connectors/mlflow/tests/auth.test.ts
Original file line number Diff line number Diff line change
@@ -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();
});
});
14 changes: 6 additions & 8 deletions packages/appkit/src/context/service-context.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}

/**
Expand Down
11 changes: 9 additions & 2 deletions packages/appkit/src/context/tests/service-context.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -411,7 +419,6 @@ describe("ServiceContext", () => {
expect.objectContaining({
path: "/api/2.0/preview/scim/v2/Me",
method: "GET",
responseHeaders: ["x-databricks-org-id"],
}),
);
});
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 };
});
Expand Down
Loading
Loading