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
98 changes: 84 additions & 14 deletions packages/appkit/src/connectors/serving/client.ts
Original file line number Diff line number Diff line change
@@ -1,13 +1,19 @@
import { createLogger } from "../../logging/logger";
import type { serving, WorkspaceClient } from "../../workspace-client";
import type {
serving,
WorkspaceClient,
WorkspaceRequest,
} from "../../workspace-client";
import { contextFromAbortSignal } from "../context";

const logger = createLogger("connectors:serving");

/**
* Structural shape of a Databricks SDK client we need for the low-level
* `apiClient.request` call. Lets `streamPath` be reused by adapters that
* don't want a hard dependency on the concrete `WorkspaceClient` type.
* request call. Lets `streamPath` be reused by adapters that don't want a
* hard dependency on the concrete `WorkspaceClient` type. AppKit's own client
* provides `request` (modular transport); a caller-supplied legacy SDK client
* only has `apiClient.request`, which stays supported.
*/
export interface ApiClientLike {
apiClient: {
Expand All @@ -16,8 +22,29 @@ export interface ApiClientLike {
context?: unknown,
): Promise<unknown>;
};
request?(req: WorkspaceRequest): Promise<Response>;
}

// The legacy SDK's `servingEndpoints.query` copied only these fields into the
// request body and dropped the rest; kept so invocations send the same payload.
const QUERY_FIELDS = [
"client_request_id",
"dataframe_records",
"dataframe_split",
"extra_params",
"input",
"inputs",
"instances",
"max_tokens",
"messages",
"n",
"prompt",
"stop",
"stream",
"temperature",
"usage_context",
];

