diff --git a/src/lib/adapters/http/auth-config.test.ts b/src/lib/adapters/http/auth-config.test.ts index 7040028ff2c..95d3c1ddcc0 100644 --- a/src/lib/adapters/http/auth-config.test.ts +++ b/src/lib/adapters/http/auth-config.test.ts @@ -11,6 +11,7 @@ import { createOpenAiLikeAuthConfig, createQueryParamAuthConfig, createXApiKeyAuthConfig, + parseOpenAiLikeExtraHeaders, } from "./auth-config"; describe("curl auth config helper", () => { @@ -96,6 +97,24 @@ describe("curl auth config helper", () => { } }); + it.each([ + " : value", + "Bad Header: value", + "missing-colon", + ])("rejects invalid OpenAI-like provider header %j", (header) => { + expect(() => parseOpenAiLikeExtraHeaders([header])).toThrow( + "invalid OpenAI-like provider header", + ); + }); + + it("accepts every HTTP token character in an OpenAI-like provider header name", () => { + const tokenChars = + "!#$%&'*+-.^_`|~0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"; + expect(parseOpenAiLikeExtraHeaders([`${tokenChars}: value`])).toEqual([ + { name: tokenChars, value: "value" }, + ]); + }); + it("honours a caller-supplied tmpfile prefix so health probes are identifiable in /proc", () => { const config = createBearerAuthConfig("nvapi-test", { prefix: "nemoclaw-kimi-health-curl" }); try { diff --git a/src/lib/adapters/http/auth-config.ts b/src/lib/adapters/http/auth-config.ts index 6dbf473de99..5ed1f4e3f55 100644 --- a/src/lib/adapters/http/auth-config.ts +++ b/src/lib/adapters/http/auth-config.ts @@ -7,6 +7,7 @@ import path from "node:path"; const CURL_AUTH_CONFIG_PREFIX = "nemoclaw-curl-auth"; const CURL_AUTH_CONFIG_NAME_PATTERN = /^[a-zA-Z][a-zA-Z0-9._-]{0,127}$/; +const HTTP_HEADER_NAME_PATTERN = /^[!#$%&'*+\-.^_`|~0-9A-Za-z]+$/; function resolveCurlAuthConfigPrefix(prefix: string | undefined): string { if (prefix === undefined) return CURL_AUTH_CONFIG_PREFIX; @@ -133,6 +134,28 @@ export interface CreateOpenAiLikeAuthConfigOptions extends CreateCurlAuthConfigO extraHeaders?: readonly string[]; } +export interface OpenAiLikeExtraHeader { + name: string; + value: string; +} + +export function parseOpenAiLikeExtraHeaders( + extraHeaders: readonly string[] = [], +): OpenAiLikeExtraHeader[] { + return extraHeaders.map((header) => { + const sanitized = header.replace(/[\r\n]+/g, " "); + const separator = sanitized.indexOf(":"); + const name = separator < 0 ? "" : sanitized.slice(0, separator).trim(); + if (!HTTP_HEADER_NAME_PATTERN.test(name)) { + throw new Error("invalid OpenAI-like provider header"); + } + return { + name, + value: sanitized.slice(separator + 1).trim(), + }; + }); +} + export function createOpenAiLikeAuthConfig( apiKey: string, authMode?: OpenAiLikeAuthMode, @@ -144,8 +167,8 @@ export function createOpenAiLikeAuthConfig( } else if (apiKey) { entries.push({ kind: "header", value: `Authorization: Bearer ${apiKey}` }); } - for (const header of options.extraHeaders ?? []) { - entries.push({ kind: "header", value: header }); + for (const { name, value } of parseOpenAiLikeExtraHeaders(options.extraHeaders)) { + entries.push({ kind: "header", value: `${name}: ${value}` }); } return createCurlAuthConfig(entries, options); } diff --git a/src/lib/adapters/http/validation-session.test.ts b/src/lib/adapters/http/validation-session.test.ts new file mode 100644 index 00000000000..74a56f7cd21 --- /dev/null +++ b/src/lib/adapters/http/validation-session.test.ts @@ -0,0 +1,313 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import http from "node:http"; +import { afterEach, describe, expect, it, vi } from "vitest"; +import { + buildLookup, + createValidationSession, + getValidationSessionIneligibility, +} from "./validation-session"; + +const servers: http.Server[] = []; + +afterEach(async () => { + await Promise.all( + servers.splice(0).map( + (server) => + new Promise((resolve) => { + server.close(() => resolve()); + server.closeAllConnections(); + }), + ), + ); +}); + +async function listen(server: http.Server): Promise { + servers.push(server); + await new Promise((resolve) => server.listen(0, "127.0.0.1", resolve)); + const address = server.address(); + expect(address).toBeTruthy(); + expect(typeof address).toBe("object"); + return (address as import("node:net").AddressInfo).port; +} + +describe("provider validation session", () => { + it("resolves once and reuses one TCP connection for sequential requests", async () => { + let connections = 0; + const requests: string[] = []; + const server = http.createServer((request, response) => { + requests.push(request.url ?? ""); + request.resume(); + response.setHeader("content-type", "application/json"); + response.end('{"ok":true}'); + }); + server.on("connection", () => { + connections += 1; + }); + const port = await listen(server); + const lookup = vi.fn(async () => [{ address: "127.0.0.1", family: 4 }]); + const sockets: import("node:net").Socket[] = []; + const session = await createValidationSession(`http://provider.example.test:${port}/v1`, { + env: {}, + lookup, + onSocket: (socket) => sockets.push(socket), + allowPrivateAddressesForTesting: true, + }); + + expect(session).not.toBeNull(); + const first = await session!.request({ + url: `http://provider.example.test:${port}/v1/responses`, + body: "{}", + timeoutMs: 1_000, + }); + const second = await session!.request({ + url: `http://provider.example.test:${port}/v1/chat/completions`, + body: "{}", + timeoutMs: 1_000, + }); + session!.close(); + + expect(first.ok).toBe(true); + expect(second.ok).toBe(true); + expect(lookup).toHaveBeenCalledTimes(1); + expect(connections).toBe(1); + expect(sockets).toHaveLength(1); + expect(sockets[0].destroyed).toBe(true); + expect(requests).toEqual(["/v1/responses", "/v1/chat/completions"]); + }); + + it("uses preflight-pinned addresses without another DNS lookup", async () => { + const server = http.createServer((request, response) => { + request.resume(); + response.end('{"ok":true}'); + }); + const port = await listen(server); + const lookup = vi.fn(); + const session = await createValidationSession(`http://provider.example.test:${port}/v1`, { + env: {}, + lookup, + pinnedAddresses: ["127.0.0.1"], + allowPrivateAddressesForTesting: true, + }); + + await expect( + session!.request({ + url: `http://provider.example.test:${port}/v1/responses`, + body: "{}", + timeoutMs: 1_000, + }), + ).resolves.toMatchObject({ ok: true }); + expect(lookup).not.toHaveBeenCalled(); + session!.close(); + }); + + it("returns ENOTFOUND when IPv6 is requested from IPv4-only pinned addresses", async () => { + const lookup = buildLookup([{ address: "93.184.216.34", family: 4 }]); + const result = new Promise((resolve, reject) => { + lookup("provider.example.test", { all: true, family: 6 }, (error, addresses) => + error ? reject(error) : resolve(addresses), + ); + }); + + await expect(result).rejects.toMatchObject({ code: "ENOTFOUND" }); + }); + + it("falls back when pre-resolution fails", async () => { + const lookup = vi.fn(async () => { + throw Object.assign(new Error("temporary DNS failure"), { code: "EAI_AGAIN" }); + }); + + await expect( + createValidationSession("https://provider.example.test/v1", { env: {}, lookup }), + ).resolves.toBeNull(); + expect(lookup).toHaveBeenCalledTimes(1); + }); + + it("falls back when DNS pre-resolution exceeds its deadline", async () => { + const lookup = vi.fn(() => new Promise>(() => {})); + + await expect( + createValidationSession("https://provider.example.test/v1", { + env: {}, + lookup, + dnsTimeoutMs: 10, + }), + ).resolves.toBeNull(); + expect(lookup).toHaveBeenCalledTimes(1); + }); + + it.each([ + "127.0.0.1", + "10.0.0.1", + ])("falls back when DNS resolves to private address %s", async (address) => { + await expect( + createValidationSession("https://provider.example.test/v1", { + env: {}, + lookup: async () => [{ address, family: 4 }], + }), + ).resolves.toBeNull(); + }); + + it("falls back before DNS pre-resolution when a proxy applies", async () => { + const lookup = vi.fn(); + + await expect( + createValidationSession("https://provider.example.test/v1", { + env: { HTTPS_PROXY: "http://proxy.example.test:8080" }, + lookup, + }), + ).resolves.toBeNull(); + expect(lookup).not.toHaveBeenCalled(); + }); + + it("enforces a total deadline while a response trickles data", async () => { + const server = http.createServer((request, response) => { + request.resume(); + response.writeHead(200, { "content-type": "application/json" }); + const timer = setInterval(() => response.write(" "), 5); + response.on("close", () => clearInterval(timer)); + }); + const port = await listen(server); + const session = await createValidationSession(`http://provider.example.test:${port}/v1`, { + env: {}, + lookup: async () => [{ address: "127.0.0.1", family: 4 }], + allowPrivateAddressesForTesting: true, + }); + + await expect( + session!.request({ + url: `http://provider.example.test:${port}/v1/responses`, + body: "{}", + timeoutMs: 30, + }), + ).resolves.toMatchObject({ ok: false, curlStatus: 28 }); + session!.close(); + }); + + it("rejects a response larger than 8 MiB", async () => { + const server = http.createServer((request, response) => { + request.resume(); + response.end(Buffer.alloc(8 * 1024 * 1024 + 1)); + }); + const port = await listen(server); + const session = await createValidationSession(`http://provider.example.test:${port}/v1`, { + env: {}, + lookup: async () => [{ address: "127.0.0.1", family: 4 }], + allowPrivateAddressesForTesting: true, + }); + + await expect( + session!.request({ + url: `http://provider.example.test:${port}/v1/responses`, + body: "{}", + timeoutMs: 1_000, + }), + ).resolves.toMatchObject({ + ok: false, + stderr: "validation response exceeded 8 MiB", + }); + session!.close(); + }); + + it("reconnects without another DNS lookup when the server closes keepalive", async () => { + let connections = 0; + const server = http.createServer((request, response) => { + request.resume(); + response.setHeader("connection", "close"); + response.end('{"ok":true}'); + }); + server.on("connection", () => { + connections += 1; + }); + const port = await listen(server); + const lookup = vi.fn(async () => [{ address: "127.0.0.1", family: 4 }]); + const session = await createValidationSession(`http://provider.example.test:${port}/v1`, { + env: {}, + lookup, + allowPrivateAddressesForTesting: true, + }); + + for (const path of ["responses", "chat/completions"]) { + await expect( + session!.request({ + url: `http://provider.example.test:${port}/v1/${path}`, + body: "{}", + timeoutMs: 1_000, + }), + ).resolves.toMatchObject({ ok: true }); + } + session!.close(); + + expect(lookup).toHaveBeenCalledTimes(1); + expect(connections).toBe(2); + }); + + it("refuses to send a session request to another origin", async () => { + const server = http.createServer((_request, response) => response.end('{"ok":true}')); + const port = await listen(server); + const session = await createValidationSession(`http://provider.example.test:${port}/v1`, { + env: {}, + lookup: async () => [{ address: "127.0.0.1", family: 4 }], + allowPrivateAddressesForTesting: true, + }); + + await expect( + session!.request({ + url: "http://different.example.test/v1/responses", + body: "{}", + timeoutMs: 1_000, + }), + ).resolves.toMatchObject({ ok: false, message: "validation session origin mismatch" }); + session!.close(); + }); + + it("keeps proxy, curl-specific TLS, IP, local, and sandbox endpoints on curl", () => { + expect( + getValidationSessionIneligibility("https://provider.example.test/v1", { + HTTPS_PROXY: "http://proxy.example.test:8080", + }), + ).toBe("proxy_configured"); + expect( + getValidationSessionIneligibility("https://provider.example.test/v1", { + HTTPS_PROXY: "http://proxy.example.test:8080", + NO_PROXY: "provider.example.test", + }), + ).toBeNull(); + expect( + getValidationSessionIneligibility("https://api.provider.example.test/v1", { + HTTPS_PROXY: "http://proxy.example.test:8080", + NO_PROXY: ".provider.example.test", + }), + ).toBeNull(); + expect(getValidationSessionIneligibility("https://127.0.0.1/v1", {})).toBe("ip_literal"); + expect(getValidationSessionIneligibility("https://[::1]/v1", {})).toBe("ip_literal"); + expect(getValidationSessionIneligibility("http://localhost:8000/v1", {})).toBe( + "local_endpoint", + ); + expect(getValidationSessionIneligibility("http://localhost.:8000/v1", {})).toBe( + "local_endpoint", + ); + expect(getValidationSessionIneligibility("http://host.openshell.internal/v1", {})).toBe( + "sandbox_internal_endpoint", + ); + expect(getValidationSessionIneligibility("http://host.docker.internal/v1", {})).toBe( + "docker_internal_endpoint", + ); + expect(getValidationSessionIneligibility("http://host.docker.internal./v1", {})).toBe( + "docker_internal_endpoint", + ); + }); + + it.each([ + "CURL_CA_BUNDLE", + "SSL_CERT_FILE", + "SSL_CERT_DIR", + ] as const)("keeps the %s curl-specific TLS override on curl", (envName) => { + expect( + getValidationSessionIneligibility("https://provider.example.test/v1", { + [envName]: "/tmp/corporate-ca", + }), + ).toBe("curl_tls_configured"); + }); +}); diff --git a/src/lib/adapters/http/validation-session.ts b/src/lib/adapters/http/validation-session.ts new file mode 100644 index 00000000000..c82a261578a --- /dev/null +++ b/src/lib/adapters/http/validation-session.ts @@ -0,0 +1,390 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import dns from "node:dns/promises"; +import http from "node:http"; +import https from "node:https"; +import net from "node:net"; +import { isPrivateHostname, isPrivateIp } from "../../private-networks"; +import { addTraceEvent, withTraceSpan } from "../../trace"; +import type { CurlProbeResult } from "./probe"; +import { summarizeProbeFailure } from "./probe"; + +export interface ValidationSessionRequest { + url: string; + headers?: Record; + body: string; + timeoutMs: number; +} + +export interface ValidationSession { + request(input: ValidationSessionRequest): Promise; + close(): void; +} + +export type ValidationDnsLookup = ( + hostname: string, + options: { all: true; verbatim: true }, +) => Promise>; + +export interface ValidationSessionOptions { + env?: NodeJS.ProcessEnv; + lookup?: ValidationDnsLookup; + /** Addresses approved by the custom-endpoint SSRF preflight. */ + pinnedAddresses?: readonly string[]; + onSocket?: (socket: net.Socket) => void; + dnsTimeoutMs?: number; + /** @internal Allows local HTTP servers in transport unit tests. */ + allowPrivateAddressesForTesting?: boolean; +} + +const PROXY_ENV_NAMES = [ + "HTTP_PROXY", + "HTTPS_PROXY", + "ALL_PROXY", + "http_proxy", + "https_proxy", + "all_proxy", +] as const; + +const CURL_TLS_ENV_NAMES = ["CURL_CA_BUNDLE", "SSL_CERT_FILE", "SSL_CERT_DIR"] as const; +const MAX_RESPONSE_BYTES = 8 * 1024 * 1024; + +function safeErrorMessage(error: unknown): string { + const message = error instanceof Error ? error.message : String(error); + return message.replace(/(https?:\/\/[^\s?]+)\?[^\s]*/gi, "$1?[redacted]").slice(0, 256); +} + +function configured(env: NodeJS.ProcessEnv, names: readonly string[]): boolean { + return names.some((name) => Boolean(env[name]?.trim())); +} + +function isNoProxyEndpoint(endpoint: URL, env: NodeJS.ProcessEnv): boolean { + const raw = env.NO_PROXY ?? env.no_proxy ?? ""; + const hostname = endpoint.hostname.toLowerCase(); + const port = endpoint.port || (endpoint.protocol === "https:" ? "443" : "80"); + return raw + .split(",") + .map((entry) => entry.trim().toLowerCase()) + .filter(Boolean) + .some((entry) => { + if (entry === "*") return true; + const lastColon = entry.lastIndexOf(":"); + const includesPort = lastColon > 0 && /^\d+$/.test(entry.slice(lastColon + 1)); + const candidateHost = includesPort ? entry.slice(0, lastColon) : entry; + const candidatePort = includesPort ? entry.slice(lastColon + 1) : null; + if (candidatePort !== null && candidatePort !== port) return false; + if (candidateHost.startsWith(".")) { + return hostname === candidateHost.slice(1) || hostname.endsWith(candidateHost); + } + return hostname === candidateHost; + }); +} + +function withTimeout(operation: Promise, timeoutMs: number, message: string): Promise { + return new Promise((resolve, reject) => { + const timer = setTimeout(() => { + reject(Object.assign(new Error(message), { code: "ETIMEDOUT" })); + }, timeoutMs); + operation.then( + (value) => { + clearTimeout(timer); + resolve(value); + }, + (error) => { + clearTimeout(timer); + reject(error); + }, + ); + }); +} + +export function getValidationSessionIneligibility( + endpointUrl: string, + env: NodeJS.ProcessEnv = process.env, +): string | null { + let parsed: URL; + try { + parsed = new URL(endpointUrl); + } catch { + return "invalid_url"; + } + if (parsed.protocol !== "http:" && parsed.protocol !== "https:") return "unsupported_protocol"; + if (parsed.username || parsed.password) return "embedded_credentials"; + if (parsed.search || parsed.hash) return "endpoint_query_or_fragment"; + const hostname = parsed.hostname.replace(/\.$/, "").toLowerCase(); + const literalHostname = + hostname.startsWith("[") && hostname.endsWith("]") ? hostname.slice(1, -1) : hostname; + if (net.isIP(literalHostname)) return "ip_literal"; + if (hostname === "localhost" || hostname.endsWith(".localhost")) { + return "local_endpoint"; + } + if (hostname === "host.openshell.internal") return "sandbox_internal_endpoint"; + if (hostname === "host.docker.internal") return "docker_internal_endpoint"; + if (isPrivateHostname(parsed.hostname)) return "private_endpoint"; + // Node's built-in agents do not implement forward-proxy tunnelling. Keep curl + // authoritative whenever proxy routing may be part of endpoint reachability. + if (configured(env, PROXY_ENV_NAMES) && !isNoProxyEndpoint(parsed, env)) { + return "proxy_configured"; + } + // NODE_EXTRA_CA_CERTS is consumed by Node itself at process startup. Curl-only + // CA overrides are not, so those requests remain on the compatibility path. + if (configured(env, CURL_TLS_ENV_NAMES)) return "curl_tls_configured"; + return null; +} + +function curlStatusForError(error: NodeJS.ErrnoException): number { + if (error.name === "AbortError" || error.code === "ETIMEDOUT") return 28; + if (error.code === "ENOTFOUND" || error.code === "EAI_AGAIN") return 6; + if ( + error.code === "ECONNREFUSED" || + error.code === "ECONNRESET" || + error.code === "EHOSTUNREACH" || + error.code === "ENETUNREACH" + ) { + return 7; + } + if (error.code?.startsWith("CERT_") || error.code?.startsWith("ERR_TLS_")) return 60; + return 1; +} + +export function buildLookup( + addresses: Array<{ address: string; family: number }>, +): net.LookupFunction { + return (_hostname, options, callback) => { + const requestedFamily = typeof options === "object" ? options.family : undefined; + const eligible = requestedFamily + ? addresses.filter((entry) => entry.family === requestedFamily) + : addresses; + if (requestedFamily && eligible.length === 0) { + callback( + Object.assign(new Error(`no resolved IPv${requestedFamily} address is available`), { + code: "ENOTFOUND", + }), + [], + ); + return; + } + const selected = eligible; + if (options?.all) { + callback(null, selected); + return; + } + const first = selected[0]; + callback(null, first.address, first.family); + }; +} + +export async function createValidationSession( + endpointUrl: string, + options: ValidationSessionOptions = {}, +): Promise { + const env = options.env ?? process.env; + const ineligible = getValidationSessionIneligibility(endpointUrl, env); + if (ineligible) { + addTraceEvent("validation_transport_fallback", { reason: ineligible }); + return null; + } + + const endpoint = new URL(endpointUrl); + const lookup = options.lookup ?? (dns.lookup as ValidationDnsLookup); + let addresses: Array<{ address: string; family: number }>; + if (options.pinnedAddresses && options.pinnedAddresses.length > 0) { + addresses = options.pinnedAddresses.map((address) => ({ address, family: net.isIP(address) })); + if (addresses.some(({ family }) => family === 0)) { + addTraceEvent("validation_transport_fallback", { reason: "invalid_pinned_address" }); + return null; + } + } else { + try { + addresses = await withTraceSpan( + "nemoclaw.inference.validation_dns_lookup", + { "server.address": endpoint.hostname }, + () => + withTimeout( + lookup(endpoint.hostname, { all: true, verbatim: true }), + options.dnsTimeoutMs ?? 5_000, + "validation DNS lookup timed out", + ), + ); + } catch (error) { + addTraceEvent("validation_transport_fallback", { + reason: "dns_lookup_failed", + error_code: (error as NodeJS.ErrnoException).code ?? "unknown", + error_message: safeErrorMessage(error), + }); + return null; + } + } + if (addresses.length === 0) { + addTraceEvent("validation_transport_fallback", { reason: "dns_lookup_empty" }); + return null; + } + if ( + !options.allowPrivateAddressesForTesting && + addresses.some(({ address }) => isPrivateIp(address)) + ) { + addTraceEvent("validation_transport_fallback", { reason: "private_dns_address" }); + return null; + } + + const sharedOptions = { + keepAlive: true, + maxSockets: 1, + maxFreeSockets: 1, + // Ask Node's net stack to consume the full lookup result and race address + // families rather than pinning validation to the first IPv4/IPv6 answer. + autoSelectFamily: true, + lookup: buildLookup(addresses), + }; + const agent = + endpoint.protocol === "https:" ? new https.Agent(sharedOptions) : new http.Agent(sharedOptions); + const seenSockets = new WeakSet(); + const endpointOrigin = endpoint.origin; + let closed = false; + + addTraceEvent("validation_transport_selected", { + transport: "node_keepalive", + address_count: addresses.length, + }); + + return { + request(input) { + const requestUrl = new URL(input.url); + if (requestUrl.origin !== endpointOrigin) { + return Promise.resolve({ + ok: false, + httpStatus: 0, + curlStatus: 1, + body: "", + stderr: "validation session origin mismatch", + message: "validation session origin mismatch", + }); + } + return withTraceSpan( + "nemoclaw.inference.node_validation_request", + { "http.url": requestUrl.origin, transport: "node_keepalive" }, + () => + new Promise((resolve) => { + const target = new URL(input.url); + let settled = false; + let responseStarted = false; + let status = 0; + const chunks: Buffer[] = []; + let receivedBytes = 0; + let overallTimer: NodeJS.Timeout | undefined; + let terminalError: NodeJS.ErrnoException | undefined; + const finish = (result: CurlProbeResult) => { + if (settled) return; + settled = true; + if (overallTimer) clearTimeout(overallTimer); + resolve(result); + }; + const requestImpl = target.protocol === "https:" ? https.request : http.request; + const request = requestImpl( + target, + { + agent, + method: "POST", + headers: { + "content-type": "application/json", + "content-length": Buffer.byteLength(input.body).toString(), + ...input.headers, + }, + }, + (response) => { + responseStarted = true; + status = response.statusCode ?? 0; + const failResponse = (rawError: NodeJS.ErrnoException) => { + request.destroy(); + const body = Buffer.concat(chunks).toString("utf8"); + const curlStatus = curlStatusForError(rawError); + finish({ + ok: false, + httpStatus: status, + curlStatus, + body, + stderr: rawError.message, + message: summarizeProbeFailure(body, status, curlStatus, rawError.message), + }); + }; + response.on("data", (chunk: Buffer | string) => { + const buffer = Buffer.isBuffer(chunk) ? chunk : Buffer.from(chunk); + receivedBytes += buffer.length; + if (receivedBytes > MAX_RESPONSE_BYTES) { + terminalError = Object.assign(new Error("validation response exceeded 8 MiB"), { + code: "EFBIG", + }); + request.destroy(terminalError); + return; + } + chunks.push(buffer); + }); + response.on("end", () => { + const body = Buffer.concat(chunks).toString("utf8"); + const ok = status >= 200 && status < 300; + finish({ + ok, + httpStatus: status, + curlStatus: 0, + body, + stderr: "", + message: ok ? `HTTP ${status}` : summarizeProbeFailure(body, status, 0, ""), + }); + }); + response.on("close", () => { + if (response.complete || settled) return; + failResponse( + terminalError ?? + Object.assign(new Error("validation response closed early"), { + code: "ECONNRESET", + }), + ); + }); + response.on("error", failResponse); + }, + ); + request.on("socket", (socket) => { + if (!seenSockets.has(socket)) { + seenSockets.add(socket); + options.onSocket?.(socket); + addTraceEvent("validation_socket_opened", { reused: false }); + } else { + addTraceEvent("validation_socket_reused", { reused: true }); + } + }); + overallTimer = setTimeout(() => { + terminalError = Object.assign(new Error("validation request timed out"), { + code: "ETIMEDOUT", + }); + request.destroy(terminalError); + }, input.timeoutMs); + request.on("error", (rawError: NodeJS.ErrnoException) => { + const body = Buffer.concat(chunks).toString("utf8"); + const curlStatus = curlStatusForError(rawError); + finish({ + ok: false, + httpStatus: status, + curlStatus, + body, + stderr: rawError.message, + message: summarizeProbeFailure( + body, + responseStarted ? status : 0, + curlStatus, + rawError.message, + ), + }); + }); + request.end(input.body); + }), + ); + }, + close() { + if (closed) return; + closed = true; + agent.destroy(); + addTraceEvent("validation_transport_closed", { transport: "node_keepalive" }); + }, + }; +} diff --git a/src/lib/inference/onboard-probes.ts b/src/lib/inference/onboard-probes.ts index add2bb6eb73..a3558b01250 100644 --- a/src/lib/inference/onboard-probes.ts +++ b/src/lib/inference/onboard-probes.ts @@ -45,12 +45,19 @@ const { runChatCompletionsRetryLoop, } = require("./probe-retry"); const { probeAnthropicEndpoint } = require("./probe-anthropic"); +const { probeOpenAiLikeEndpointWithValidationSession } = require("./openai-validation-session"); +const { + getChatCompletionsProbePayload, + isDeepSeekV4ProModel, + isKimiK26Model, +} = require("./openai-probe-models"); const { buildValidationProbeTimingProfile, getValidationProbeCurlArgs, getDeepSeekV4ProValidationProbeCurlArgs, getKimiK26ValidationProbeCurlArgs, getExtendedNvidiaEndpointValidationProbeCurlArgs, + getCurlMaxTimeSeconds, getProbeProcessTimeoutMs, } = require("./probe-http-helpers"); @@ -293,6 +300,15 @@ function calibrateOpenAiLikeValidationTiming(baseUrl, options = {}) { }); } +function resolveOpenAiLikeValidationTiming(baseUrl, options = {}) { + return ( + options.validationTiming ?? + (options.calibrateTimeouts === true + ? calibrateOpenAiLikeValidationTiming(baseUrl, options) + : undefined) + ); +} + // ── Responses API probe ────────────────────────────────────────── function probeResponsesToolCalling(endpointUrl, model, apiKey, options = {}) { @@ -482,14 +498,6 @@ function probeChatCompletionsToolCalling(endpointUrl, model, apiKey, options = { } // ── OpenAI-like probe ──────────────────────────────────────────── -function isDeepSeekV4ProModel(model) { - return String(model || "").toLowerCase() === "deepseek-ai/deepseek-v4-pro"; -} - -function isKimiK26Model(model) { - return String(model || "").toLowerCase() === "moonshotai/kimi-k2.6"; -} - function needsExtendedNvidiaEndpointValidationBudget(model) { return EXTENDED_NVIDIA_ENDPOINT_VALIDATION_MODELS.has(String(model || "").toLowerCase()); } @@ -503,35 +511,6 @@ function getChatCompletionsProbeTimingArgs(model, opts) { return getValidationProbeCurlArgs(opts); } -function getChatCompletionsProbePayload(model) { - const payload = { - model, - messages: [{ role: "user", content: "Reply with exactly: OK" }], - max_tokens: 8, - }; - - if (isDeepSeekV4ProModel(model)) { - return { - ...payload, - temperature: 1, - top_p: 0.95, - max_tokens: 8192, - chat_template_kwargs: { thinking: false }, - stream: true, - }; - } - - if (isKimiK26Model(model)) { - return { - ...payload, - max_tokens: 8, - chat_template_kwargs: { thinking: false }, - }; - } - - return payload; -} - // credentialArgs is the curl argument slice that carries the auth credential // for the probe — typically ["--config", ] from the auth-config // module. The parameter used to be named `authHeader` and used to receive a @@ -729,11 +708,7 @@ function probeOpenAiLikeEndpoint(endpointUrl, model, apiKey, options = {}) { } const baseUrl = String(endpointUrl).replace(/\/+$/, ""); - const validationTiming = - options.validationTiming ?? - (options.calibrateTimeouts === true - ? calibrateOpenAiLikeValidationTiming(baseUrl, options) - : undefined); + const validationTiming = resolveOpenAiLikeValidationTiming(baseUrl, options); if (validationTiming) { options = { ...options, validationTiming }; } @@ -977,6 +952,37 @@ function probeOpenAiLikeEndpoint(endpointUrl, model, apiKey, options = {}) { } } +async function probeOpenAiLikeEndpointOptimized(endpointUrl, model, apiKey, options = {}) { + const normalizedKey = apiKey ? normalizeCredentialValue(apiKey) : ""; + const baseUrl = String(endpointUrl).replace(/\/+$/, ""); + const validationTiming = resolveOpenAiLikeValidationTiming(baseUrl, options); + const sessionProbeOptions = validationTiming ? { ...options, validationTiming } : options; + return probeOpenAiLikeEndpointWithValidationSession( + endpointUrl, + model, + normalizedKey, + sessionProbeOptions, + { + legacyProbe: probeOpenAiLikeEndpoint, + hasResponsesToolCall, + hasChatCompletionsToolCall, + hasChatCompletionsToolCallLeak, + getChatPayload: getChatCompletionsProbePayload, + getResponsesTimeoutMs: (probeOptions) => + getCurlMaxTimeSeconds(getValidationProbeCurlArgs(getProbeTimingOptions(probeOptions))) * + 1000, + getChatTimeoutMs: (probeModel, probeOptions) => { + const platformOptions = getProbeTimingOptions(probeOptions); + return ( + getCurlMaxTimeSeconds(getChatCompletionsProbeTimingArgs(probeModel, platformOptions)) * + 1000 + ); + }, + sessionOptions: sessionProbeOptions.validationSessionOptions, + }, + ); +} + // ── Anthropic probe ────────────────────────────────────────────── module.exports = { @@ -997,6 +1003,7 @@ module.exports = { probeResponsesToolCalling, probeChatCompletionsToolCalling, probeOpenAiLikeEndpoint, + probeOpenAiLikeEndpointOptimized, probeAnthropicEndpoint, RETRIABLE_HTTP_PROBE_STATUSES, }; @@ -1028,7 +1035,7 @@ export function shouldSmokeOpenAiLikeOnboardRoute( ); } -export function verifyOnboardInferenceSmoke(options: any) { +export async function verifyOnboardInferenceSmoke(options: any) { if ( !options.forceOpenAiLike && !shouldSmokeOpenAiLikeOnboardRoute(options.provider, options.credentialEnv) @@ -1042,7 +1049,7 @@ export function verifyOnboardInferenceSmoke(options: any) { const apiKey = credentialEnv ? resolveProviderCredential(credentialEnv) || getCredential(credentialEnv) || "" : ""; - const probe = probeOpenAiLikeEndpoint(endpointUrl, options.model, apiKey, { + const probe = await probeOpenAiLikeEndpointOptimized(endpointUrl, options.model, apiKey, { authMode: getProbeAuthMode(options.provider), extraHeaders: getProbeExtraHeaders(options.provider), skipResponsesProbe: true, diff --git a/src/lib/inference/openai-probe-models.ts b/src/lib/inference/openai-probe-models.ts new file mode 100644 index 00000000000..06a58df62cd --- /dev/null +++ b/src/lib/inference/openai-probe-models.ts @@ -0,0 +1,39 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +export function isDeepSeekV4ProModel(model: unknown): boolean { + return String(model || "").toLowerCase() === "deepseek-ai/deepseek-v4-pro"; +} + +export function isKimiK26Model(model: unknown): boolean { + return String(model || "").toLowerCase() === "moonshotai/kimi-k2.6"; +} + +export function getChatCompletionsProbePayload(model: string): Record { + const payload = { + model, + messages: [{ role: "user", content: "Reply with exactly: OK" }], + max_tokens: 8, + }; + + if (isDeepSeekV4ProModel(model)) { + return { + ...payload, + temperature: 1, + top_p: 0.95, + max_tokens: 8192, + chat_template_kwargs: { thinking: false }, + stream: true, + }; + } + + if (isKimiK26Model(model)) { + return { + ...payload, + max_tokens: 8, + chat_template_kwargs: { thinking: false }, + }; + } + + return payload; +} diff --git a/src/lib/inference/openai-validation-session-auth.test.ts b/src/lib/inference/openai-validation-session-auth.test.ts new file mode 100644 index 00000000000..c6c8c35d59b --- /dev/null +++ b/src/lib/inference/openai-validation-session-auth.test.ts @@ -0,0 +1,73 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import http from "node:http"; +import { describe, expect, it } from "vitest"; +import { probeOpenAiLikeEndpointWithValidationSession } from "./openai-validation-session"; +import { + createOpenAiValidationTestDeps, + useOpenAiValidationTestServers, +} from "./openai-validation-session.test-helpers"; + +const listen = useOpenAiValidationTestServers(); + +describe("OpenAI validation authentication and headers", () => { + it("keeps query-parameter authentication out of request headers", async () => { + let observedUrl = ""; + let observedAuthorization: string | undefined; + const server = http.createServer((request, response) => { + observedUrl = request.url ?? ""; + observedAuthorization = request.headers.authorization; + request.resume(); + response.end('{"choices":[{"message":{"content":"OK"}}]}'); + }); + const port = await listen(server); + const harness = createOpenAiValidationTestDeps(); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "query-secret", + { authMode: "query-param", skipResponsesProbe: true }, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(observedAuthorization).toBeUndefined(); + expect(new URL(observedUrl, "http://provider.example.test").searchParams.get("key")).toBe( + "query-secret", + ); + }); + + it("preserves provider extra headers on the native validation request", async () => { + let observedReferer: string | undefined; + let observedTitle: string | undefined; + const server = http.createServer((request, response) => { + observedReferer = request.headers["http-referer"] as string | undefined; + observedTitle = request.headers["x-openrouter-title"] as string | undefined; + request.resume(); + response.end('{"choices":[{"message":{"content":"OK"}}]}'); + }); + const port = await listen(server); + const harness = createOpenAiValidationTestDeps(); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { + skipResponsesProbe: true, + extraHeaders: [ + "HTTP-Referer: https://github.com/NVIDIA/NemoClaw", + "X-OpenRouter-Title: NemoClaw", + ], + }, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(observedReferer).toBe("https://github.com/NVIDIA/NemoClaw"); + expect(observedTitle).toBe("NemoClaw"); + expect(harness.legacyProbe).not.toHaveBeenCalled(); + }); +}); diff --git a/src/lib/inference/openai-validation-session-fallback.test.ts b/src/lib/inference/openai-validation-session-fallback.test.ts new file mode 100644 index 00000000000..05b13c4d5b6 --- /dev/null +++ b/src/lib/inference/openai-validation-session-fallback.test.ts @@ -0,0 +1,275 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import http from "node:http"; +import { afterEach, describe, expect, it, vi } from "vitest"; +import { + type OpenAiValidationSessionDeps, + probeOpenAiLikeEndpointWithValidationSession, +} from "./openai-validation-session"; +import { + createOpenAiValidationTestDeps, + useOpenAiValidationTestServers, +} from "./openai-validation-session.test-helpers"; + +const listen = useOpenAiValidationTestServers(); + +afterEach(() => { + vi.unstubAllEnvs(); +}); + +describe("OpenAI validation curl fallback", () => { + it("recovers natively after transient HTTP failures", async () => { + vi.stubEnv("NEMOCLAW_TEST_NO_SLEEP", "1"); + const responsePlan = [ + [503, '{"error":{"message":"retry"}}'], + [429, '{"error":{"message":"retry"}}'], + [200, '{"choices":[{"message":{"content":"OK"}}]}'], + ] as const; + let requests = 0; + const server = http.createServer((request, response) => { + request.resume(); + const [statusCode, body] = responsePlan[requests] ?? responsePlan.at(-1)!; + requests += 1; + response.statusCode = statusCode; + response.end(body); + }); + const port = await listen(server); + const harness = createOpenAiValidationTestDeps(); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { skipResponsesProbe: true }, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(requests).toBe(3); + expect(harness.legacyProbe).not.toHaveBeenCalled(); + }); + + it("falls back once after transient HTTP retries are exhausted", async () => { + vi.stubEnv("NEMOCLAW_TEST_NO_SLEEP", "1"); + let requests = 0; + const server = http.createServer((request, response) => { + request.resume(); + requests += 1; + response.statusCode = 503; + response.end('{"error":{"message":"still unavailable"}}'); + }); + const port = await listen(server); + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: false, + message: "curl retry diagnostic", + })); + const harness = createOpenAiValidationTestDeps(legacyProbe); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { skipResponsesProbe: true }, + harness, + ); + + expect(result).toEqual({ ok: false, message: "curl retry diagnostic" }); + expect(requests).toBe(4); + expect(legacyProbe).toHaveBeenCalledTimes(1); + }); + + it("replays through curl after a terminal native failure", async () => { + const server = http.createServer((request, response) => { + request.resume(); + response.statusCode = 401; + response.end('{"error":{"message":"invalid key"}}'); + }); + const port = await listen(server); + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: false, + message: "curl diagnostic", + })); + const harness = createOpenAiValidationTestDeps(legacyProbe); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "bad-key", + { skipResponsesProbe: true }, + harness, + ); + + expect(result).toEqual({ ok: false, message: "curl diagnostic" }); + expect(legacyProbe).toHaveBeenCalledTimes(1); + }); + + it("replays through curl once after a native connection reset", async () => { + const server = http.createServer((request) => { + request.socket.destroy(); + }); + const port = await listen(server); + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: false, + message: "curl connection diagnostic", + })); + const harness = createOpenAiValidationTestDeps(legacyProbe); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { skipResponsesProbe: true }, + harness, + ); + + expect(result).toEqual({ ok: false, message: "curl connection diagnostic" }); + expect(legacyProbe).toHaveBeenCalledTimes(1); + }); + + it("replays through curl when DNS pre-resolution exceeds its deadline", async () => { + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: false, + message: "curl DNS diagnostic", + })); + const lookup = vi.fn(() => new Promise>(() => {})); + const harness = createOpenAiValidationTestDeps(legacyProbe); + harness.sessionOptions = { env: {}, lookup, dnsTimeoutMs: 10 }; + + const result = await probeOpenAiLikeEndpointWithValidationSession( + "https://provider.example.test/v1", + "test-model", + "test-key", + {}, + harness, + ); + + expect(result).toEqual({ ok: false, message: "curl DNS diagnostic" }); + expect(lookup).toHaveBeenCalledTimes(1); + expect(legacyProbe).toHaveBeenCalledTimes(1); + }); + + it("replays through curl after a connection reset during Responses streaming", async () => { + const handleRequest = vi + .fn() + .mockImplementationOnce((_request, response) => { + response.end('{"output":[{"type":"message"}]}'); + }) + .mockImplementationOnce((request) => { + request.socket.destroy(); + }); + const server = http.createServer((request, response) => { + request.resume(); + handleRequest(request, response); + }); + const port = await listen(server); + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: false, + message: "curl streaming diagnostic", + })); + const harness = createOpenAiValidationTestDeps(legacyProbe); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { probeStreaming: true }, + harness, + ); + + expect(result).toEqual({ ok: false, message: "curl streaming diagnostic" }); + expect(handleRequest).toHaveBeenCalledTimes(2); + expect(legacyProbe).toHaveBeenCalledTimes(1); + }); + + it("uses curl without DNS pre-resolution when a proxy is configured", async () => { + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: true, + api: "openai-completions", + })); + const lookup = vi.fn(); + const harness = createOpenAiValidationTestDeps(legacyProbe); + harness.sessionOptions = { + env: { HTTPS_PROXY: "http://proxy.example.test:8080" }, + lookup, + }; + + const result = await probeOpenAiLikeEndpointWithValidationSession( + "https://provider.example.test/v1", + "test-model", + "test-key", + {}, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(legacyProbe).toHaveBeenCalledTimes(1); + expect(lookup).not.toHaveBeenCalled(); + }); + + it.each([ + "CURL_CA_BUNDLE", + "SSL_CERT_FILE", + "SSL_CERT_DIR", + ])("uses curl without DNS pre-resolution when %s is configured", async (envName) => { + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: true, + api: "openai-completions", + })); + const lookup = vi.fn(); + const harness = createOpenAiValidationTestDeps(legacyProbe); + harness.sessionOptions = { env: { [envName]: "/tmp/provider-tls-config" }, lookup }; + + const result = await probeOpenAiLikeEndpointWithValidationSession( + "https://provider.example.test/v1", + "test-model", + "test-key", + {}, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(legacyProbe).toHaveBeenCalledTimes(1); + expect(lookup).not.toHaveBeenCalled(); + }); + + it("keeps preflight-pinned endpoints on curl without native DNS", async () => { + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: true, + api: "openai-completions", + })); + const harness = createOpenAiValidationTestDeps(legacyProbe); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + "https://provider.example.test/v1", + "test-model", + "test-key", + { pinnedAddresses: ["203.0.113.10"] }, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(legacyProbe).toHaveBeenCalledTimes(1); + expect(harness.sessionOptions!.lookup).not.toHaveBeenCalled(); + }); + + it("keeps DeepSeek V4 Pro on its specialized legacy streaming probe", async () => { + const legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(() => ({ + ok: true, + api: "openai-completions", + })); + const harness = createOpenAiValidationTestDeps(legacyProbe); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + "https://provider.example.test/v1", + "deepseek-ai/deepseek-v4-pro", + "test-key", + {}, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(legacyProbe).toHaveBeenCalledTimes(1); + expect(harness.sessionOptions!.lookup).not.toHaveBeenCalled(); + }); +}); diff --git a/src/lib/inference/openai-validation-session.test-helpers.ts b/src/lib/inference/openai-validation-session.test-helpers.ts new file mode 100644 index 00000000000..ccaf5b68e78 --- /dev/null +++ b/src/lib/inference/openai-validation-session.test-helpers.ts @@ -0,0 +1,72 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import http from "node:http"; +import { afterEach, expect, vi } from "vitest"; +import type { CurlProbeResult } from "../adapters/http/probe"; +import type { OpenAiValidationSessionDeps } from "./openai-validation-session"; + +export function useOpenAiValidationTestServers(): (server: http.Server) => Promise { + const servers: http.Server[] = []; + + afterEach(async () => { + await Promise.all( + servers.splice(0).map( + (server) => + new Promise((resolve) => { + server.close(() => resolve()); + server.closeAllConnections(); + }), + ), + ); + }); + + return async (server: http.Server): Promise => { + await new Promise((resolve, reject) => { + const onError = (error: Error) => { + server.off("listening", onListening); + reject(error); + }; + const onListening = () => { + server.off("error", onError); + resolve(); + }; + server.once("error", onError); + server.once("listening", onListening); + server.listen(0, "127.0.0.1"); + }); + servers.push(server); + const address = server.address(); + expect(address).toBeTruthy(); + expect(typeof address).toBe("object"); + return (address as import("node:net").AddressInfo).port; + }; +} + +const legacySuccess = (): CurlProbeResult => ({ + ok: true, + httpStatus: 200, + curlStatus: 0, + body: "", + stderr: "", + message: "legacy", +}); + +export function createOpenAiValidationTestDeps( + legacyProbe: OpenAiValidationSessionDeps["legacyProbe"] = vi.fn(legacySuccess), +): OpenAiValidationSessionDeps { + return { + legacyProbe, + hasResponsesToolCall: (body: string) => body.includes('"type":"function_call"'), + hasChatCompletionsToolCall: (body: string) => body.includes('"tool_calls"'), + hasChatCompletionsToolCallLeak: () => false, + getChatPayload: (model: string) => ({ model, messages: [] }), + getResponsesTimeoutMs: () => 1_000, + getChatTimeoutMs: () => 1_000, + sessionOptions: { + env: {}, + lookup: vi.fn(async () => [{ address: "127.0.0.1", family: 4 }]), + allowPrivateAddressesForTesting: true, + }, + }; +} diff --git a/src/lib/inference/openai-validation-session.test.ts b/src/lib/inference/openai-validation-session.test.ts new file mode 100644 index 00000000000..64c0eeeff77 --- /dev/null +++ b/src/lib/inference/openai-validation-session.test.ts @@ -0,0 +1,116 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import http from "node:http"; +import { describe, expect, it } from "vitest"; +import { probeOpenAiLikeEndpointWithValidationSession } from "./openai-validation-session"; +import { + createOpenAiValidationTestDeps, + useOpenAiValidationTestServers, +} from "./openai-validation-session.test-helpers"; + +const listen = useOpenAiValidationTestServers(); + +describe("OpenAI validation keepalive sequence", () => { + it("uses one connection for Responses semantic fallback and Chat success", async () => { + let connections = 0; + const paths: string[] = []; + const server = http.createServer((request, response) => { + paths.push(request.url ?? ""); + request.resume(); + response.setHeader("content-type", "application/json"); + response.end( + request.url?.endsWith("/responses") + ? '{"output":[{"type":"message"}]}' + : '{"choices":[{"message":{"content":"OK"}}]}', + ); + }); + server.on("connection", () => { + connections += 1; + }); + const port = await listen(server); + const harness = createOpenAiValidationTestDeps(); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { requireResponsesToolCalling: true }, + harness, + ); + + expect(result).toMatchObject({ + ok: true, + api: "openai-completions", + label: "Chat Completions API", + }); + expect(harness.legacyProbe).not.toHaveBeenCalled(); + expect(harness.sessionOptions!.lookup).toHaveBeenCalledTimes(1); + expect(connections).toBe(1); + expect(paths).toEqual(["/v1/responses", "/v1/chat/completions"]); + }); + + it("reuses the connection across non-streaming, streaming, and Chat fallback", async () => { + let connections = 0; + let responsesCalls = 0; + const paths: string[] = []; + const server = http.createServer((request, response) => { + paths.push(request.url ?? ""); + request.resume(); + const isResponses = request.url?.endsWith("/responses") === true; + responsesCalls += Number(isResponses); + response.end( + isResponses + ? responsesCalls === 1 + ? '{"output":[{"type":"function_call"}]}' + : "event: response.completed\ndata: {}\n\n" + : '{"choices":[{"message":{"content":"OK"}}]}', + ); + }); + server.on("connection", () => { + connections += 1; + }); + const port = await listen(server); + const harness = createOpenAiValidationTestDeps(); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { requireResponsesToolCalling: true, probeStreaming: true }, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-completions" }); + expect(harness.legacyProbe).not.toHaveBeenCalled(); + expect(connections).toBe(1); + expect(paths).toEqual(["/v1/responses", "/v1/responses", "/v1/chat/completions"]); + }); + + it("returns Responses success when native streaming emits the required event", async () => { + const paths: string[] = []; + const server = http.createServer((request, response) => { + paths.push(request.url ?? ""); + request.resume(); + response.end( + paths.length === 1 + ? '{"output":[{"type":"message"}]}' + : "event: response.output_text.delta\ndata: {}\n\n", + ); + }); + const port = await listen(server); + const harness = createOpenAiValidationTestDeps(); + + const result = await probeOpenAiLikeEndpointWithValidationSession( + `http://provider.example.test:${port}/v1`, + "test-model", + "test-key", + { probeStreaming: true }, + harness, + ); + + expect(result).toMatchObject({ ok: true, api: "openai-responses" }); + expect(paths).toEqual(["/v1/responses", "/v1/responses"]); + expect(harness.legacyProbe).not.toHaveBeenCalled(); + }); +}); diff --git a/src/lib/inference/openai-validation-session.ts b/src/lib/inference/openai-validation-session.ts new file mode 100644 index 00000000000..90b3b0590d3 --- /dev/null +++ b/src/lib/inference/openai-validation-session.ts @@ -0,0 +1,344 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import { parseOpenAiLikeExtraHeaders } from "../adapters/http/auth-config"; +import type { CurlProbeResult } from "../adapters/http/probe"; +import { + createValidationSession, + type ValidationSessionOptions, +} from "../adapters/http/validation-session"; +import { addTraceEvent, withTraceSpan } from "../trace"; +import { isDeepSeekV4ProModel } from "./openai-probe-models"; + +const RETRIABLE_HTTP_STATUSES = new Set([429, 502, 503, 504]); +const RETRY_DELAYS_MS = [5_000, 15_000, 30_000]; + +export interface OpenAiValidationOptions { + authMode?: "bearer" | "query-param"; + extraHeaders?: readonly string[]; + requireResponsesToolCalling?: boolean; + requireChatCompletionsToolCalling?: boolean; + skipResponsesProbe?: boolean; + probeStreaming?: boolean; + isWsl?: boolean; + pinnedAddresses?: readonly string[]; + validationSessionOptions?: ValidationSessionOptions; +} + +export interface OpenAiValidationResult { + ok: boolean; + api?: string | null; + label?: string | null; + message?: string; + failures?: unknown[]; +} + +export interface OpenAiValidationSessionDeps { + legacyProbe( + endpointUrl: string, + model: string, + apiKey: string, + options: OpenAiValidationOptions, + ): OpenAiValidationResult; + hasResponsesToolCall(body: string): boolean; + hasChatCompletionsToolCall(body: string): boolean; + hasChatCompletionsToolCallLeak(body: string): boolean; + getChatPayload(model: string): Record; + getResponsesTimeoutMs(options: OpenAiValidationOptions): number; + getChatTimeoutMs(model: string, options: OpenAiValidationOptions): number; + sessionOptions?: ValidationSessionOptions; +} + +function responsesPayload(model: string, requireToolCall: boolean, stream = false): string { + if (!requireToolCall) { + return JSON.stringify({ + model, + input: "Reply with exactly: OK", + ...(stream ? { stream } : {}), + }); + } + return JSON.stringify({ + model, + input: "Call the emit_ok function with value OK. Do not answer with plain text.", + tool_choice: "required", + tools: [ + { + type: "function", + name: "emit_ok", + description: "Returns the probe value for validation.", + parameters: { + type: "object", + properties: { value: { type: "string" } }, + required: ["value"], + additionalProperties: false, + }, + }, + ], + ...(stream ? { stream } : {}), + }); +} + +function chatToolPayload(model: string): string { + return JSON.stringify({ + model, + messages: [ + { + role: "system", + content: + "You are a tool-calling assistant. When tools are available and the user asks for an action, call a tool.", + }, + { + role: "user", + content: + "Send hello to the current session. Use the sessions_send tool and do not answer in plain text.", + }, + ], + tools: [ + { + type: "function", + function: { + name: "sessions_send", + description: "Send a message to the active chat session.", + parameters: { + type: "object", + properties: { message: { type: "string" } }, + required: ["message"], + additionalProperties: false, + }, + }, + }, + { + type: "function", + function: { + name: "memory_search", + description: "Search memory for relevant prior context.", + parameters: { + type: "object", + properties: { query: { type: "string" } }, + required: ["query"], + additionalProperties: false, + }, + }, + }, + { + type: "function", + function: { + name: "web_fetch", + description: "Fetch a URL and summarize the result.", + parameters: { + type: "object", + properties: { url: { type: "string" } }, + required: ["url"], + additionalProperties: false, + }, + }, + }, + ], + tool_choice: "required", + temperature: 0, + max_tokens: 256, + stream: false, + }); +} + +function requestAuth( + rawUrl: string, + apiKey: string, + options: OpenAiValidationOptions, +): { url: string; headers: Record } { + const url = new URL(rawUrl); + const headers = Object.fromEntries( + parseOpenAiLikeExtraHeaders(options.extraHeaders).map(({ name, value }) => [name, value]), + ); + if (options.authMode === "query-param") { + if (apiKey) url.searchParams.set("key", apiKey); + return { url: url.toString(), headers }; + } + if (apiKey) headers.authorization = `Bearer ${apiKey}`; + return { url: url.toString(), headers }; +} + +function streamingEventTypes(body: string): Set { + const events = new Set(); + for (const line of body.split("\n")) { + const match = /^event:\s*(.+)$/i.exec(line.trim()); + if (match) events.add(match[1].trim()); + } + return events; +} + +async function waitForRetry(ms: number): Promise { + if (process.env.NEMOCLAW_TEST_NO_SLEEP === "1") return; + await new Promise((resolve) => setTimeout(resolve, ms)); +} + +function safeErrorDetails(error: unknown): { error_code: string; error_message: string } { + const value = error as NodeJS.ErrnoException; + const message = error instanceof Error ? error.message : String(error); + return { + error_code: value?.code ?? "unknown", + error_message: message.replace(/(https?:\/\/[^\s?]+)\?[^\s]*/gi, "$1?[redacted]").slice(0, 256), + }; +} + +async function requestWithHttpRetry( + name: string, + request: () => Promise, +): Promise { + let result = await request(); + let attempt = 1; + addTraceEvent("probe_result", { + attempt, + ok: result.ok, + http_status: result.httpStatus, + curl_status: result.curlStatus, + }); + for (const delayMs of RETRY_DELAYS_MS) { + if (result.curlStatus !== 0 || !RETRIABLE_HTTP_STATUSES.has(result.httpStatus)) break; + console.log( + ` ${name} validation returned HTTP ${result.httpStatus}; retrying in ${Math.round(delayMs / 1000)}s...`, + ); + await waitForRetry(delayMs); + attempt += 1; + result = await request(); + addTraceEvent("probe_result", { + attempt, + ok: result.ok, + http_status: result.httpStatus, + curl_status: result.curlStatus, + }); + } + return result; +} + +function shouldUseLegacyForModel(model: string): boolean { + // Invalid state: the native session does not reproduce DeepSeek V4 Pro's + // accepted late-first-token timeout result. The source of truth remains the + // specialized streaming Chat Completions path in onboard-probes.ts, which + // owns its payload, extended timeout, warning, and validated:false result. + // Trying native first would add a long duplicate request before curl fallback + // and could turn that accepted warning into a validation failure. The + // "keeps DeepSeek V4 Pro on its specialized legacy streaming probe" test + // locks direct legacy dispatch without native DNS. Remove this exception once + // both transports share the streaming timeout-continuation helper and return + // the same validation result. + return isDeepSeekV4ProModel(model); +} + +export async function probeOpenAiLikeEndpointWithValidationSession( + endpointUrl: string, + model: string, + apiKey: string, + options: OpenAiValidationOptions, + deps: OpenAiValidationSessionDeps, +): Promise { + if (shouldUseLegacyForModel(model)) { + addTraceEvent("validation_transport_fallback", { reason: "special_streaming_model" }); + return deps.legacyProbe(endpointUrl, model, apiKey, options); + } + // Custom-endpoint SSRF preflight pins approved addresses through curl's + // reviewed --resolve boundary. Keep that security path authoritative until + // native address pinning has equivalent end-to-end rebinding coverage. + if (options.pinnedAddresses && options.pinnedAddresses.length > 0) { + addTraceEvent("validation_transport_fallback", { reason: "preflight_address_pinning" }); + return deps.legacyProbe(endpointUrl, model, apiKey, options); + } + + const session = await createValidationSession(endpointUrl, { + ...deps.sessionOptions, + pinnedAddresses: options.pinnedAddresses ?? deps.sessionOptions?.pinnedAddresses, + }); + if (!session) return deps.legacyProbe(endpointUrl, model, apiKey, options); + + const baseUrl = endpointUrl.replace(/\/+$/, ""); + const nativeFailureFallback = async (reason: string): Promise => { + addTraceEvent("validation_transport_fallback", { reason }); + session.close(); + return deps.legacyProbe(endpointUrl, model, apiKey, options); + }; + + try { + if (!options.skipResponsesProbe) { + const auth = requestAuth(`${baseUrl}/responses`, apiKey, options); + const responses = await withTraceSpan( + "nemoclaw.inference.validation_probe", + { probe_name: "Responses API", api: "openai-responses" }, + () => + requestWithHttpRetry("Responses API", () => + session.request({ + ...auth, + body: responsesPayload(model, options.requireResponsesToolCalling === true), + timeoutMs: deps.getResponsesTimeoutMs(options), + }), + ), + ); + if (responses.curlStatus !== 0) return nativeFailureFallback("native_responses_failure"); + const responsesSemanticallyValid = + responses.ok && + (options.requireResponsesToolCalling !== true || deps.hasResponsesToolCall(responses.body)); + if (responsesSemanticallyValid) { + if (options.probeStreaming === true) { + const streamResult = await session.request({ + ...auth, + body: responsesPayload(model, false, true), + timeoutMs: deps.getResponsesTimeoutMs(options), + }); + const events = streamingEventTypes(streamResult.body); + if (streamResult.curlStatus !== 0 && streamResult.curlStatus !== 28) { + return nativeFailureFallback("native_streaming_failure"); + } + // Match onboard-probes.ts: a successful Responses payload without + // response.output_text.delta falls through to Chat Completions. This + // duplicate can be removed once both transports share event parsing. + if (!events.has("response.output_text.delta")) { + console.log( + " ℹ Responses API streaming response is missing required event: response.output_text.delta", + ); + } else { + return { ok: true, api: "openai-responses", label: "Responses API" }; + } + } else { + return { ok: true, api: "openai-responses", label: "Responses API" }; + } + } + } + + const auth = requestAuth(`${baseUrl}/chat/completions`, apiKey, options); + const chatBody = + options.requireChatCompletionsToolCalling === true + ? chatToolPayload(model) + : JSON.stringify(deps.getChatPayload(model)); + const chat = await withTraceSpan( + "nemoclaw.inference.validation_probe", + { probe_name: "Chat Completions API", api: "openai-completions" }, + () => + requestWithHttpRetry("Chat Completions API", () => + session.request({ + ...auth, + body: chatBody, + timeoutMs: deps.getChatTimeoutMs(model, options), + }), + ), + ); + if (chat.curlStatus !== 0) return nativeFailureFallback("native_chat_failure"); + if (!chat.ok) return nativeFailureFallback("native_terminal_http_failure"); + if (options.requireChatCompletionsToolCalling === true) { + if (!deps.hasChatCompletionsToolCall(chat.body)) { + return nativeFailureFallback( + deps.hasChatCompletionsToolCallLeak(chat.body) + ? "native_chat_tool_call_leak" + : "native_chat_tool_call_missing", + ); + } + } + return { ok: true, api: "openai-completions", label: "Chat Completions API" }; + } catch (error) { + addTraceEvent("validation_transport_error", { + reason: "native_unexpected_failure", + ...safeErrorDetails(error), + }); + return nativeFailureFallback("native_unexpected_failure"); + } finally { + session.close(); + } +} diff --git a/src/lib/onboard/bedrock-runtime.test.ts b/src/lib/onboard/bedrock-runtime.test.ts index 58942de3af9..7e8fe7f327b 100644 --- a/src/lib/onboard/bedrock-runtime.test.ts +++ b/src/lib/onboard/bedrock-runtime.test.ts @@ -3,10 +3,26 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { selectBedrockRuntimeCustomAnthropic } from "./bedrock-runtime"; +import { + selectBedrockRuntimeCustomAnthropic, + setupBedrockRuntimeInference, +} from "./bedrock-runtime"; import { BACK_TO_SELECTION } from "./credential-navigation"; const BEDROCK_URL = "https://bedrock-runtime.us-east-1.amazonaws.com"; +const BEDROCK_PROVIDER = "compatible-anthropic-endpoint"; +const BEDROCK_MODEL = "anthropic.claude"; +const BEDROCK_SUCCESS_LOG = ` ✓ Inference route set: ${BEDROCK_PROVIDER} / ${BEDROCK_MODEL}`; + +type BedrockSetupOptions = Parameters[0]; + +function deferred(): { promise: Promise; resolve: () => void } { + let resolve!: () => void; + const promise = new Promise((done) => { + resolve = done; + }); + return { promise, resolve }; +} function createBedrockRuntimeDependencies() { return { @@ -30,6 +46,38 @@ function clearBedrockAuthEnv(): void { delete process.env.COMPATIBLE_ANTHROPIC_API_KEY; } +function createBedrockSetupHarness( + verifyOnboardInferenceSmoke: BedrockSetupOptions["verifyOnboardInferenceSmoke"], +) { + process.env.COMPATIBLE_ANTHROPIC_API_KEY = "bedrock-compatible-token"; + const log = vi.fn(); + const updateSandbox = vi.fn(() => true); + const options: BedrockSetupOptions = { + ...createBedrockRuntimeDependencies(), + log, + sandboxName: "alpha", + provider: BEDROCK_PROVIDER, + model: BEDROCK_MODEL, + endpointUrl: BEDROCK_URL, + credentialEnv: "COMPATIBLE_ANTHROPIC_API_KEY", + isNonInteractive: () => false, + runOpenshell: vi.fn(() => ({ status: 0, stdout: "", stderr: "" })), + upsertProvider: vi.fn(() => ({ ok: true })), + verifyInferenceRoute: vi.fn(), + verifyOnboardInferenceSmoke, + ensureAdapter: vi.fn(async () => ({ + baseUrl: "http://host.openshell.internal:18081/v1", + localBaseUrl: "http://127.0.0.1:18081/v1", + logPath: "/tmp/bedrock-runtime-adapter.log", + credentialEnv: "NEMOCLAW_BEDROCK_RUNTIME_ADAPTER_TOKEN", + token: "adapter-token", + region: "us-east-1", + })), + updateSandbox, + }; + return { log, options, updateSandbox }; +} + afterEach(() => { clearBedrockAuthEnv(); vi.restoreAllMocks(); @@ -163,4 +211,35 @@ describe("Bedrock Runtime onboarding helper", () => { preferredInferenceApi: "openai-completions", }); }); + + it("waits for async smoke validation before persisting Bedrock route success (#3771)", async () => { + const smoke = deferred(); + const verifyOnboardInferenceSmoke = vi.fn(() => smoke.promise); + const { log, options, updateSandbox } = createBedrockSetupHarness(verifyOnboardInferenceSmoke); + + const setup = setupBedrockRuntimeInference(options); + await vi.waitFor(() => expect(verifyOnboardInferenceSmoke).toHaveBeenCalledOnce()); + + expect(updateSandbox).not.toHaveBeenCalled(); + expect(log).not.toHaveBeenCalledWith(BEDROCK_SUCCESS_LOG); + + smoke.resolve(); + await expect(setup).resolves.toEqual({ handled: true, result: { ok: true } }); + expect(updateSandbox).toHaveBeenCalledWith("alpha", { + model: BEDROCK_MODEL, + provider: BEDROCK_PROVIDER, + }); + expect(log).toHaveBeenCalledWith(BEDROCK_SUCCESS_LOG); + }); + + it("does not persist Bedrock route success when async smoke validation rejects (#3771)", async () => { + const verifyOnboardInferenceSmoke = vi.fn(async () => { + throw new Error("bedrock smoke rejected"); + }); + const { log, options, updateSandbox } = createBedrockSetupHarness(verifyOnboardInferenceSmoke); + + await expect(setupBedrockRuntimeInference(options)).rejects.toThrow("bedrock smoke rejected"); + expect(updateSandbox).not.toHaveBeenCalled(); + expect(log).not.toHaveBeenCalledWith(BEDROCK_SUCCESS_LOG); + }); }); diff --git a/src/lib/onboard/bedrock-runtime.ts b/src/lib/onboard/bedrock-runtime.ts index b90d2188ee6..6b8c36a09c1 100644 --- a/src/lib/onboard/bedrock-runtime.ts +++ b/src/lib/onboard/bedrock-runtime.ts @@ -139,7 +139,7 @@ export async function setupBedrockRuntimeInference( endpointUrl?: string | null; credentialEnv?: string | null; forceOpenAiLike?: boolean; - }) => void; + }) => void | Promise; ensureAdapter?: typeof ensureBedrockRuntimeAdapter; updateSandbox?: typeof registry.updateSandbox; } & BedrockRuntimeDependencies, @@ -213,7 +213,7 @@ export async function setupBedrockRuntimeInference( } options.verifyInferenceRoute(options.provider, options.model); - options.verifyOnboardInferenceSmoke({ + await options.verifyOnboardInferenceSmoke({ provider: options.provider, model: options.model, endpointUrl: adapter.localBaseUrl, diff --git a/src/lib/onboard/inference-providers/hermes.test.ts b/src/lib/onboard/inference-providers/hermes.test.ts index d0c2f7b6191..0f60fbe9ae2 100644 --- a/src/lib/onboard/inference-providers/hermes.test.ts +++ b/src/lib/onboard/inference-providers/hermes.test.ts @@ -59,6 +59,54 @@ function publicLookup() { return vi.fn(async () => [{ address: "8.8.8.8", family: 4 }]); } +function deferred(): { promise: Promise; resolve: () => void } { + let resolve!: () => void; + const promise = new Promise((done) => { + resolve = done; + }); + return { promise, resolve }; +} + +describe("setupHermesProviderInference smoke verification", () => { + it("waits for the smoke check before persisting or logging success (#3771)", async () => { + const smoke = deferred(); + const deps = makeDeps({ + isNonInteractive: vi.fn(() => true), + verifyOnboardInferenceSmoke: vi.fn(() => smoke.promise), + }); + + const setup = setupHermesProviderInference(makeArgs(null), deps as never); + + expect(deps.verifyOnboardInferenceSmoke).toHaveBeenCalledOnce(); + expect(deps.registry.updateSandbox).not.toHaveBeenCalled(); + expect(deps.log).not.toHaveBeenCalled(); + + smoke.resolve(); + await expect(setup).resolves.toEqual({ ok: true }); + + expect(deps.registry.updateSandbox).toHaveBeenCalledWith("alpha", { + model: "m", + provider: "p", + }); + expect(deps.log).toHaveBeenCalledWith(" ✓ Inference route set: p / m"); + }); + + it("rejects setup without persisting or logging success when the smoke check rejects (#3771)", async () => { + const smokeError = new Error("Hermes smoke failed"); + const deps = makeDeps({ + isNonInteractive: vi.fn(() => true), + verifyOnboardInferenceSmoke: vi.fn(() => Promise.reject(smokeError)), + }); + + await expect(setupHermesProviderInference(makeArgs(null), deps as never)).rejects.toThrow( + smokeError, + ); + + expect(deps.registry.updateSandbox).not.toHaveBeenCalled(); + expect(deps.log).not.toHaveBeenCalled(); + }); +}); + describe("setupHermesProviderInference SSRF guard (#6072)", () => { it("rejects loopback address", async () => { await expect( diff --git a/src/lib/onboard/inference-providers/hermes.ts b/src/lib/onboard/inference-providers/hermes.ts index 92c502101b7..3f5fb848e09 100644 --- a/src/lib/onboard/inference-providers/hermes.ts +++ b/src/lib/onboard/inference-providers/hermes.ts @@ -169,7 +169,12 @@ export async function setupHermesProviderInference( } verifyInferenceRoute(provider, model); - verifyOnboardInferenceSmoke({ provider, model, endpointUrl: resolvedEndpointUrl, credentialEnv }); + await verifyOnboardInferenceSmoke({ + provider, + model, + endpointUrl: resolvedEndpointUrl, + credentialEnv, + }); if (sandboxName) { registry.updateSandbox(sandboxName, { model, provider }); } diff --git a/src/lib/onboard/inference-providers/remote.ts b/src/lib/onboard/inference-providers/remote.ts index 1fe14c7ed15..0e9c4a71ea7 100644 --- a/src/lib/onboard/inference-providers/remote.ts +++ b/src/lib/onboard/inference-providers/remote.ts @@ -16,13 +16,13 @@ import { } from "./compatible-endpoint-gateway-route"; import type { RemoteProviderDeps, SetupInferenceResult } from "./types"; -const { probeOpenAiLikeEndpoint } = require("../../inference/onboard-probes") as { - probeOpenAiLikeEndpoint: ( +const { probeOpenAiLikeEndpointOptimized } = require("../../inference/onboard-probes") as { + probeOpenAiLikeEndpointOptimized: ( endpointUrl: string, model: string, apiKey: string, options?: Record, - ) => { ok: boolean; message?: string }; + ) => Promise<{ ok: boolean; message?: string }>; }; type StaleProviderReplaceResult = { ok: boolean; status?: number | null; message?: string }; @@ -220,7 +220,7 @@ export async function setupRemoteProviderInference( // Bedrock endpoints never reach here — the adapter branch above returns first. const useOpenAiSurface = provider === "compatible-anthropic-endpoint" && preferredInferenceApi === "openai-completions"; - const probeOpenAiSurface = deps.probeOpenAiLikeEndpoint ?? probeOpenAiLikeEndpoint; + const probeOpenAiSurface = deps.probeOpenAiLikeEndpoint ?? probeOpenAiLikeEndpointOptimized; // The concrete modules type their openshell runners independently; the deps // runner is call-compatible with both, so bridge the nominal mismatch here. const readProviderMetadata = @@ -268,10 +268,15 @@ export async function setupRemoteProviderInference( // route exercise the identical URL. const openAiSurfaceBaseUrl = getCompatibleAnthropicOpenAiSurfaceBaseUrl(resolvedEndpointUrl); - const surfaceProbe = probeOpenAiSurface(openAiSurfaceBaseUrl, model, credentialValue, { - skipResponsesProbe: true, - pinnedAddresses, - }); + const surfaceProbe = await probeOpenAiSurface( + openAiSurfaceBaseUrl, + model, + credentialValue, + { + skipResponsesProbe: true, + pinnedAddresses, + }, + ); if (!surfaceProbe.ok) { providerResult = { ok: false, diff --git a/src/lib/onboard/inference-providers/types.ts b/src/lib/onboard/inference-providers/types.ts index 7701612ba78..28fa3f8e509 100644 --- a/src/lib/onboard/inference-providers/types.ts +++ b/src/lib/onboard/inference-providers/types.ts @@ -66,7 +66,7 @@ export type VerifyOnboardInferenceSmoke = (input: { credentialEnv?: string | null; forceOpenAiLike?: boolean; pinnedAddresses?: readonly string[]; -}) => void; +}) => void | Promise; export type PromptValidationRecovery = ( label: string, @@ -109,7 +109,7 @@ export type RemoteProviderDeps = CommonDeps & { model: string, apiKey: string, options?: Record, - ) => { ok: boolean; message?: string }; + ) => { ok: boolean; message?: string } | Promise<{ ok: boolean; message?: string }>; readGatewayProviderMetadata?: ( name: string, runOpenshell: RunOpenshell, diff --git a/src/lib/onboard/inference-selection-validation.test.ts b/src/lib/onboard/inference-selection-validation.test.ts index 7476aa81788..c7777507f91 100644 --- a/src/lib/onboard/inference-selection-validation.test.ts +++ b/src/lib/onboard/inference-selection-validation.test.ts @@ -274,7 +274,7 @@ describe("inference selection validation", () => { }, ], })); - const probeOpenAiLikeEndpoint = vi.fn(() => ({ + const probeOpenAiLikeEndpoint = vi.fn(async () => ({ ok: true, api: "openai-completions", label: "Chat Completions API", diff --git a/src/lib/onboard/inference-selection-validation.ts b/src/lib/onboard/inference-selection-validation.ts index a50e019aaab..6badc85f4b4 100644 --- a/src/lib/onboard/inference-selection-validation.ts +++ b/src/lib/onboard/inference-selection-validation.ts @@ -4,7 +4,7 @@ import { getCredential } from "../credentials/store"; import { getCompatibleAnthropicOpenAiSurfaceBaseUrl } from "../inference/config"; -const { probeAnthropicEndpoint, probeOpenAiLikeEndpoint } = +const { probeAnthropicEndpoint, probeOpenAiLikeEndpointOptimized } = require("../inference/onboard-probes") as { probeAnthropicEndpoint( endpointUrl: string, @@ -12,14 +12,21 @@ const { probeAnthropicEndpoint, probeOpenAiLikeEndpoint } = apiKey: string | null | undefined, options?: { probeStreaming?: boolean; pinnedAddresses?: readonly string[] }, ): any; - probeOpenAiLikeEndpoint( + probeOpenAiLikeEndpointOptimized( endpointUrl: string, model: string, apiKey: string | null | undefined, options?: Record, - ): any; + ): Promise; }; +type OpenAiLikeProbe = ( + endpointUrl: string, + model: string, + apiKey: string | null | undefined, + options?: Record, +) => any | Promise; + import { assertEndpointResolvesPublic, type EndpointDnsLookupFn, @@ -44,7 +51,7 @@ export interface InferenceSelectionValidationDeps { agentProductName(): string; getCredential?: typeof getCredential; probeAnthropicEndpoint?: typeof probeAnthropicEndpoint; - probeOpenAiLikeEndpoint?: typeof probeOpenAiLikeEndpoint; + probeOpenAiLikeEndpoint?: OpenAiLikeProbe; /** Injectable DNS resolver for the custom-endpoint SSRF preflight (tests). */ resolveEndpointHost?: EndpointDnsLookupFn; promptValidationRecovery( @@ -105,7 +112,7 @@ export function createInferenceSelectionValidationHelpers( ): InferenceSelectionValidationHelpers { const resolveCredential = deps.getCredential ?? getCredential; const runAnthropicProbe = deps.probeAnthropicEndpoint ?? probeAnthropicEndpoint; - const runOpenAiLikeProbe = deps.probeOpenAiLikeEndpoint ?? probeOpenAiLikeEndpoint; + const runOpenAiLikeProbe = deps.probeOpenAiLikeEndpoint ?? probeOpenAiLikeEndpointOptimized; function exitNonInteractiveValidationFailure(): never { process.exitCode = 1; @@ -192,7 +199,7 @@ export function createInferenceSelectionValidationHelpers( } = {}, ): Promise { const apiKey = credentialEnv ? resolveCredential(credentialEnv) : ""; - const probe = runOpenAiLikeProbe(endpointUrl, model, apiKey, { + const probe = await runOpenAiLikeProbe(endpointUrl, model, apiKey, { ...options, calibrateTimeouts: true, }); @@ -270,7 +277,7 @@ export function createInferenceSelectionValidationHelpers( const apiKey = resolveCredential(credentialEnv); const reasoningEnabled = normalizeReasoningFlag(process.env.NEMOCLAW_REASONING) === "true"; // Reasoning-only compatible endpoints often reject Responses, tool-call, and streaming probes. - const probe = runOpenAiLikeProbe(endpointUrl, model, apiKey, { + const probe = await runOpenAiLikeProbe(endpointUrl, model, apiKey, { calibrateTimeouts: true, requireResponsesToolCalling: !reasoningEnabled, skipResponsesProbe: @@ -332,7 +339,7 @@ export function createInferenceSelectionValidationHelpers( // for duplicate/missing/out-of-order events (#6289). const probe = intendedApi === "openai-completions" - ? runOpenAiLikeProbe( + ? await runOpenAiLikeProbe( getCompatibleAnthropicOpenAiSurfaceBaseUrl(endpointUrl), model, apiKey, diff --git a/src/lib/onboard/openrouter-runtime.ts b/src/lib/onboard/openrouter-runtime.ts index 2a30fb4b7d0..3ca8dcc89c3 100644 --- a/src/lib/onboard/openrouter-runtime.ts +++ b/src/lib/onboard/openrouter-runtime.ts @@ -52,7 +52,7 @@ export async function setupOpenRouterRuntimeInference( endpointUrl?: string | null; credentialEnv?: string | null; forceOpenAiLike?: boolean; - }) => void; + }) => void | Promise; ensureAdapter?: typeof ensureOpenRouterRuntimeAdapter; updateSandbox?: typeof registry.updateSandbox; } & OpenRouterRuntimeDependencies, @@ -126,7 +126,7 @@ export async function setupOpenRouterRuntimeInference( if (options.skipHostInferenceSmoke === true || !options.credentialValue) { log(" Reusing existing gateway credential; skipping host inference smoke."); } else { - options.verifyOnboardInferenceSmoke({ + await options.verifyOnboardInferenceSmoke({ provider: options.provider, model: options.model, endpointUrl: adapter.localBaseUrl, diff --git a/src/lib/onboard/setup-inference-route-containment.test.ts b/src/lib/onboard/setup-inference-route-containment.test.ts index 0b719fe4fbb..792e934a999 100644 --- a/src/lib/onboard/setup-inference-route-containment.test.ts +++ b/src/lib/onboard/setup-inference-route-containment.test.ts @@ -205,8 +205,12 @@ describe("onboard shared gateway route containment", () => { expect(exitProcess).toHaveBeenCalledWith(1); }); - it("reserves a fresh route before smoke failure lets another setup mutate it (#6315)", async () => { + it("keeps a pending reservation while async smoke failure blocks another setup (#6315)", async () => { const reservations: SandboxEntry[] = []; + let rejectSmoke!: (reason?: unknown) => void; + const smokePending = new Promise((_resolve, reject) => { + rejectSmoke = reject; + }); let lockTail = Promise.resolve(); const withGatewayRouteMutationLock = async ( _gatewayName: string, @@ -226,11 +230,13 @@ describe("onboard shared gateway route containment", () => { }; const updateSandbox = vi.fn( (name: string, route: Parameters[1]) => { - reservations.push({ name, ...route }); + reservations.push({ name, pendingRouteReservation: true, ...route }); return true; }, ); const runOpenshell = vi.fn(() => ({ status: 0 })); + const verifyOnboardInferenceSmoke = vi.fn(() => smokePending); + const log = vi.fn(); const exitProcess = vi.fn((code: number): never => { throw new Error(`exit ${code}`); }); @@ -247,9 +253,7 @@ describe("onboard shared gateway route containment", () => { updateSandbox, upsertProvider: vi.fn(() => ({ ok: true })), verifyInferenceRoute: vi.fn(), - verifyOnboardInferenceSmoke: vi.fn(() => { - throw new Error("smoke failed"); - }), + verifyOnboardInferenceSmoke, isNonInteractive: () => true, hermesProviderAuth: { HERMES_PROVIDER_NAME: "hermes-provider" }, isRoutedInferenceProvider: () => true, @@ -270,15 +274,40 @@ describe("onboard shared gateway route containment", () => { }, redact: (value: string) => value, compactText: (value: string) => value, - log: vi.fn(), + log, error: vi.fn(), exitProcess, } as unknown as SetupInferenceDeps); - const results = await Promise.allSettled([ - setupInference("alpha", "model-a", "router-a", "http://router-a.test/v1", "ROUTER_KEY"), - setupInference("beta", "model-b", "router-b", "http://router-b.test/v1", "ROUTER_KEY"), + const firstSetup = setupInference( + "alpha", + "model-a", + "router-a", + "http://router-a.test/v1", + "ROUTER_KEY", + ); + await vi.waitFor(() => expect(verifyOnboardInferenceSmoke).toHaveBeenCalledOnce()); + expect(reservations).toEqual([ + expect.objectContaining({ + name: "alpha", + pendingRouteReservation: true, + provider: "router-a", + model: "model-a", + }), ]); + expect(log).not.toHaveBeenCalledWith(expect.stringContaining("Inference route set")); + + const secondSetup = setupInference( + "beta", + "model-b", + "router-b", + "http://router-b.test/v1", + "ROUTER_KEY", + ); + const resultsPending = Promise.allSettled([firstSetup, secondSetup]); + expect(runOpenshell).toHaveBeenCalledTimes(1); + rejectSmoke(new Error("smoke failed")); + const results = await resultsPending; expect(results).toEqual([ { status: "rejected", reason: expect.objectContaining({ message: "smoke failed" }) }, @@ -294,6 +323,8 @@ describe("onboard shared gateway route containment", () => { gatewayName: "nemoclaw", }); expect(reservations).toHaveLength(1); + expect(updateSandbox).toHaveBeenCalledOnce(); + expect(log).not.toHaveBeenCalledWith(expect.stringContaining("Inference route set")); expect(exitProcess).toHaveBeenCalledWith(1); }); diff --git a/src/lib/onboard/setup-inference.ts b/src/lib/onboard/setup-inference.ts index 485c1062339..e8038322985 100644 --- a/src/lib/onboard/setup-inference.ts +++ b/src/lib/onboard/setup-inference.ts @@ -443,7 +443,7 @@ export function createSetupInference( if (options.skipHostInferenceSmoke === true) deps.log(" Reusing existing gateway credential; skipping host inference smoke."); else - deps.verifyOnboardInferenceSmoke({ + await deps.verifyOnboardInferenceSmoke({ provider, model, endpointUrl, diff --git a/test/helpers/onboard-smoke-verifier-harness.ts b/test/helpers/onboard-smoke-verifier-harness.ts index a7370d672fc..78ef18d0143 100644 --- a/test/helpers/onboard-smoke-verifier-harness.ts +++ b/test/helpers/onboard-smoke-verifier-harness.ts @@ -14,9 +14,9 @@ type VerifyOnboardSmokeInvocation = { provider?: string; }; -export function runVerifyOnboardSmokeHarness( +export async function runVerifyOnboardSmokeHarness( invocations: VerifyOnboardSmokeInvocation[], -): SmokeVerifierHarnessCall[] { +): Promise { const harness = String.raw` const Module = require("node:module"); const originalLoad = Module._load; @@ -87,16 +87,20 @@ const { verifyOnboardInferenceSmoke } = require(process.env.PROBES_MODULE); const invocations = JSON.parse(process.env.SMOKE_INVOCATIONS || "[]"); console.log = (...args) => calls.push(["log", args.join(" ")]); -for (const invocation of invocations) { - verifyOnboardInferenceSmoke({ - endpointUrl: "https://api.example.com/v1", - model: "nous/test-model", - provider: "hermes-provider", - ...invocation, - }); -} - -process.stdout.write(JSON.stringify(calls)); +(async () => { + for (const invocation of invocations) { + await verifyOnboardInferenceSmoke({ + endpointUrl: "https://api.example.com/v1", + model: "nous/test-model", + provider: "hermes-provider", + ...invocation, + }); + } + process.stdout.write(JSON.stringify(calls)); +})().catch((error) => { + process.stderr.write(String(error && error.stack ? error.stack : error)); + process.exit(1); +}); `; const result = spawnSync(process.execPath, ["-e", harness], { cwd: process.cwd(), diff --git a/test/onboard-openrouter-inference.test.ts b/test/onboard-openrouter-inference.test.ts index 44ffa56347f..747dbf50db6 100644 --- a/test/onboard-openrouter-inference.test.ts +++ b/test/onboard-openrouter-inference.test.ts @@ -81,6 +81,44 @@ describe("OpenRouter onboarding inference setup", () => { }); }); + it("waits for host smoke verification before reporting OpenRouter success", async () => { + let finishSmoke: (() => void) | undefined; + const smokePending = new Promise((resolve) => { + finishSmoke = resolve; + }); + const log = vi.fn(); + + const setup = openrouterRuntimeOnboard.setupOpenRouterRuntimeInference({ + sandboxName: null, + provider: "openrouter-api", + model: "test-model", + credentialEnv: "OPENROUTER_API_KEY", + credentialValue: "sk-or-test", + isNonInteractive: () => true, + runOpenshell: () => ({ status: 0 }), + upsertProvider: () => ({ ok: true }), + verifyInferenceRoute: vi.fn(), + verifyOnboardInferenceSmoke: vi.fn(() => smokePending), + ensureAdapter: vi.fn(async () => ({ + baseUrl: "http://host.openshell.internal:11437/v1", + localBaseUrl: "http://127.0.0.1:11437/v1", + credentialEnv: "OPENROUTER_API_KEY", + logPath: "/tmp/openrouter-runtime-adapter.log", + })), + exitProcess: ((code: number) => { + throw new Error(`unexpected exit ${code}`); + }) as never, + error: vi.fn(), + log, + }); + + await vi.waitFor(() => expect(log).toHaveBeenCalledTimes(1)); + expect(log).not.toHaveBeenCalledWith(expect.stringContaining("Inference route set")); + finishSmoke?.(); + await setup; + expect(log).toHaveBeenCalledWith(" ✓ Inference route set: openrouter-api / test-model"); + }); + it("updates OpenRouter adapter config while reusing a gateway-held credential (#5826)", async () => { await withProcessEnv({ OPENROUTER_API_KEY: undefined }, async () => { const ensureAdapter = vi.fn(async () => ({ diff --git a/test/onboard-smoke-verifier.test.ts b/test/onboard-smoke-verifier.test.ts index 38a3906b737..8fc8d4ca98f 100644 --- a/test/onboard-smoke-verifier.test.ts +++ b/test/onboard-smoke-verifier.test.ts @@ -12,8 +12,8 @@ describe("Hermes onboard smoke verification", () => { expect(shouldSmokeOpenAiLikeOnboardRoute("openai-api")).toBe(true); }); - it("skips only the Hermes OAuth smoke path in the runtime verifier", () => { - const calls = runVerifyOnboardSmokeHarness([ + it("skips only the Hermes OAuth smoke path in the runtime verifier", async () => { + const calls = await runVerifyOnboardSmokeHarness([ { credentialEnv: "OPENAI_API_KEY" }, { credentialEnv: "NOUS_API_KEY" }, { credentialEnv: "OPENAI_API_KEY", forceOpenAiLike: true },