Skip to content

Built for people who want to own their automations. Join the waitlist for an invite.

Package listing

@kody/ai

src/complete.ts

573 lines · 14.4 KB · TypeScript
import {
	chatUrl,
	defaultSecretName,
	modelsUrl,
	providerCatalog,
	providerSecretUrl,
	requestHeaders,
} from './providers.ts'
import { getSettings } from './settings.ts'
import {
	asRecord,
	assertHttpsUrl,
	boundedMaxTokens,
	inputRecord,
	isPreviewDryRun,
	optionalString,
	parseProviderId,
	stringify,
	type ProviderId,
} from './validation.ts'

export type ChatMessage = {
	role: 'system' | 'user' | 'assistant' | 'tool'
	content?: string | null
	tool_calls?: unknown[]
	tool_call_id?: string
	name?: string
}

export type CompleteInput = {
	messages?: Array<ChatMessage>
	prompt?: string
	system?: string
	provider?: ProviderId | string
	model?: string
	modelId?: string
	baseUrl?: string
	cloudflareAccountId?: string
	cloudflareAiGatewayId?: string
	account?: string
	apiKeySecret?: string
	maxTokens?: number
	temperature?: number
	tools?: unknown[]
	toolChoice?: unknown
	responseFormat?: unknown
	dryRun?: boolean
}

export type ToolCall = {
	toolCallId: string
	toolName: string
	input: Record<string, unknown>
}

export type CompleteResult = {
	ok: boolean
	error?: string
	text: string
	reasoningText: string
	finishReason: string
	toolCalls: Array<ToolCall>
	provider: ProviderId
	model: string
	host: string
}

export type CompletePreview = {
	dryRun: true
	provider: ProviderId
	model: string
	host: string
	href: string
	apiKeySecret: string
	secretUrl: string
	maxTokens: number
	messageCount: number
	body: Record<string, unknown>
}

export type ResolvedComplete = {
	provider: ProviderId
	model: string
	href: string
	host: string
	apiKeySecret: string
	headers: Record<string, string>
	maxTokens: number
	body: Record<string, unknown>
	messages: Array<ChatMessage>
}

function parseArguments(value: unknown): Record<string, unknown> {
	if (!value) return {}
	if (typeof value === 'object' && !Array.isArray(value)) {
		return value as Record<string, unknown>
	}
	if (typeof value !== 'string') return {}
	try {
		const parsed = JSON.parse(value)
		return parsed && typeof parsed === 'object' && !Array.isArray(parsed)
			? (parsed as Record<string, unknown>)
			: {}
	} catch {
		return {}
	}
}

function normalizeMessages(input: CompleteInput): Array<ChatMessage> {
	const messages: Array<ChatMessage> = []
	if (input.system) messages.push({ role: 'system', content: input.system })
	if (Array.isArray(input.messages)) {
		for (const message of input.messages) {
			messages.push({
				role: message.role || 'user',
				content: message.content ?? '',
				tool_calls: message.tool_calls,
				tool_call_id: message.tool_call_id,
				name: message.name,
			})
		}
	}
	if (input.prompt) messages.push({ role: 'user', content: input.prompt })
	return messages
}

function openaiBody(
	model: string,
	messages: Array<ChatMessage>,
	maxTokens: number,
	input: CompleteInput,
): Record<string, unknown> {
	const body: Record<string, unknown> = {
		model,
		messages,
		max_tokens: maxTokens,
	}
	if (typeof input.temperature === 'number') body.temperature = input.temperature
	if (Array.isArray(input.tools) && input.tools.length > 0) {
		body.tools = input.tools
		body.tool_choice = input.toolChoice ?? 'auto'
	}
	if (input.responseFormat) body.response_format = input.responseFormat
	return body
}

