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
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ describe("createMockWorkspaceClient — build the fake client yourself", () => {
});
expect(getMock(client, "jobs.getRun")).toHaveBeenCalledWith({ run_id: 1 });
await expect(
client.genie.getMessage({ id: "m-1" }),
client.genie.genieGetConversationMessage({ messageId: "m-1" }),
).resolves.toBeUndefined();
});
});
Expand Down
4 changes: 2 additions & 2 deletions docs/docs/api/appkit/Interface.WorkspaceClient.md

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

4 changes: 2 additions & 2 deletions docs/docs/plugins/testing.md
Original file line number Diff line number Diff line change
Expand Up @@ -343,7 +343,7 @@ await mock.attach(plugin);
```

`options` is:
- `responses` — seed the mock workspace client with responses keyed by dotted path (`"jobs.getRun"`, `"genie.getMessage"`). A value can be static or a function of call arguments and the abort signal.
- `responses` — seed the mock workspace client with responses keyed by dotted path (`"jobs.getRun"`, `"genie.genieGetConversationMessage"`). A value can be static or a function of call arguments and the abort signal.
- `env` — set environment variables scoped to the test; they are restored on plugin detach.
- `strict` — throw if a handler calls an undeclared workspace-client path (instead of silently resolving `undefined`). The built-in defaults still count as declared.

Expand Down Expand Up @@ -513,7 +513,7 @@ const client = createMockWorkspaceClient({
});

await client.jobs.getRun({ run_id: 1 }); // → { state: "TERMINATED" }
await client.genie.getMessage({ id: "m-1" }); // → undefined, does not throw
await client.genie.genieGetConversationMessage({ messageId: "m-1" }); // → undefined, does not throw
```

