@dbx-tools/appkit-model-gateway 0.9.58 → 0.9.61

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/src/registry.ts CHANGED
@@ -5,7 +5,13 @@
5
5
  */
6
6
 
7
7
  import { getExecutionContext, type WorkspaceClient } from "@databricks/appkit";
8
- import { metadata, modelCatalog, policy, resolve as modelResolve } from "@dbx-tools/model";
8
+ import {
9
+ metadata,
10
+ modelCatalog,
11
+ policy,
12
+ resolve as modelResolve,
13
+ type ListServingEndpointsOptions,
14
+ } from "@dbx-tools/model";
9
15
  import { log } from "@dbx-tools/shared-core";
10
16
  import {
11
17
  ModelClass,
@@ -36,31 +42,36 @@ export interface ModelRegistry {
36
42
  export interface ModelRegistryOptions {
37
43
  readonly ttlMs?: number;
38
44
  readonly overrides?: readonly ModelCapabilityOverride[];
45
+ /** Discover models with this client instead of the AppKit execution context. */
46
+ readonly client?: WorkspaceClient;
39
47
  }
40
48
 
41
49
  interface RegistryContext {
42
50
  readonly client: WorkspaceClient;
43
51
  readonly host: string;
44
- readonly identity: string;
52
+ readonly identity?: string;
45
53
  }
46
54
 
47
55
  /** Model registry backed by the active AppKit execution context. */
48
56
  export class DatabricksModelRegistry implements ModelRegistry {
49
57
  private readonly ttlMs: number;
50
58
  private readonly overrides: readonly ModelCapabilityOverride[];
59
+ private readonly client: WorkspaceClient | undefined;
51
60
 
52
61
  constructor(options: ModelRegistryOptions = {}) {
53
62
  this.ttlMs = positiveTtl(options.ttlMs);
54
63
  this.overrides = options.overrides ?? [];
64
+ this.client = options.client;
55
65
  }
56
66
 
57
67
  async list(): Promise<ModelTarget[]> {
58
- const context = await registryContext();
68
+ const context = await this.context();
59
69
  logger.debug("loading catalogue", { host: context.host, ttlMs: this.ttlMs });
60
- const endpoints = await modelCatalog.listServingEndpoints(context.client, context.host, {
61
- cacheIdentity: context.identity,
62
- ttlMs: this.ttlMs,
63
- });
70
+ const endpoints = await modelCatalog.listServingEndpoints(
71
+ context.client,
72
+ context.host,
73
+ catalogueOptions(context, this.ttlMs),
74
+ );
64
75
  const targets = endpoints.map((endpoint) => this.target(endpoint));
65
76
  logger.debug("loaded catalogue", {
66
77
  endpointCount: endpoints.length,
@@ -124,13 +135,16 @@ export class DatabricksModelRegistry implements ModelRegistry {
124
135
  }
125
136
 
126
137
  async refresh(): Promise<void> {
127
- const context = await registryContext();
138
+ const context = await this.context();
128
139
  logger.debug("clearing catalogue", { host: context.host });
129
- await modelCatalog.clearServingEndpointsCache(context.host, context.identity);
130
- await modelCatalog.listServingEndpoints(context.client, context.host, {
131
- cacheIdentity: context.identity,
132
- ttlMs: this.ttlMs,
133
- });
140
+ if (context.identity) {
141
+ await modelCatalog.clearServingEndpointsCache(context.host, context.identity);
142
+ }
143
+ await modelCatalog.listServingEndpoints(
144
+ context.client,
145
+ context.host,
146
+ catalogueOptions(context, this.ttlMs),
147
+ );
134
148
  logger.debug("refreshed catalogue", { host: context.host });
135
149
  }
136
150
 
@@ -182,6 +196,21 @@ export class DatabricksModelRegistry implements ModelRegistry {
182
196
  reasoningEfforts: endpoint.reasoningEfforts ?? [],
183
197
  };
184
198
  }
199
+
200
+ private async context(): Promise<RegistryContext> {
201
+ if (this.client) {
202
+ const host = (await this.client.config.getHost()).toString();
203
+ return { client: this.client, host };
204
+ }
205
+ return registryContext();
206
+ }
207
+ }
208
+
209
+ function catalogueOptions(context: RegistryContext, ttlMs: number): ListServingEndpointsOptions {
210
+ return {
211
+ ttlMs,
212
+ ...(context.identity ? { cacheIdentity: context.identity } : {}),
213
+ };
185
214
  }
186
215
 
187
216
  async function registryContext(): Promise<RegistryContext> {
package/src/router.ts CHANGED
@@ -4,6 +4,7 @@
4
4
  * @module
5
5
  */
6
6
 
7
+ import { object } from "@dbx-tools/shared-core";
7
8
  import type {
8
9
  GatewayRoute,
9
10
  ModelTarget,
@@ -35,21 +36,16 @@ export function resolveRoute(input: ResolveRouteInput): GatewayRoute {
35
36
  rejectStatefulResponses(features);
36
37
  if (
37
38
  isCodexOriginator(input.originator) &&
38
- targetsModelService(requestedModel, target) &&
39
- target.capabilities.aiGatewayCodex &&
40
- supportsDirectResponses(target, features)
39
+ canUseAiGatewayCodex(target, features) &&
40
+ !requestedServingEndpoint(requestedModel, target)
41
41
  ) {
42
42
  return directRoute(input, "databricks-ai-gateway-codex", target.modelServiceName!);
43
43
  }
44
44
  if (target.capabilities.responses && supportsDirectResponses(target, features)) {
45
45
  return directRoute(input, "databricks-responses", target.id);
46
46
  }
47
- if (
48
- target.capabilities.aiGatewayCodex &&
49
- target.modelServiceName &&
50
- supportsDirectResponses(target, features)
51
- ) {
52
- return directRoute(input, "databricks-ai-gateway-codex", target.modelServiceName);
47
+ if (canUseAiGatewayCodex(target, features)) {
48
+ return directRoute(input, "databricks-ai-gateway-codex", target.modelServiceName!);
53
49
  }
54
50
  if (target.capabilities.openResponses && supportsOpenResponses(target, features)) {
55
51
  return directRoute(input, "databricks-open-responses", target.id);
@@ -74,17 +70,17 @@ export function resolveRoute(input: ResolveRouteInput): GatewayRoute {
74
70
  export function requestedFeatures(body: Readonly<Record<string, unknown>>): RequestedFeatures {
75
71
  const tools = Array.isArray(body.tools) ? body.tools : [];
76
72
  const toolTypes = tools
77
- .filter((tool): tool is Record<string, unknown> => isRecord(tool))
73
+ .filter((tool): tool is Record<string, unknown> => object.isRecord(tool))
78
74
  .map((tool) => tool.type);
79
- const text = isRecord(body.text) ? body.text : {};
80
- const format = isRecord(text.format) ? text.format.type : text.format;
75
+ const text = object.isRecord(body.text) ? body.text : {};
76
+ const format = object.isRecord(text.format) ? text.format.type : text.format;
81
77
  return {
82
78
  background: body.background === true,
83
79
  customTools: toolTypes.some((type) => type === "custom"),
84
80
  parallelTools: body.parallel_tool_calls === true,
85
81
  previousResponse:
86
82
  typeof body.previous_response_id === "string" && body.previous_response_id.length > 0,
87
- reasoning: isRecord(body.reasoning) || body.reasoning_effort !== undefined,
83
+ reasoning: object.isRecord(body.reasoning) || body.reasoning_effort !== undefined,
88
84
  storage: body.store === true,
89
85
  structuredOutput: format !== undefined && format !== "text",
90
86
  tools: tools.length > 0,
@@ -163,12 +159,14 @@ function rejectStatefulResponses(features: RequestedFeatures): void {
163
159
  if (unsupported.length > 0) throw new UnsupportedGatewayFeatureError(unsupported);
164
160
  }
165
161
 
166
- function targetsModelService(requestedModel: string, target: ModelTarget): boolean {
167
- if (!target.modelServiceName) return false;
168
- const unqualified = requestedModel.trim().replace(/^(?:dbx|databricks)\//i, "");
169
- return unqualified === target.modelServiceName;
162
+ function canUseAiGatewayCodex(target: ModelTarget, features: RequestedFeatures): boolean {
163
+ return (
164
+ Boolean(target.capabilities.aiGatewayCodex && target.modelServiceName) &&
165
+ supportsDirectResponses(target, features)
166
+ );
170
167
  }
171
168
 
172
- function isRecord(value: unknown): value is Record<string, unknown> {
173
- return value !== null && typeof value === "object" && !Array.isArray(value);
169
+ /** Serving-endpoint ids stay on Databricks Responses; aliases and model-service slugs use Codex. */
170
+ function requestedServingEndpoint(requestedModel: string, target: ModelTarget): boolean {
171
+ return requestedModel.trim().replace(/^(?:dbx|databricks)\//i, "") === target.id;
174
172
  }
package/src/transport.ts CHANGED
@@ -7,6 +7,7 @@
7
7
  import { getExecutionContext } from "@databricks/appkit";
8
8
  import { workspaceClient } from "@dbx-tools/databricks";
9
9
  import { log } from "@dbx-tools/shared-core";
10
+ import { openaiChat } from "@dbx-tools/shared-model";
10
11
  import { upstreamUrl, type GatewayRoute } from "@dbx-tools/shared-model-gateway";
11
12
 
12
13
  const logger = log.logger("appkit/model-gateway/transport");
@@ -60,7 +61,7 @@ export async function fetchDatabricks(request: DatabricksTransportRequest): Prom
60
61
  await client.config.authenticate(headers);
61
62
  const url = upstreamUrl(host, request.route);
62
63
  const startedAt = Date.now();
63
- const body = JSON.stringify({ ...request.body, model: request.route.upstreamModel });
64
+ const body = JSON.stringify(directDatabricksRequestBody(request.route, request.body));
64
65
  logger.debug("sending upstream request", {
65
66
  model: request.route.upstreamModel,
66
67
  path: new URL(url).pathname,
@@ -119,6 +120,20 @@ export async function fetchDatabricks(request: DatabricksTransportRequest): Prom
119
120
  }
120
121
  }
121
122
 
123
+ /** Remove client compatibility fields rejected by direct Databricks protocols. */
124
+ export function directDatabricksRequestBody(
125
+ route: Pick<GatewayRoute, "upstreamProtocol" | "upstreamModel">,
126
+ body: Readonly<Record<string, unknown>>,
127
+ ): Record<string, unknown> {
128
+ const sanitized: Record<string, unknown> = { ...body, model: route.upstreamModel };
129
+ if (route.upstreamProtocol === "databricks-chat") {
130
+ openaiChat.stripUnsupportedChatFields(sanitized);
131
+ } else {
132
+ delete sanitized.parallel_tool_calls;
133
+ }
134
+ return sanitized;
135
+ }
136
+
122
137
  /** Copy only response headers that belong to the public model protocol. */
123
138
  export function gatewayResponseHeaders(upstream: Headers): Headers {
124
139
  const headers = new Headers();