function anthropicBody(
	model: string,
	messages: Array<ChatMessage>,
	maxTokens: number,
	input: CompleteInput,
): Record<string, unknown> {
	const system = messages
		.filter((message) => message.role === 'system')
		.map((message) => message.content ?? '')
		.filter(Boolean)
		.join('\n\n')
	const converted = messages
		.filter((message) => message.role !== 'system')
		.map((message) => {
			if (message.role === 'tool') {
				return {
					role: 'user',
					content: [
						{
							type: 'tool_result',
							tool_use_id: message.tool_call_id,
							content: message.content ?? '',
						},
					],
				}
			}
			if (message.role === 'assistant' && message.tool_calls) {
				const text = message.content
					? [{ type: 'text', text: message.content }]
					: []
				const tools = (message.tool_calls as Array<Record<string, unknown>>).map(
					(call) => {
						const fn = asRecord(call.function)
						return {
							type: 'tool_use',
							id: call.id,
							name: fn.name ?? call.name,
							input: parseArguments(fn.arguments ?? call.arguments),
						}
					},
				)
				return { role: 'assistant', content: [...text, ...tools] }
			}
			return { role: message.role, content: message.content ?? '' }
		})
	const body: Record<string, unknown> = {
		model,
		max_tokens: maxTokens,
		messages: converted,
	}
	if (system) body.system = system
	if (typeof input.temperature === 'number') body.temperature = input.temperature
	if (Array.isArray(input.tools) && input.tools.length > 0) {
		body.tools = input.tools.map((tool) => {
			const record = asRecord(tool)
			const fn = asRecord(record.function)
			if (record.type === 'function' || fn.name) {
				return {
					name: fn.name,
					description: fn.description,
					input_schema: fn.parameters ?? { type: 'object', properties: {} },
				}
			}
			return tool
		})
		body.tool_choice = input.toolChoice ?? { type: 'auto' }
	}
	return body
}

export async function resolveComplete(
	input: CompleteInput = {},
): Promise<ResolvedComplete> {
	const parsed = inputRecord(input ?? {})
	const settings = await getSettings()
	const provider =
		parseProviderId(optionalString(parsed, 'provider')) ??
		settings.provider ??
		'openai'
	const catalog = providerCatalog(provider)
	const model =
		optionalString(parsed, 'model') ??
		optionalString(parsed, 'modelId') ??
		settings.model ??
		catalog.defaultModel
	if (!model) {
		throw new Error(
			'model is required. Pass model or store a default in ./settings on your fork.',
		)
	}
	const baseUrlRaw =
		optionalString(parsed, 'baseUrl') ?? settings.baseUrl ?? undefined
	const baseUrl = baseUrlRaw ? assertHttpsUrl(baseUrlRaw, 'baseUrl') : undefined
	const cloudflareAccountId =
		optionalString(parsed, 'cloudflareAccountId') ??
		settings.cloudflareAccountId ??
		undefined
	const cloudflareAiGatewayId =
		optionalString(parsed, 'cloudflareAiGatewayId') ??
		settings.cloudflareAiGatewayId ??
		undefined
	const apiKeySecret = defaultSecretName(
		provider,
		optionalString(parsed, 'account'),
		optionalString(parsed, 'apiKeySecret') ?? settings.apiKeySecret ?? undefined,
	)
	const messages = normalizeMessages(input)
	const maxTokens = boundedMaxTokens(parsed.maxTokens)
	const endpoint = chatUrl({
		provider,
		baseUrl,
		cloudflareAccountId,
		cloudflareAiGatewayId,
	})
	const body =
		provider === 'anthropic'
			? anthropicBody(model, messages, maxTokens, input)
			: openaiBody(model, messages, maxTokens, input)
	return {
		provider,
		model,
		href: endpoint.href,
		host: endpoint.host,
		apiKeySecret,
		headers: requestHeaders(provider, apiKeySecret),
		maxTokens,
		body,
		messages,
	}
}

function choice(payload: unknown) {
	const choices = asRecord(payload).choices
	return Array.isArray(choices) && choices[0] && typeof choices[0] === 'object'
		? (choices[0] as Record<string, unknown>)
		: null
}

function openaiMessage(payload: unknown) {
	const msg = choice(payload)?.message
	return msg && typeof msg === 'object' ? (msg as Record<string, unknown>) : null
}

function textFromOpenAI(payload: unknown) {
	const msg = openaiMessage(payload)
	if (typeof msg?.content === 'string') return msg.content
	return ''
}