`createTestApp` installs one of these for you, so reach for it directly only when you're driving a plugin through `createTestPluginContext` or `mockServiceContext`.
Expand Down
1 change: 1 addition & 0 deletions knip.json
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
"@databricks/sdk-auth",
"@databricks/sdk-core",
"@databricks/sdk-experimental",
"@databricks/sdk-genie",
"@databricks/sdk-options",
"@databricks/sdk-scim",
"@databricks/sdk-statementexecution",
Expand Down
1 change: 1 addition & 0 deletions packages/appkit/package.json
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,7 @@
"@databricks/sdk-auth": "0.51.0",
"@databricks/sdk-core": "0.51.0",
"@databricks/sdk-experimental": "0.17.0",
"@databricks/sdk-genie": "0.54.0",
"@databricks/sdk-options": "0.51.0",
"@databricks/sdk-scim": "0.51.0",
"@databricks/sdk-statementexecution": "0.52.0",
Expand Down
212 changes: 143 additions & 69 deletions packages/appkit/src/connectors/genie/client.ts
Original file line number Diff line number Diff line change
@@ -1,13 +1,7 @@
import { createLogger } from "../../logging";
import {
type GenieMessage,
Time,
TimeUnits,
type Waiter,
type WorkspaceClient,
} from "../../workspace-client";
import type { GenieMessage, WorkspaceClient } from "../../workspace-client";
import { genieConnectorDefaults } from "./defaults";
import { pollWaiter } from "./poll-waiter";
import { type Pollable, pollWaiter } from "./poll-waiter";
import type {
GenieAttachmentResponse,
GenieConversationHistoryResponse,
Expand All @@ -26,7 +20,11 @@ const GenieErrors = {
QUERY_RESULT_FAILED: "Failed to fetch query result",
} as const;

type CreateMessageWaiter = Waiter<GenieMessage, GenieMessage>;
type CreateMessageWaiter = Pollable<GenieMessage>;

// Legacy SDK waiter defaults, kept so polling cadence is unchanged.
const DEFAULT_WAIT_TIMEOUT_MS = 10 * 60_000;
const MAX_POLL_INTERVAL_MS = 10_000;

interface GenieConnectorConfig {
timeout?: number;
Expand All @@ -35,38 +33,72 @@ interface GenieConnectorConfig {

function mapAttachments(message: GenieMessage): GenieAttachmentResponse[] {
return (
message.attachments?.map((att) => ({
attachmentId: att.attachment_id,
query: att.query
? {
title: att.query.title,
description: att.query.description,
query: att.query.query,
statementId: att.query.statement_id,
}
: undefined,
text: att.text ? { content: att.text.content } : undefined,
suggestedQuestions: att.suggested_questions?.questions,
message.attachments?.map(({ attachmentId, attachment }) => ({
attachmentId,
query:
attachment?.$case === "query"
? {
title: attachment.query.title,
description: attachment.query.description,
query: attachment.query.query,
statementId: attachment.query.statementId,
}
: undefined,
text:
attachment?.$case === "text"
? { content: attachment.text.content }
: undefined,
suggestedQuestions:
attachment?.$case === "suggestedQuestions"
? attachment.suggestedQuestions.questions
: undefined,
})) ?? []
);
}

function toMessageResponse(message: GenieMessage): GenieMessageResponse {
return {
messageId: message.message_id,
conversationId: message.conversation_id,
spaceId: message.space_id,
messageId: message.messageId ?? "",
conversationId: message.conversationId ?? "",
spaceId: message.spaceId ?? "",
status: message.status ?? "COMPLETED",
content: message.content,
content: message.content ?? "",
attachments: mapAttachments(message),
error: message.error?.error,
};
}

/**
* The modular SDK returns the statement response camelCased with int64 fields as
* `bigint`. The SSE contract (`GenieStatementResponse`, read by appkit-ui) is the
* raw snake_case API shape, and `JSON.stringify` throws on `bigint`, so convert
* back here. Recurses into objects only; `data_array` rows pass through as-is.
*/
function toWireShape(value: unknown): unknown {
if (typeof value === "bigint") return Number(value);
if (Array.isArray(value)) return value.map(toWireShape);
if (value === null || typeof value !== "object") return value;
return Object.fromEntries(
Object.entries(value).map(([key, v]) => [
key.replace(/[A-Z]/g, (c) => `_${c.toLowerCase()}`),
toWireShape(v),
]),
);
}

function classifyGenieError(error: unknown): string {
const message = error instanceof Error ? error.message : String(error);
// Modular ApiError carries the code on `.code`, legacy on `.errorCode`.
const { code, errorCode } = (error ?? {}) as {
code?: unknown;
errorCode?: unknown;
};

if (message.includes("RESOURCE_DOES_NOT_EXIST")) {
if (
code === "RESOURCE_DOES_NOT_EXIST" ||
errorCode === "RESOURCE_DOES_NOT_EXIST" ||
message.includes("RESOURCE_DOES_NOT_EXIST")
) {
return GenieErrors.SPACE_ACCESS_DENIED;
}

Expand Down Expand Up @@ -100,26 +132,74 @@ export class GenieConnector {
conversationId: string;
messageId: string;
}> {
if (conversationId) {
const waiter = await workspaceClient.genie.createMessage({
space_id: spaceId,
conversation_id: conversationId,
content,
});
return {
messageWaiter: waiter,
conversationId,
messageId: waiter.message_id ?? "",
};
}
const start = await workspaceClient.genie.startConversation({
space_id: spaceId,
content,
});
const started = conversationId
? await workspaceClient.genie.genieCreateConversationMessage({
spaceId,
conversationId,
content,
})
: await workspaceClient.genie.genieStartConversation({
spaceId,
content,
});
return {
messageWaiter: start as unknown as CreateMessageWaiter,
conversationId: start.conversation_id,
messageId: start.message_id,
messageWaiter: this.messagePoller(
workspaceClient,
spaceId,
started.conversationId,
started.messageId,
),
conversationId: started.conversationId,
messageId: started.messageId,
};
}

/**
* Polls `getConversationMessage` until COMPLETED. Replaces the SDK waiter: the
* modular `wait()` has no `onProgress`, which the SSE `status` events need. It
* mirrors the legacy waiter: progress on every poll, backoff of `attempt`
* seconds + 50-750ms jitter capped at 10s, a 10 minute default timeout, and the
* same `failed to reach COMPLETED state` errors `classifyGenieError` matches.
*/
private messagePoller(
workspaceClient: WorkspaceClient,
spaceId: string,
conversationId: string,
messageId: string,
): CreateMessageWaiter {
return {
async wait(options) {
const timeout =
typeof options?.timeout === "number"
? options.timeout
: DEFAULT_WAIT_TIMEOUT_MS;
const deadline = Date.now() + timeout;
let lastStatus: string | undefined;
for (let attempt = 1; Date.now() < deadline; attempt++) {
const message =
await workspaceClient.genie.genieGetConversationMessage({
spaceId,
conversationId,
messageId,
});
await options?.onProgress?.(message);
lastStatus = message.status;
if (lastStatus === "COMPLETED") return message;
if (lastStatus === "FAILED") {
throw new Error("failed to reach COMPLETED state, got FAILED");
}
const jitter = 50 + Math.random() * 700;
await new Promise((resolve) =>
setTimeout(
resolve,
Math.min(attempt * 1000 + jitter, MAX_POLL_INTERVAL_MS),
),
);
}
throw new Error(
`timed out: failed to reach COMPLETED state, got ${lastStatus}`,
);
},
};
}

Expand All @@ -128,9 +208,7 @@ export class GenieConnector {
options?: { timeout?: number },
): Promise<GenieMessage> {
const timeout = options?.timeout ?? this.config.timeout;
const waitOptions =
timeout > 0 ? { timeout: new Time(timeout, TimeUnits.milliseconds) } : {};
return messageWaiter.wait(waitOptions);
return messageWaiter.wait(timeout > 0 ? { timeout } : {});
}

async listConversationMessages(
Expand All @@ -145,18 +223,18 @@ export class GenieConnector {
const pageSize =
options?.pageSize ?? genieConnectorDefaults.initialPageSize;

const response = await workspaceClient.genie.listConversationMessages({
space_id: spaceId,
conversation_id: conversationId,
page_size: pageSize,
...(options?.pageToken ? { page_token: options.pageToken } : {}),
const response = await workspaceClient.genie.genieListConversationMessages({
spaceId,
conversationId,
pageSize,
...(options?.pageToken ? { pageToken: options.pageToken } : {}),
});

const messages = (response.messages ?? []).reverse().map(toMessageResponse);

return {
messages,
nextPageToken: response.next_page_token ?? null,
nextPageToken: response.nextPageToken ?? null,
};
}

Expand All @@ -169,13 +247,13 @@ export class GenieConnector {
_signal?: AbortSignal,
): Promise<GenieStatementResponse> {
const response =
await workspaceClient.genie.getMessageAttachmentQueryResult({
space_id: spaceId,
conversation_id: conversationId,
message_id: messageId,
attachment_id: attachmentId,
await workspaceClient.genie.genieGetMessageAttachmentQueryResult({
spaceId,
conversationId,
messageId,
attachmentId,
});
return response.statement_response as GenieStatementResponse;
return toWireShape(response.statementResponse) as GenieStatementResponse;
}

async *streamSendMessage(
Expand Down Expand Up @@ -206,10 +284,7 @@ export class GenieConnector {

const timeout =
options?.timeout != null ? options.timeout : this.config.timeout;
const waitOptions =
timeout > 0
? { timeout: new Time(timeout, TimeUnits.milliseconds) }
: {};
const waitOptions = timeout > 0 ? { timeout } : {};

let completedMessage!: GenieMessage;
for await (const event of pollWaiter(messageWaiter, waitOptions)) {
Expand Down Expand Up @@ -402,11 +477,10 @@ export class GenieConnector {
while (true) {
if (signal?.aborted) return;

const message = await workspaceClient.genie.getMessage({
space_id: spaceId,
conversation_id: conversationId,
message_id: messageId,
});
const message = await workspaceClient.genie.genieGetConversationMessage(
{ spaceId, conversationId, messageId },
{ signal },
);

if (message.status && message.status !== lastStatus) {
lastStatus = message.status;
Expand Down
Loading
Loading