/**
* Transport shim shared by the agent adapters: given a request body, returns
* the raw SSE byte stream from a serving / AI-gateway endpoint. Injected at
Expand All @@ -31,8 +58,12 @@ export type StreamBody = (
) => Promise<ReadableStream<Uint8Array>>;

/**
* Invokes a serving endpoint using the SDK's high-level query API.
* Returns a typed QueryEndpointResponse.
* Invokes a serving endpoint. Returns the endpoint's JSON response as-is
* (model-specific: chat, completions, embeddings, custom), plus the
* `served-model-name` response header when present, like the legacy SDK's
* `servingEndpoints.query`. Sent raw via `client.request` because the modular
* serving SDK has no query method, and a generated unmarshal would strip
* model-specific fields.
*/
export async function invoke(
client: WorkspaceClient,
Expand All @@ -44,22 +75,44 @@ export async function invoke(

logger.debug("Invoking endpoint %s", endpointName);

return client.servingEndpoints.query({
name: endpointName,
...cleanBody,
} as serving.QueryEndpointInput);
const payload: Record<string, unknown> = {};
for (const key of QUERY_FIELDS) {
if (Object.hasOwn(cleanBody, key)) payload[key] = cleanBody[key];
}

const response = await client.request({
method: "POST",
path: `/serving-endpoints/${endpointName}/invocations`,
headers: {
Accept: "application/json",
"Content-Type": "application/json",
},
body: JSON.stringify(payload),
});

const text = await response.text();
let json: serving.QueryEndpointResponse;
try {
json = text.length === 0 ? {} : JSON.parse(text);
} catch {
// Same message (and typo) as the legacy SDK.
throw new Error(`Can't parse reponse as JSON: ${text}`);
}
const servedModelName = response.headers.get("served-model-name");
return servedModelName === null
? json
: { ...json, "served-model-name": servedModelName };
}

/**
* POSTs `body` as JSON to an arbitrary workspace API path and returns the raw
* SSE byte stream. No parsing is performed — bytes are passed through as-is.
*
* Uses the SDK's low-level `apiClient.request({ raw: true })` so callers
* inherit URL resolution, the SDK credential chain (PAT/OAuth/OIDC), and
* any future retries/telemetry baked into the SDK transport.
* Uses the client's `request` (modular transport) when available, else the
* legacy SDK's `apiClient.request({ raw: true })`, so callers inherit URL
* resolution and the SDK credential chain (PAT/OAuth/OIDC).
*
* When `signal` is provided it is bridged to the SDK's `Context` /
* `CancellationToken` so aborts cancel the outbound HTTP request.
* When `signal` is provided it aborts the outbound HTTP request.
*
* @internal
*
Expand All @@ -78,6 +131,23 @@ export async function streamPath(
): Promise<ReadableStream<Uint8Array>> {
logger.debug("Streaming from path %s", path);

if (client.request) {
const response = await client.request({
method: "POST",
path,
headers: {
"Content-Type": "application/json",
Accept: "text/event-stream",
},
body: JSON.stringify(body),
signal,
});
if (!response.body) {
throw new Error("Response body is null — streaming not supported");
}
return response.body;
}

const context = contextFromAbortSignal(signal);

const response = (await client.apiClient.request(
Expand Down
142 changes: 120 additions & 22 deletions packages/appkit/src/connectors/serving/tests/client.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,79 +6,177 @@ import { invoke, stream } from "../client";
function createMockClient(host = "https://test.databricks.com") {
return {
config: { host },
servingEndpoints: {
query: vi.fn(),
},
request: vi.fn(),
apiClient: {
request: vi.fn(),
},
} as any;
}

function createLegacyClient() {
return { apiClient: { request: vi.fn() } } as any;
}

function jsonResponse(body: unknown, headers?: Record<string, string>) {
return new Response(JSON.stringify(body), { headers });
}

function sentBody(client: any) {
return JSON.parse(client.request.mock.calls[0][0].body);
}

describe("Serving Connector", () => {
afterEach(() => {
vi.restoreAllMocks();
});

describe("invoke", () => {
test("calls servingEndpoints.query with endpoint name and body", async () => {
test("POSTs the body to the endpoint's invocations path", async () => {
const client = createMockClient();
const mockResponse = { choices: [{ message: { content: "Hello" } }] };
client.servingEndpoints.query.mockResolvedValue(mockResponse);
client.request.mockResolvedValue(jsonResponse(mockResponse));

const result = await invoke(client, "my-endpoint", {
messages: [{ role: "user", content: "Hi" }],
temperature: 0.7,
});

expect(client.servingEndpoints.query).toHaveBeenCalledWith({
name: "my-endpoint",
messages: [{ role: "user", content: "Hi" }],
temperature: 0.7,
expect(client.request).toHaveBeenCalledWith({
method: "POST",
path: "/serving-endpoints/my-endpoint/invocations",
headers: {
Accept: "application/json",
"Content-Type": "application/json",
},
body: JSON.stringify({
messages: [{ role: "user", content: "Hi" }],
temperature: 0.7,
}),
});
expect(result).toEqual(mockResponse);
});

test("strips stream property from body", async () => {
const client = createMockClient();
client.servingEndpoints.query.mockResolvedValue({});
client.request.mockResolvedValue(jsonResponse({}));

await invoke(client, "my-endpoint", {
messages: [],
stream: true,
temperature: 0.7,
});

const queryArg = client.servingEndpoints.query.mock.calls[0][0];
const queryArg = sentBody(client);
expect(queryArg.stream).toBeUndefined();
expect(queryArg.temperature).toBe(0.7);
});

test("returns typed QueryEndpointResponse", async () => {
// The legacy SDK's query() copied only its known request fields.
test("sends only the fields the legacy query sent", async () => {
const client = createMockClient();
client.request.mockResolvedValue(jsonResponse({}));

await invoke(client, "my-endpoint", {
messages: [],
max_tokens: 5,
top_p: 0.9,
});

expect(sentBody(client)).toEqual({ messages: [], max_tokens: 5 });
});

// Responses are model-specific: nothing may be stripped.
test("returns the raw JSON response, unknown fields included", async () => {
const client = createMockClient();
const responseData = {
choices: [{ message: { content: "Hello" } }],
usage: { prompt_tokens: 3, completion_tokens: 1 },
model: "test-model",
custom_field: { nested: [1, 2] },
};
client.servingEndpoints.query.mockResolvedValue(responseData);
client.request.mockResolvedValue(jsonResponse(responseData));

const result = await invoke(client, "my-endpoint", { messages: [] });
expect(result).toEqual(responseData);
});

test("propagates SDK errors", async () => {
test("merges the served-model-name header like the legacy query", async () => {
const client = createMockClient();
client.servingEndpoints.query.mockRejectedValue(
new Error("Endpoint not found"),
client.request.mockResolvedValue(
jsonResponse({ predictions: [1] }, { "served-model-name": "m-1" }),
);

const result = await invoke(client, "my-endpoint", { inputs: [1] });
expect(result).toEqual({ predictions: [1], "served-model-name": "m-1" });
});

test("returns {} for an empty body", async () => {
const client = createMockClient();
client.request.mockResolvedValue(new Response(""));

expect(await invoke(client, "my-endpoint", {})).toEqual({});
});

test("throws the legacy message for a non-JSON body", async () => {
const client = createMockClient();
client.request.mockResolvedValue(new Response("oops"));

await expect(invoke(client, "my-endpoint", {})).rejects.toThrow(
"Can't parse reponse as JSON: oops",
);
});

test("propagates SDK errors", async () => {
const client = createMockClient();
client.request.mockRejectedValue(new Error("Endpoint not found"));

await expect(
invoke(client, "my-endpoint", { messages: [] }),
).rejects.toThrow("Endpoint not found");
});
});

describe("stream", () => {
describe("stream via client.request", () => {
test("POSTs with stream: true and returns the response body", async () => {
const client = createMockClient();
const body = new ReadableStream<Uint8Array>();
client.request.mockResolvedValue(new Response(body));
const controller = new AbortController();

const result = await stream(
client,
"my endpoint",
{ messages: [], stream: false },
controller.signal,
);

expect(result).toBeInstanceOf(ReadableStream);
expect(client.request).toHaveBeenCalledWith({
method: "POST",
path: "/serving-endpoints/my%20endpoint/invocations",
headers: {
"Content-Type": "application/json",
Accept: "text/event-stream",
},
body: JSON.stringify({ messages: [], stream: true }),
signal: controller.signal,
});
expect(client.apiClient.request).not.toHaveBeenCalled();
});

test("throws when the response has no body", async () => {
const client = createMockClient();
client.request.mockResolvedValue(new Response(null));

await expect(
stream(client, "my-endpoint", { messages: [] }),
).rejects.toThrow("streaming not supported");
});
});

// Caller-supplied legacy SDK clients (agents' public `WorkspaceClientLike`)
// have no `request`, so they keep the `apiClient.request` path.
describe("stream via a legacy apiClient", () => {
test("returns a ReadableStream from apiClient.request", async () => {
const encoder = new TextEncoder();
const mockContents = new ReadableStream<Uint8Array>({
Expand All @@ -88,7 +186,7 @@ describe("Serving Connector", () => {
},
});

const client = createMockClient();
const client = createLegacyClient();
client.apiClient.request.mockResolvedValue({ contents: mockContents });

const result = await stream(client, "my-endpoint", { messages: [] });
Expand All @@ -97,7 +195,7 @@ describe("Serving Connector", () => {
});

test("sends stream: true in payload via apiClient.request", async () => {
const client = createMockClient();
const client = createLegacyClient();
client.apiClient.request.mockResolvedValue({
contents: new ReadableStream(),
});
Expand All @@ -116,7 +214,7 @@ describe("Serving Connector", () => {
});

test("passes SDK Context when AbortSignal is provided", async () => {
const client = createMockClient();
const client = createLegacyClient();
client.apiClient.request.mockResolvedValue({
contents: new ReadableStream(),
});
Expand All @@ -133,7 +231,7 @@ describe("Serving Connector", () => {
});

test("strips user-provided stream and re-injects", async () => {
const client = createMockClient();
const client = createLegacyClient();
client.apiClient.request.mockResolvedValue({
contents: new ReadableStream(),
});
Expand All @@ -148,7 +246,7 @@ describe("Serving Connector", () => {
});

test("throws when response has no contents", async () => {
const client = createMockClient();
const client = createLegacyClient();
client.apiClient.request.mockResolvedValue({ contents: null });

await expect(
Expand Down
Loading
Loading