function reasoningFromOpenAI(payload: unknown) {
	const msg = openaiMessage(payload)
	if (typeof msg?.reasoning_content === 'string') return msg.reasoning_content
	if (typeof msg?.reasoning === 'string') return msg.reasoning
	return ''
}

function finishFromOpenAI(payload: unknown) {
	const reason = choice(payload)?.finish_reason
	return typeof reason === 'string' ? reason : 'stop'
}

function toolCallsFromOpenAI(payload: unknown): Array<ToolCall> {
	const msg = openaiMessage(payload)
	const calls = Array.isArray(msg?.tool_calls) ? msg.tool_calls : []
	return calls.map((call, index) => {
		const record = asRecord(call)
		const fn = asRecord(record.function)
		return {
			toolCallId:
				typeof record.id === 'string' && record.id
					? record.id
					: 'tool_' + index,
			toolName: typeof fn.name === 'string' ? fn.name : String(record.name ?? ''),
			input: parseArguments(fn.arguments ?? record.arguments),
		}
	})
}

function textFromAnthropic(payload: unknown) {
	const content = asRecord(payload).content
	if (!Array.isArray(content)) {
		return typeof asRecord(payload).content === 'string'
			? String(asRecord(payload).content)
			: ''
	}
	return content
		.map((block) => {
			const record = asRecord(block)
			return record.type === 'text' && typeof record.text === 'string'
				? record.text
				: ''
		})
		.filter(Boolean)
		.join('\n')
}

function toolCallsFromAnthropic(payload: unknown): Array<ToolCall> {
	const content = asRecord(payload).content
	if (!Array.isArray(content)) return []
	const calls: Array<ToolCall> = []
	for (const block of content) {
		const record = asRecord(block)
		if (record.type !== 'tool_use') continue
		calls.push({
			toolCallId: typeof record.id === 'string' ? record.id : 'tool_' + calls.length,
			toolName: typeof record.name === 'string' ? record.name : '',
			input: asRecord(record.input),
		})
	}
	return calls
}

function normalizePayload(
	provider: ProviderId,
	payload: unknown,
): Pick<CompleteResult, 'text' | 'reasoningText' | 'finishReason' | 'toolCalls'> {
	if (provider === 'anthropic') {
		const stop = asRecord(payload).stop_reason
		return {
			text: textFromAnthropic(payload),
			reasoningText: '',
			finishReason: typeof stop === 'string' ? stop : 'stop',
			toolCalls: toolCallsFromAnthropic(payload),
		}
	}
	return {
		text: textFromOpenAI(payload),
		reasoningText: reasoningFromOpenAI(payload),
		finishReason: finishFromOpenAI(payload),
		toolCalls: toolCallsFromOpenAI(payload),
	}
}

export function completePreview(resolved: ResolvedComplete): CompletePreview {
	return {
		dryRun: true,
		provider: resolved.provider,
		model: resolved.model,
		host: resolved.host,
		href: resolved.href,
		apiKeySecret: resolved.apiKeySecret,
		secretUrl: providerSecretUrl(
			resolved.provider,
			resolved.apiKeySecret,
			resolved.host,
		),
		maxTokens: resolved.maxTokens,
		messageCount: resolved.messages.length,
		body: resolved.body,
	}
}

async function postComplete(resolved: ResolvedComplete): Promise<CompleteResult> {
	const response = await fetch(resolved.href, {
		method: 'POST',
		headers: resolved.headers,
		body: JSON.stringify(resolved.body),
	})
	const raw = await response.text()
	let payload: unknown = raw
	try {
		payload = raw ? JSON.parse(raw) : null
	} catch {
		payload = raw
	}
	if (!response.ok) {
		throw new Error(
			resolved.provider +
				' chat request failed: ' +
				response.status +
				' ' +
				stringify(payload) +
				'. Save ' +
				resolved.apiKeySecret +
				' and approve host ' +
				resolved.host +
				': ' +
				providerSecretUrl(resolved.provider, resolved.apiKeySecret, resolved.host),
		)
	}
	const normalized = normalizePayload(resolved.provider, payload)
	return {
		ok: true,
		...normalized,
		provider: resolved.provider,
		model: resolved.model,
		host: resolved.host,
	}
}

