Luigit
repositories / pi-ext

pi-ext

bugabingas pi extensions

owned by admin

extensions/web/zai.ts

Raw
/**
 * Z.AI Coding Plan Web Search client.
 * Uses the Coding Plan MCP endpoint, not the paid REST Web Search API.
 * Docs: https://docs.z.ai/devpack/mcp/search-mcp-server
 */

import type { ExtensionContext } from "@earendil-works/pi-coding-agent";
import type { SearchResult } from "./constants.js";
import { callMcpTool, McpError, type McpToolResult } from "./mcp.js";
import {
	type SettingsContext,
	type ZaiSearchFreshness,
	zaiContentSize,
	zaiFreshness,
	zaiLocation,
	zaiMcpUrl,
} from "./settings.ts";

type ZaiRecencyFilter =
	| "oneDay"
	| "oneWeek"
	| "oneMonth"
	| "oneYear"
	| "noLimit";

interface ZaiSearchObjectResponse {
	title?: string;
	content?: string;
	link?: string;
	url?: string;
	media?: string;
	icon?: string;
	refer?: string;
	publish_date?: string;
	date?: string;
}

interface ZaiErrorResponse {
	code?: number | string;
	message?: string;
	error?: { code?: number | string; message?: string };
}

export class ZaiApiError extends Error {
	readonly code?: string;
	readonly status?: number;
	readonly rawMessage: string;

	constructor(
		message: string,
		code?: string,
		status?: number,
		rawMessage = "unknown error",
	) {
		super(message);
		this.name = "ZaiApiError";
		this.code = code;
		this.status = status;
		this.rawMessage = rawMessage;
	}
}

const ZAI_PROVIDER_NAMES = ["zai", "z-ai", "z.ai", "bigmodel"] as const;
const ZAI_API_KEY_ENV = ["ZAI_API_KEY", "Z_AI_API_KEY"] as const;

const FRESHNESS_TO_RECENCY: Record<ZaiSearchFreshness, ZaiRecencyFilter> = {
	day: "oneDay",
	week: "oneWeek",
	month: "oneMonth",
	year: "oneYear",
};

export async function resolveZaiApiKey(
	ctx?: ExtensionContext,
): Promise<string | undefined> {
	if (ctx) {
		for (const provider of ZAI_PROVIDER_NAMES) {
			const key = await ctx.modelRegistry.getApiKeyForProvider(provider);
			if (key) return key;
		}
	}

	for (const envName of ZAI_API_KEY_ENV) {
		const value = process.env[envName]?.trim();
		if (value) return value;
	}
	return undefined;
}

export async function isZaiAvailable(ctx?: ExtensionContext): Promise<boolean> {
	return !!(await resolveZaiApiKey(ctx));
}

function normalizeMcpUrl(input: string): string {
	const value = input.trim().replace(/\/+$/, "");
	try {
		const url = new URL(value);
		if (url.pathname.endsWith("/mcp")) return value;
		return `${url.origin}/api/mcp/web_search_prime/mcp`;
	} catch {}
	return value.endsWith("/mcp")
		? value
		: `${value}/api/mcp/web_search_prime/mcp`;
}

function makeZaiApiError(
	status: number | undefined,
	code: unknown,
	message: unknown,
): ZaiApiError {
	const codeText =
		code === undefined || code === null ? undefined : String(code);
	const rawMessage =
		typeof message === "string" && message ? message : "unknown error";
	const tag = codeText
		? ` (${codeText})`
		: status !== undefined && status >= 400
			? ` (HTTP ${status})`
			: "";
	return new ZaiApiError(
		`Z.AI API error${tag}: ${rawMessage}`,
		codeText,
		status,
		rawMessage,
	);
}

async function parseZaiError(response: Response): Promise<ZaiApiError> {
	const text = await response.text().catch(() => "");
	if (!text)
		return makeZaiApiError(
			response.status,
			undefined,
			response.statusText || "request failed",
		);
	try {
		const data = JSON.parse(text) as ZaiErrorResponse;
		return makeZaiApiError(
			response.status,
			data.error?.code ?? data.code,
			data.error?.message ?? data.message ?? response.statusText,
		);
	} catch {}
	return makeZaiApiError(
		response.status,
		undefined,
		`${response.statusText}: ${text.slice(0, 500)}`,
	);
}

function decodeSearchItems(text: string): ZaiSearchObjectResponse[] {
	let value: unknown = text.trim();

	for (let i = 0; i < 3 && typeof value === "string"; i++) {
		try {
			value = JSON.parse(value);
		} catch {
			break;
		}
	}

	if (Array.isArray(value)) return value as ZaiSearchObjectResponse[];
	if (value && typeof value === "object") {
		const record = value as {
			search_result?: ZaiSearchObjectResponse[];
			results?: ZaiSearchObjectResponse[];
		};
		if (Array.isArray(record.search_result)) return record.search_result;
		if (Array.isArray(record.results)) return record.results;
	}

	return [];
}

function mcpResultToSearchResults(result: McpToolResult): SearchResult[] {
	const text =
		result.content?.find((item) => item.type === "text" && item.text)?.text ??
		result.content?.find((item) => item.text)?.text;

	if (result.isError) {
		throw makeZaiApiError(200, undefined, text ?? "MCP tool returned an error");
	}
	if (!text) return [];

	return decodeSearchItems(text)
		.map((r) => ({
			title: r.title ?? r.link ?? r.url ?? "Untitled",
			url: r.link ?? r.url ?? "",
			snippet: r.content ?? r.media ?? "",
			date: r.publish_date ?? r.date,
		}))
		.filter((r) => r.url);
}

export async function zaiSearch(
	query: string,
	ctx: ExtensionContext | undefined,
	signal?: AbortSignal,
	options?: {
		count?: number;
		freshness?: ZaiSearchFreshness;
		timeoutMs?: number;
	},
): Promise<SearchResult[]> {
	const apiKey = await resolveZaiApiKey(ctx);
	if (!apiKey) {
		throw new Error(
			"Z.AI API key not configured. Set ZAI_API_KEY or configure a zai provider.",
		);
	}

	// Without a session, project settings stay unread.
	const settings: SettingsContext = ctx ?? {
		cwd: process.cwd(),
		isProjectTrusted: () => false,
	};
	const freshness = options?.freshness ?? zaiFreshness(settings);
	const toolArgs: Record<string, unknown> = {
		search_query: query,
		content_size: zaiContentSize(settings),
		location: zaiLocation(settings),
	};
	if (freshness)
		toolArgs.search_recency_filter = FRESHNESS_TO_RECENCY[freshness];

	const timeoutSignal = AbortSignal.timeout(options?.timeoutMs ?? 30_000);
	const requestSignal = signal
		? AbortSignal.any([signal, timeoutSignal])
		: timeoutSignal;
	const result = await callMcpTool({
		url: normalizeMcpUrl(zaiMcpUrl(settings)),
		protocolVersion: "2025-03-26",
		headers: { Authorization: `Bearer ${apiKey}` },
		signal: requestSignal,
		name: "web_search_prime",
		arguments: toolArgs,
	}).catch(async (error: unknown) => {
		if (!(error instanceof McpError)) throw error;
		if (error.response) throw await parseZaiError(error.response);
		throw makeZaiApiError(error.status, error.code, error.message);
	});

	const results = mcpResultToSearchResults(result);
	const count = Math.min(
		Math.max(Math.trunc(options?.count ?? results.length), 1),
		50,
	);
	return results.slice(0, count);
}