/**
 * One provider-agnostic chat completion. Pass `dryRun: true` to preview.
 *
 * @example
 * import { generateText } from 'kody:@kody/ai/complete'
 * const preview = await generateText({
 *   provider: 'openai',
 *   messages: [{ role: 'user', content: 'Hello' }],
 *   dryRun: true,
 * })
 */
export async function generateText(
	input: CompleteInput = {},
): Promise<CompleteResult | CompletePreview> {
	const resolved = await resolveComplete(input)
	if (resolved.messages.length === 0) {
		throw new Error('messages or prompt is required.')
	}
	if (isPreviewDryRun(input)) return completePreview(resolved)
	try {
		return await postComplete(resolved)
	} catch (error) {
		return {
			ok: false,
			error: error instanceof Error ? error.message : String(error),
			text: '',
			reasoningText: '',
			finishReason: 'error',
			toolCalls: [],
			provider: resolved.provider,
			model: resolved.model,
			host: resolved.host,
		}
	}
}

export async function resolveProvider(input: CompleteInput = {}) {
	const parsed = inputRecord(input ?? {})
	const settings = await getSettings()
	const provider =
		parseProviderId(optionalString(parsed, 'provider')) ??
		settings.provider ??
		'openai'
	const catalog = providerCatalog(provider)
	const model =
		optionalString(parsed, 'model') ??
		optionalString(parsed, 'modelId') ??
		settings.model ??
		catalog.defaultModel
	const baseUrlRaw =
		optionalString(parsed, 'baseUrl') ?? settings.baseUrl ?? undefined
	const baseUrl = baseUrlRaw ? assertHttpsUrl(baseUrlRaw, 'baseUrl') : undefined
	const apiKeySecret = defaultSecretName(
		provider,
		optionalString(parsed, 'account'),
		optionalString(parsed, 'apiKeySecret') ?? settings.apiKeySecret ?? undefined,
	)
	return {
		provider,
		model,
		baseUrl,
		apiKeySecret,
		cloudflareAccountId:
			optionalString(parsed, 'cloudflareAccountId') ??
			settings.cloudflareAccountId ??
			undefined,
	}
}

export async function listModels(input: CompleteInput = {}) {
	const resolved = await resolveProvider(input)
	const models = modelsUrl({
		provider: resolved.provider,
		baseUrl: resolved.baseUrl,
		cloudflareAccountId: resolved.cloudflareAccountId,
	})
	if (!models) {
		return {
			ok: false,
			error:
				'Cannot list models for this provider until baseUrl or cloudflareAccountId is set on your fork.',
			provider: resolved.provider,
		}
	}
	if (isPreviewDryRun(input)) {
		return {
			dryRun: true as const,
			provider: resolved.provider,
			href: models.href,
			host: models.host,
			apiKeySecret: resolved.apiKeySecret,
		}
	}
	const response = await fetch(models.href, {
		method: 'GET',
		headers: requestHeaders(resolved.provider, resolved.apiKeySecret),
	})
	const raw = await response.text()
	let payload: unknown = raw
	try {
		payload = raw ? JSON.parse(raw) : null
	} catch {
		payload = raw
	}
	if (!response.ok) {
		throw new Error(
			'Models list failed: ' +
				response.status +
				' ' +
				stringify(payload) +
				'. ' +
				providerSecretUrl(resolved.provider, resolved.apiKeySecret, models.host),
		)
	}
	const record = asRecord(payload)
	const data = Array.isArray(record.data)
		? record.data
		: Array.isArray(record.result)
			? record.result
			: []
	const ids = data
		.slice(0, 8)
		.map((item) => {
			const row = asRecord(item)
			return typeof row.id === 'string' ? row.id : null
		})
		.filter((id): id is string => Boolean(id))
	return {
		ok: true,
		provider: resolved.provider,
		host: models.host,
		modelCount: data.length,
		sampleIds: ids,
	}
}

export default generateText