mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:40:50 +08:00
chore: update dependencies and improve chart components
- Updated `foreign-types` to version 0.5.0 and added new dependencies in `Cargo.lock`. - Removed unused dependencies `chart.js` and `react-chartjs-2` from `package.json` and `package-lock.json`. - Enhanced `CacheHitRateChart` to use `EChart` for rendering, improving performance and visual fidelity. - Adjusted styles for `CacheHitRateChart` to ensure proper layout and responsiveness. - Refactored API query construction in `api.ts` for better readability. - Updated `HomeMetrics` to support refresh functionality for dynamic data updates.
This commit is contained in:
Generated
+100
-2
@@ -926,7 +926,7 @@ dependencies = [
|
||||
"bitflags 2.13.1",
|
||||
"core-foundation 0.10.1",
|
||||
"core-graphics-types",
|
||||
"foreign-types",
|
||||
"foreign-types 0.5.0",
|
||||
"libc",
|
||||
]
|
||||
|
||||
@@ -1905,6 +1905,15 @@ version = "0.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb"
|
||||
|
||||
[[package]]
|
||||
name = "foreign-types"
|
||||
version = "0.3.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f6f339eb8adc052cd2ca78910fda869aefa38d22d5cb648e6485e4d3fc06f3b1"
|
||||
dependencies = [
|
||||
"foreign-types-shared 0.1.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "foreign-types"
|
||||
version = "0.5.0"
|
||||
@@ -1912,7 +1921,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d737d9aa519fb7b749cbc3b962edcf310a8dd1f4b67c91c4f83975dbdd17d965"
|
||||
dependencies = [
|
||||
"foreign-types-macros",
|
||||
"foreign-types-shared",
|
||||
"foreign-types-shared 0.3.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -1926,6 +1935,12 @@ dependencies = [
|
||||
"syn 3.0.3",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "foreign-types-shared"
|
||||
version = "0.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "00b0228411908ca8685dba7fc2cdd70ec9990a6e753e89b6ac91a84c40fbaf4b"
|
||||
|
||||
[[package]]
|
||||
name = "foreign-types-shared"
|
||||
version = "0.3.1"
|
||||
@@ -2727,6 +2742,22 @@ dependencies = [
|
||||
"webpki-roots 1.0.9",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hyper-tls"
|
||||
version = "0.6.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "70206fc6890eaca9fde8a0bf71caa2ddfc9fe045ac9e5c70df101a7dbde866e0"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"http-body-util",
|
||||
"hyper",
|
||||
"hyper-util",
|
||||
"native-tls",
|
||||
"tokio",
|
||||
"tokio-native-tls",
|
||||
"tower-service",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hyper-tungstenite"
|
||||
version = "0.30.0"
|
||||
@@ -3640,6 +3671,23 @@ version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1d87ecb2933e8aeadb3e3a02b828fed80a7528047e68b4f424523a0981a3a084"
|
||||
|
||||
[[package]]
|
||||
name = "native-tls"
|
||||
version = "0.2.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "465500e14ea162429d264d44189adc38b199b62b1c21eea9f69e4b73cb03bbf2"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"log",
|
||||
"openssl",
|
||||
"openssl-probe",
|
||||
"openssl-sys",
|
||||
"schannel",
|
||||
"security-framework",
|
||||
"security-framework-sys",
|
||||
"tempfile",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ndk"
|
||||
version = "0.9.0"
|
||||
@@ -4038,12 +4086,49 @@ dependencies = [
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "openssl"
|
||||
version = "0.10.81"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "77823a27f0babb03091cb9ed9ef80af3b39dbc82f97e8fa530374b7dafd87a45"
|
||||
dependencies = [
|
||||
"bitflags 2.13.1",
|
||||
"cfg-if",
|
||||
"foreign-types 0.3.2",
|
||||
"libc",
|
||||
"openssl-macros",
|
||||
"openssl-sys",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "openssl-macros"
|
||||
version = "0.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a948666b637a0f465e8564c73e89d4dde00d72d4d473cc972f390fc3dcee7d9c"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "openssl-probe"
|
||||
version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe"
|
||||
|
||||
[[package]]
|
||||
name = "openssl-sys"
|
||||
version = "0.9.117"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b47e7e6bb2c38cd930d25a23b40fa52e068c10e85f3e03a7f5ba5aaca5713695"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"libc",
|
||||
"pkg-config",
|
||||
"vcpkg",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "option-ext"
|
||||
version = "0.2.0"
|
||||
@@ -5043,9 +5128,11 @@ dependencies = [
|
||||
"http-body-util",
|
||||
"hyper",
|
||||
"hyper-rustls",
|
||||
"hyper-tls",
|
||||
"hyper-util",
|
||||
"js-sys",
|
||||
"log",
|
||||
"native-tls",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"quinn",
|
||||
@@ -5056,6 +5143,7 @@ dependencies = [
|
||||
"serde_urlencoded",
|
||||
"sync_wrapper",
|
||||
"tokio",
|
||||
"tokio-native-tls",
|
||||
"tokio-rustls",
|
||||
"tokio-util",
|
||||
"tower",
|
||||
@@ -6924,6 +7012,16 @@ dependencies = [
|
||||
"syn 3.0.3",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-native-tls"
|
||||
version = "0.3.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "bbae76ab933c85776efabc971569dd6119c580d8f5d448769dec1764bf796ef2"
|
||||
dependencies = [
|
||||
"native-tls",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tokio-rustls"
|
||||
version = "0.26.4"
|
||||
|
||||
Generated
-30
@@ -18,13 +18,11 @@
|
||||
"@tauri-apps/plugin-opener": "^2.5.4",
|
||||
"@tauri-apps/plugin-process": "^2.3.1",
|
||||
"@tauri-apps/plugin-updater": "^2.10.1",
|
||||
"chart.js": "^4.5.1",
|
||||
"echarts": "^6.1.0",
|
||||
"keepalive-for-react": "^5.0.11",
|
||||
"keepalive-for-react-router": "^5.0.7",
|
||||
"monaco-editor": "0.56.0",
|
||||
"react": "^19.2.8",
|
||||
"react-chartjs-2": "^5.3.1",
|
||||
"react-dom": "^19.2.8",
|
||||
"react-router-dom": "^7.18.2",
|
||||
"sortablejs": "^1.15.7",
|
||||
@@ -427,12 +425,6 @@
|
||||
"@jridgewell/sourcemap-codec": "^1.4.14"
|
||||
}
|
||||
},
|
||||
"node_modules/@kurkle/color": {
|
||||
"version": "0.3.4",
|
||||
"resolved": "https://registry.npmjs.org/@kurkle/color/-/color-0.3.4.tgz",
|
||||
"integrity": "sha512-M5UknZPHRu3DEDWoipU6sE8PdkZ6Z/S+v4dD+Ke8IaNlpdSQah50lz1KtcFBa2vsdOnwbbnxJwVM4wty6udA5w==",
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@oxc-project/types": {
|
||||
"version": "0.142.0",
|
||||
"resolved": "https://registry.npmjs.org/@oxc-project/types/-/types-0.142.0.tgz",
|
||||
@@ -1817,18 +1809,6 @@
|
||||
"node": ">=8"
|
||||
}
|
||||
},
|
||||
"node_modules/chart.js": {
|
||||
"version": "4.5.1",
|
||||
"resolved": "https://registry.npmjs.org/chart.js/-/chart.js-4.5.1.tgz",
|
||||
"integrity": "sha512-GIjfiT9dbmHRiYi6Nl2yFCq7kkwdkp1W/lp2J99rX0yo9tgJGn3lKQATztIjb5tVtevcBtIdICNWqlq5+E8/Pw==",
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@kurkle/color": "^0.3.0"
|
||||
},
|
||||
"engines": {
|
||||
"pnpm": ">=8"
|
||||
}
|
||||
},
|
||||
"node_modules/chokidar": {
|
||||
"version": "5.0.0",
|
||||
"resolved": "https://registry.npmjs.org/chokidar/-/chokidar-5.0.0.tgz",
|
||||
@@ -2998,16 +2978,6 @@
|
||||
"node": ">=0.10.0"
|
||||
}
|
||||
},
|
||||
"node_modules/react-chartjs-2": {
|
||||
"version": "5.3.1",
|
||||
"resolved": "https://registry.npmjs.org/react-chartjs-2/-/react-chartjs-2-5.3.1.tgz",
|
||||
"integrity": "sha512-h5IPXKg9EXpjoBzUfyWJvllMjG2mQ4EiuHQFhms/AjUm0XSZHhyRy2xVmLXHKrtcdrPO4mnGqRtYoD0vp95A0A==",
|
||||
"license": "MIT",
|
||||
"peerDependencies": {
|
||||
"chart.js": "^4.1.1",
|
||||
"react": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0"
|
||||
}
|
||||
},
|
||||
"node_modules/react-dom": {
|
||||
"version": "19.2.8",
|
||||
"resolved": "https://registry.npmjs.org/react-dom/-/react-dom-19.2.8.tgz",
|
||||
|
||||
@@ -28,13 +28,11 @@
|
||||
"@tauri-apps/plugin-opener": "^2.5.4",
|
||||
"@tauri-apps/plugin-process": "^2.3.1",
|
||||
"@tauri-apps/plugin-updater": "^2.10.1",
|
||||
"chart.js": "^4.5.1",
|
||||
"echarts": "^6.1.0",
|
||||
"keepalive-for-react": "^5.0.11",
|
||||
"keepalive-for-react-router": "^5.0.7",
|
||||
"monaco-editor": "0.56.0",
|
||||
"react": "^19.2.8",
|
||||
"react-chartjs-2": "^5.3.1",
|
||||
"react-dom": "^19.2.8",
|
||||
"react-router-dom": "^7.18.2",
|
||||
"sortablejs": "^1.15.7",
|
||||
|
||||
@@ -300,8 +300,8 @@ export const api = {
|
||||
params.set("end_ms", String(filter.endMs));
|
||||
if (filter.modelHashes?.length) params.set("model_hashes", JSON.stringify(filter.modelHashes));
|
||||
}
|
||||
const query = params.size ? `?${params}` : "";
|
||||
return request<Overview>(`/overview${query}`);
|
||||
const query = params.toString();
|
||||
return request<Overview>(`/overview${query ? `?${query}` : ""}`);
|
||||
},
|
||||
cursorHarness: () => request<CursorHarnessStatus>("/harness/cursor/status"),
|
||||
initializeCursorCa: () => request<CursorHarnessStatus>("/harness/cursor/ca/initialize", { method: "POST" }),
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import { BarChart, LineChart } from "echarts/charts";
|
||||
import { BarChart, GaugeChart, LineChart } from "echarts/charts";
|
||||
import { GridComponent, LegendComponent, MarkLineComponent, TooltipComponent } from "echarts/components";
|
||||
import { getInstanceByDom, init, use, type EChartsCoreOption } from "echarts/core";
|
||||
import { CanvasRenderer } from "echarts/renderers";
|
||||
import { useEffect, useRef, type MouseEventHandler } from "react";
|
||||
import styles from "./EChart.module.scss";
|
||||
|
||||
use([BarChart, LineChart, GridComponent, LegendComponent, MarkLineComponent, TooltipComponent, CanvasRenderer]);
|
||||
use([BarChart, GaugeChart, LineChart, GridComponent, LegendComponent, MarkLineComponent, TooltipComponent, CanvasRenderer]);
|
||||
|
||||
type EChartProps = {
|
||||
option: EChartsCoreOption;
|
||||
|
||||
@@ -1,18 +1,22 @@
|
||||
@use "../../styles/typography" as type;
|
||||
|
||||
.root {
|
||||
--cache-hit-track-color: color-mix(in srgb, var(--vscode-foreground) 12%, transparent);
|
||||
--cache-hit-value-color: var(--vscode-gitDecoration-addedResourceForeground, #4ade80);
|
||||
position: relative;
|
||||
width: 132px;
|
||||
width: 200px;
|
||||
max-width: 100%;
|
||||
height: 82px;
|
||||
overflow: hidden;
|
||||
align-self: center;
|
||||
flex: 0 0 auto;
|
||||
}
|
||||
|
||||
.canvas {
|
||||
width: 100% !important;
|
||||
height: 100% !important;
|
||||
.chart {
|
||||
position: absolute;
|
||||
top: 0;
|
||||
left: 0;
|
||||
width: 100%;
|
||||
height: 140px;
|
||||
pointer-events: none;
|
||||
}
|
||||
|
||||
.label {
|
||||
|
||||
@@ -1,75 +1,62 @@
|
||||
import { ArcElement, Chart as ChartJS, Tooltip, type ChartOptions, type ScriptableContext } from "chart.js";
|
||||
import { useMemo } from "react";
|
||||
import { Doughnut } from "react-chartjs-2";
|
||||
import type { EChartsCoreOption } from "echarts/core";
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
import { EChart } from "../charts/EChart";
|
||||
import styles from "./CacheHitRateChart.module.scss";
|
||||
|
||||
ChartJS.register(ArcElement, Tooltip);
|
||||
const valueColor = "#40c463";
|
||||
const trackColor = "rgba(139, 148, 158, 0.20)";
|
||||
|
||||
type SegmentRadius = number | {
|
||||
outerStart: number;
|
||||
outerEnd: number;
|
||||
innerStart: number;
|
||||
innerEnd: number;
|
||||
};
|
||||
|
||||
function chartColor(context: ScriptableContext<"doughnut">) {
|
||||
const styles = getComputedStyle(context.chart.canvas);
|
||||
const variable = context.dataIndex === 0 ? "--cache-hit-value-color" : "--cache-hit-track-color";
|
||||
return styles.getPropertyValue(variable).trim();
|
||||
}
|
||||
|
||||
function segmentBorderRadius(percentage: number, dataIndex: number): SegmentRadius {
|
||||
const radius = 5;
|
||||
|
||||
if (percentage <= 0) {
|
||||
return dataIndex === 1
|
||||
? { outerStart: radius, outerEnd: radius, innerStart: radius, innerEnd: radius }
|
||||
: 0;
|
||||
}
|
||||
|
||||
if (percentage >= 100) {
|
||||
return dataIndex === 0
|
||||
? { outerStart: radius, outerEnd: radius, innerStart: radius, innerEnd: radius }
|
||||
: 0;
|
||||
}
|
||||
|
||||
return dataIndex === 0
|
||||
? { outerStart: radius, outerEnd: 0, innerStart: radius, innerEnd: 0 }
|
||||
: { outerStart: 0, outerEnd: radius, innerStart: 0, innerEnd: radius };
|
||||
}
|
||||
|
||||
const options: ChartOptions<"doughnut"> = {
|
||||
responsive: true,
|
||||
maintainAspectRatio: false,
|
||||
cutout: "82%",
|
||||
rotation: -90,
|
||||
circumference: 180,
|
||||
animation: { duration: 450 },
|
||||
events: [],
|
||||
plugins: {
|
||||
legend: { display: false },
|
||||
tooltip: { enabled: false },
|
||||
},
|
||||
};
|
||||
|
||||
export function CacheHitRateChart({ rate }: { rate: number }) {
|
||||
export function CacheHitRateChart({ rate, animationKey = 0 }: { rate: number; animationKey?: number }) {
|
||||
const finiteRate = Number.isFinite(rate) ? rate : 0;
|
||||
const percentage = Math.max(0, Math.min(100, finiteRate * 100));
|
||||
const [displayedPercentage, setDisplayedPercentage] = useState(0);
|
||||
const label = Number.isFinite(rate) ? `${percentage.toFixed(2)}%` : "--";
|
||||
const data = useMemo(() => ({
|
||||
labels: [t("命中"), t("未命中")],
|
||||
datasets: [{
|
||||
data: [percentage, Math.max(0, 100 - percentage)],
|
||||
backgroundColor: chartColor,
|
||||
borderWidth: 0,
|
||||
hoverBorderWidth: 0,
|
||||
selfJoin: false,
|
||||
borderRadius: (context: ScriptableContext<"doughnut">) => segmentBorderRadius(percentage, context.dataIndex),
|
||||
|
||||
useEffect(() => {
|
||||
setDisplayedPercentage(0);
|
||||
let frame = requestAnimationFrame(() => {
|
||||
frame = requestAnimationFrame(() => setDisplayedPercentage(percentage));
|
||||
});
|
||||
return () => cancelAnimationFrame(frame);
|
||||
}, [animationKey, percentage]);
|
||||
|
||||
const option = useMemo<EChartsCoreOption>(() => ({
|
||||
animationDuration: 0,
|
||||
animationDurationUpdate: displayedPercentage > 0 ? 1_000 : 0,
|
||||
animationEasing: "cubicOut",
|
||||
animationEasingUpdate: "cubicOut",
|
||||
series: [{
|
||||
type: "gauge",
|
||||
min: 0,
|
||||
max: 100,
|
||||
startAngle: 180,
|
||||
endAngle: 0,
|
||||
center: ["50%", "50%"],
|
||||
radius: "90%",
|
||||
silent: true,
|
||||
pointer: { show: false },
|
||||
progress: {
|
||||
show: true,
|
||||
roundCap: true,
|
||||
width: 11,
|
||||
itemStyle: { color: displayedPercentage > 0 ? valueColor : "transparent" },
|
||||
},
|
||||
axisLine: {
|
||||
roundCap: true,
|
||||
lineStyle: { width: 11, color: [[1, trackColor]] },
|
||||
},
|
||||
axisTick: { show: false },
|
||||
splitLine: { show: false },
|
||||
axisLabel: { show: false },
|
||||
anchor: { show: false },
|
||||
title: { show: false },
|
||||
detail: { show: false },
|
||||
data: [{ value: displayedPercentage }],
|
||||
}],
|
||||
}), [percentage]);
|
||||
}), [displayedPercentage]);
|
||||
|
||||
return <div className={styles.root} role="img" aria-label={t("缓存命中率 {rate}", { rate: label })}>
|
||||
<Doughnut className={styles.canvas} data={data} options={options} />
|
||||
<EChart className={styles.chart} option={option} />
|
||||
<div className={styles.label}>{label}</div>
|
||||
</div>;
|
||||
}
|
||||
|
||||
@@ -71,7 +71,7 @@ function InfoTooltip({ content }: { content: string }) {
|
||||
</>;
|
||||
}
|
||||
|
||||
export function HomeMetrics({ data }: { data: HomeMetricsData }) {
|
||||
export function HomeMetrics({ data, refreshVersion = 0 }: { data: HomeMetricsData; refreshVersion?: number }) {
|
||||
const inputTokens = Math.max(0, data.promptTokens - data.cacheReadTokens - data.cacheWriteTokens);
|
||||
const outputTokens = Math.max(0, data.tokenUsage - data.promptTokens);
|
||||
const defaultCacheHitRate = calculateRate(data.cacheReadTokens, data.cacheReadTokens + inputTokens);
|
||||
@@ -148,7 +148,7 @@ export function HomeMetrics({ data }: { data: HomeMetricsData }) {
|
||||
<section className={styles.root} aria-label={t("调用统计")}>
|
||||
<article className={styles.metric}>
|
||||
<div className={styles.label}>{t("缓存命中率")}<InfoTooltip content={cacheTooltip} /></div>
|
||||
<CacheHitRateChart rate={defaultCacheHitRate ?? 0} />
|
||||
<CacheHitRateChart rate={defaultCacheHitRate ?? 0} animationKey={refreshVersion} />
|
||||
</article>
|
||||
<article className={styles.metric}>
|
||||
<div className={styles.label}>{t("LLM 调用")}<InfoTooltip content={callsTooltip} /></div>
|
||||
|
||||
@@ -118,7 +118,7 @@ export function HomePage() {
|
||||
{
|
||||
key: "metrics",
|
||||
estimatedHeight: 130,
|
||||
content: <HomeMetrics data={metrics} />,
|
||||
content: <HomeMetrics data={metrics} refreshVersion={refreshVersion} />,
|
||||
},
|
||||
|
||||
{
|
||||
|
||||
@@ -29,24 +29,37 @@ impl ModelAssets {
|
||||
}
|
||||
|
||||
pub fn ensure(cache_root: &Path) -> Result<Self> {
|
||||
let client = reqwest::blocking::Client::builder()
|
||||
.build()
|
||||
.map_err(|error| Error::ModelAsset(error.to_string()))?;
|
||||
Self::ensure_with_client(cache_root, &client)
|
||||
}
|
||||
|
||||
pub fn ensure_with_client(
|
||||
cache_root: &Path,
|
||||
client: &reqwest::blocking::Client,
|
||||
) -> Result<Self> {
|
||||
let directory = cache_root.join("models/potion-code-16M-v2");
|
||||
fs::create_dir_all(&directory).map_err(|error| Error::io(&directory, error))?;
|
||||
let model = Self::model_path(cache_root);
|
||||
let tokenizer = directory.join("tokenizer.json");
|
||||
ensure_asset(&model, MODEL_URL, MODEL_SHA256)?;
|
||||
ensure_asset(&tokenizer, TOKENIZER_URL, TOKENIZER_SHA256)?;
|
||||
ensure_asset(client, &model, MODEL_URL, MODEL_SHA256)?;
|
||||
ensure_asset(client, &tokenizer, TOKENIZER_URL, TOKENIZER_SHA256)?;
|
||||
Ok(Self { model, tokenizer })
|
||||
}
|
||||
}
|
||||
|
||||
fn ensure_asset(path: &Path, url: &str, expected: &str) -> Result<()> {
|
||||
fn ensure_asset(
|
||||
client: &reqwest::blocking::Client,
|
||||
path: &Path,
|
||||
url: &str,
|
||||
expected: &str,
|
||||
) -> Result<()> {
|
||||
if path.is_file() && digest(path)? == expected {
|
||||
return Ok(());
|
||||
}
|
||||
let temporary = path.with_extension(format!("tmp-{}", std::process::id()));
|
||||
let response = reqwest::blocking::Client::builder()
|
||||
.build()
|
||||
.map_err(|error| Error::ModelAsset(error.to_string()))?
|
||||
let response = client
|
||||
.get(url)
|
||||
.send()
|
||||
.and_then(reqwest::blocking::Response::error_for_status)
|
||||
|
||||
@@ -35,12 +35,22 @@ static EMBEDDERS: LazyLock<Mutex<HashMap<PathBuf, Arc<StaticEmbedder>>>> =
|
||||
|
||||
impl SearchEngine {
|
||||
pub fn load_default(config: SembleConfig) -> Result<Self> {
|
||||
let client = reqwest::blocking::Client::builder()
|
||||
.build()
|
||||
.map_err(|error| Error::ModelAsset(error.to_string()))?;
|
||||
Self::load_default_with_client(config, &client)
|
||||
}
|
||||
|
||||
pub fn load_default_with_client(
|
||||
config: SembleConfig,
|
||||
client: &reqwest::blocking::Client,
|
||||
) -> Result<Self> {
|
||||
let model_path = ModelAssets::model_path(&config.cache_dir);
|
||||
let cached = { EMBEDDERS.lock().get(&model_path).cloned() };
|
||||
let embedder = if let Some(embedder) = cached {
|
||||
embedder
|
||||
} else {
|
||||
let assets = ModelAssets::ensure(&config.cache_dir)?;
|
||||
let assets = ModelAssets::ensure_with_client(&config.cache_dir, client)?;
|
||||
let embedder = Arc::new(StaticEmbedder::load(&assets.model, &assets.tokenizer)?);
|
||||
EMBEDDERS
|
||||
.lock()
|
||||
|
||||
+1
-1
@@ -33,7 +33,7 @@ parking_lot = "0.12"
|
||||
pem = "3"
|
||||
prost = "0.13"
|
||||
prost-types = "0.13"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["brotli", "deflate", "gzip", "json", "rustls-tls", "socks", "stream", "system-proxy", "zstd"] }
|
||||
reqwest = { version = "0.12", default-features = false, features = ["blocking", "brotli", "deflate", "gzip", "json", "native-tls", "socks", "stream", "system-proxy", "zstd"] }
|
||||
regex = "1"
|
||||
rcgen = { version = "0.14", features = ["aws_lc_rs", "pem", "x509-parser"] }
|
||||
scraper = "0.24"
|
||||
|
||||
+150
-13
@@ -279,10 +279,7 @@ impl ControlService {
|
||||
.model_tests
|
||||
.lock()
|
||||
.expect("model test registry mutex poisoned");
|
||||
tests
|
||||
.entry(test_id.to_owned())
|
||||
.or_insert_with(CancellationToken::new)
|
||||
.clone()
|
||||
tests.entry(test_id.to_owned()).or_default().clone()
|
||||
};
|
||||
cancellation.cancel();
|
||||
}
|
||||
@@ -739,13 +736,64 @@ fn model_discovery_url(base_url: &str) -> Result<Url> {
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
fn model_discovery_urls(base_url: &str) -> Result<Vec<Url>> {
|
||||
let mut configured = Url::parse(base_url)
|
||||
.map_err(|error| Error::Config(format!("invalid model request URL: {error}")))?;
|
||||
let path = configured.path().trim_end_matches('/');
|
||||
let tail = path.rsplit('/').next().unwrap_or_default();
|
||||
if matches!(tail.to_ascii_lowercase().as_str(), "model" | "models") {
|
||||
configured.set_query(None);
|
||||
configured.set_fragment(None);
|
||||
return Ok(vec![configured]);
|
||||
}
|
||||
|
||||
let primary = model_discovery_url(base_url)?;
|
||||
let versioned = tail.len() > 1
|
||||
&& tail.starts_with('v')
|
||||
&& tail[1..].bytes().all(|byte| byte.is_ascii_digit());
|
||||
let complete_request_url = [
|
||||
"/chat/completions",
|
||||
"/responses",
|
||||
"/messages",
|
||||
"/completions",
|
||||
]
|
||||
.iter()
|
||||
.any(|suffix| path.to_ascii_lowercase().ends_with(suffix));
|
||||
if versioned || complete_request_url {
|
||||
return Ok(vec![primary]);
|
||||
}
|
||||
|
||||
let Some(prefix) = primary.path().strip_suffix("/v1/models") else {
|
||||
return Ok(vec![primary]);
|
||||
};
|
||||
let mut fallback = primary.clone();
|
||||
fallback.set_path(&format!("{prefix}/models"));
|
||||
Ok(vec![primary, fallback])
|
||||
}
|
||||
|
||||
async fn openai_models(
|
||||
client: &reqwest::Client,
|
||||
base_url: &str,
|
||||
api_key: &str,
|
||||
custom_headers: &serde_json::Value,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut request = client.get(model_discovery_url(base_url)?);
|
||||
let mut last_error = None;
|
||||
for url in model_discovery_urls(base_url)? {
|
||||
match openai_models_at(client, url, api_key, custom_headers).await {
|
||||
Ok(models) => return Ok(models),
|
||||
Err(error) => last_error = Some(error),
|
||||
}
|
||||
}
|
||||
Err(last_error.unwrap_or_else(|| Error::Provider("no model discovery URL available".into())))
|
||||
}
|
||||
|
||||
async fn openai_models_at(
|
||||
client: &reqwest::Client,
|
||||
url: Url,
|
||||
api_key: &str,
|
||||
custom_headers: &serde_json::Value,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut request = client.get(url);
|
||||
if !api_key.is_empty() {
|
||||
request = request.bearer_auth(api_key);
|
||||
}
|
||||
@@ -767,12 +815,28 @@ async fn anthropic_models(
|
||||
base_url: &str,
|
||||
api_key: &str,
|
||||
custom_headers: &serde_json::Value,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut last_error = None;
|
||||
for url in model_discovery_urls(base_url)? {
|
||||
match anthropic_models_at(client, url, api_key, custom_headers).await {
|
||||
Ok(models) => return Ok(models),
|
||||
Err(error) => last_error = Some(error),
|
||||
}
|
||||
}
|
||||
Err(last_error.unwrap_or_else(|| Error::Provider("no model discovery URL available".into())))
|
||||
}
|
||||
|
||||
async fn anthropic_models_at(
|
||||
client: &reqwest::Client,
|
||||
url: Url,
|
||||
api_key: &str,
|
||||
custom_headers: &serde_json::Value,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut after_id = None::<String>;
|
||||
let mut found = BTreeSet::new();
|
||||
loop {
|
||||
let mut request = client
|
||||
.get(model_discovery_url(base_url)?)
|
||||
.get(url.clone())
|
||||
.query(&[("limit", "100")])
|
||||
.header("anthropic-version", "2023-06-01");
|
||||
if !api_key.is_empty() {
|
||||
@@ -871,26 +935,99 @@ mod tests {
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::{model_discovery_url, ControlService};
|
||||
use super::{model_discovery_url, model_discovery_urls, ControlService};
|
||||
|
||||
#[test]
|
||||
fn model_discovery_url_appends_to_path() {
|
||||
let cases = [
|
||||
("https://api.deepseek.com", "https://api.deepseek.com/v1/models"),
|
||||
("https://open.bigmodel.cn/api/anthropic", "https://open.bigmodel.cn/api/anthropic/v1/models"),
|
||||
("https://api.kimi.com/coding", "https://api.kimi.com/coding/v1/models"),
|
||||
("https://api.moonshot.cn/v1", "https://api.moonshot.cn/v1/models"),
|
||||
("https://ark.cn-beijing.volces.com/api/v3", "https://ark.cn-beijing.volces.com/api/v3/models"),
|
||||
(
|
||||
"https://api.deepseek.com",
|
||||
"https://api.deepseek.com/v1/models",
|
||||
),
|
||||
(
|
||||
"https://open.bigmodel.cn/api/anthropic",
|
||||
"https://open.bigmodel.cn/api/anthropic/v1/models",
|
||||
),
|
||||
(
|
||||
"https://api.kimi.com/coding",
|
||||
"https://api.kimi.com/coding/v1/models",
|
||||
),
|
||||
(
|
||||
"https://api.moonshot.cn/v1",
|
||||
"https://api.moonshot.cn/v1/models",
|
||||
),
|
||||
(
|
||||
"https://ark.cn-beijing.volces.com/api/v3",
|
||||
"https://ark.cn-beijing.volces.com/api/v3/models",
|
||||
),
|
||||
(
|
||||
"https://open.bigmodel.cn/api/coding/paas/v4/chat/completions",
|
||||
"https://open.bigmodel.cn/api/coding/paas/v4/models",
|
||||
),
|
||||
];
|
||||
for (base, expected) in cases {
|
||||
assert_eq!(model_discovery_url(base).unwrap().as_str(), expected, "base: {base}");
|
||||
assert_eq!(
|
||||
model_discovery_url(base).unwrap().as_str(),
|
||||
expected,
|
||||
"base: {base}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_discovery_urls_fall_back_without_a_version() {
|
||||
let cases = [
|
||||
(
|
||||
"https://opencode.ai/zen/go/v1",
|
||||
vec!["https://opencode.ai/zen/go/v1/models"],
|
||||
),
|
||||
(
|
||||
"https://opencode.ai/zen/go",
|
||||
vec![
|
||||
"https://opencode.ai/zen/go/v1/models",
|
||||
"https://opencode.ai/zen/go/models",
|
||||
],
|
||||
),
|
||||
(
|
||||
"https://api.example.com/openai/v1/models",
|
||||
vec!["https://api.example.com/openai/v1/models"],
|
||||
),
|
||||
];
|
||||
for (base, expected) in cases {
|
||||
let actual = model_discovery_urls(base)
|
||||
.unwrap()
|
||||
.into_iter()
|
||||
.map(|url| url.to_string())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(actual, expected, "base: {base}");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_model_discovery_uses_the_unversioned_fallback() {
|
||||
let app = axum::Router::new().route(
|
||||
"/proxy/models",
|
||||
axum::routing::get(|| async {
|
||||
axum::Json(serde_json::json!({ "data": [{ "id": "model-a" }] }))
|
||||
}),
|
||||
);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
|
||||
|
||||
let models = super::openai_models(
|
||||
&reqwest::Client::new(),
|
||||
&format!("http://{address}/proxy"),
|
||||
"secret",
|
||||
&serde_json::json!({}),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(models, vec!["model-a"]);
|
||||
server.abort();
|
||||
}
|
||||
|
||||
struct TestProvider {
|
||||
invocation: Arc<Mutex<Option<ModelInvocation>>>,
|
||||
}
|
||||
|
||||
@@ -48,7 +48,11 @@ impl CursorActor {
|
||||
let tool_runtime = CursorToolRuntime::default();
|
||||
let context_sync =
|
||||
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
|
||||
let tools = ToolDispatcher::with_results(tool_runtime.clone(), results_tx.clone());
|
||||
let tools = ToolDispatcher::with_results(
|
||||
tool_runtime.clone(),
|
||||
results_tx.clone(),
|
||||
dependencies.store.clone(),
|
||||
);
|
||||
let mut run_resources = Some((results_rx, runtime_actions_rx, dependencies));
|
||||
loop {
|
||||
let command = match receiver.recv().await {
|
||||
|
||||
@@ -10,6 +10,7 @@ use std::collections::BTreeMap;
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::ToolCall,
|
||||
store::Store,
|
||||
web::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
@@ -36,6 +37,7 @@ pub(super) async fn start(
|
||||
message_index: usize,
|
||||
dynamic_mcp: &BTreeMap<String, pb::McpToolDefinition>,
|
||||
context: &ExecContext,
|
||||
store: Option<&Store>,
|
||||
) -> Result<ToolStart> {
|
||||
if let Some(definition) = dynamic_mcp.get(&call.name) {
|
||||
return exec::start_dynamic(runtime, call, definition, context).await;
|
||||
@@ -57,7 +59,7 @@ pub(super) async fn start(
|
||||
| "generateimage" => interaction::start(runtime, call).await,
|
||||
"todowrite" | "updatecurrentstep" => local::start(call, message_index),
|
||||
"awaitshell" => await_shell::start(runtime, results, call, context).await,
|
||||
"semblesearch" | "semblefindrelated" => semble::start(results, call),
|
||||
"semblesearch" | "semblefindrelated" => semble::start(results, call, store.cloned()),
|
||||
_ => Err(Error::Protocol(format!("unsupported tool: {}", call.name))),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,7 +7,7 @@ use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::OnceCell;
|
||||
|
||||
use crate::{model::ToolCall, Error, Result};
|
||||
use crate::{model::ToolCall, store::Store, Error, Result};
|
||||
|
||||
use super::ToolStart;
|
||||
use crate::cursor::tools::{
|
||||
@@ -52,7 +52,11 @@ struct FindRelatedArguments {
|
||||
content: ContentSelection,
|
||||
}
|
||||
|
||||
pub(super) fn start(results: &ToolResultSender, call: &ToolCall) -> Result<ToolStart> {
|
||||
pub(super) fn start(
|
||||
results: &ToolResultSender,
|
||||
call: &ToolCall,
|
||||
store: Option<Store>,
|
||||
) -> Result<ToolStart> {
|
||||
let operation = match super::normalized(&call.name).as_str() {
|
||||
"semblesearch" => Operation::Search(serde_json::from_value(call.arguments.clone())?),
|
||||
"semblefindrelated" => {
|
||||
@@ -69,7 +73,7 @@ pub(super) fn start(results: &ToolResultSender, call: &ToolCall) -> Result<ToolS
|
||||
let results = results.clone();
|
||||
let started_at_ms = now_ms();
|
||||
tokio::spawn(async move {
|
||||
let output = execute(operation).await;
|
||||
let output = execute(operation, store).await;
|
||||
match result::semble(&call, started_at_ms, output) {
|
||||
Ok(completion) => results.send(completion),
|
||||
Err(error) => results.send_error(error),
|
||||
@@ -86,8 +90,8 @@ enum Operation {
|
||||
FindRelated(FindRelatedArguments),
|
||||
}
|
||||
|
||||
async fn execute(operation: Operation) -> std::result::Result<Value, String> {
|
||||
let engine = engine().await.map_err(|error| error.to_string())?;
|
||||
async fn execute(operation: Operation, store: Option<Store>) -> std::result::Result<Value, String> {
|
||||
let engine = engine(store).await.map_err(|error| error.to_string())?;
|
||||
tokio::task::spawn_blocking(move || match operation {
|
||||
Operation::Search(arguments) => engine
|
||||
.search(SearchRequest {
|
||||
@@ -114,14 +118,21 @@ async fn execute(operation: Operation) -> std::result::Result<Value, String> {
|
||||
.map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
async fn engine() -> Result<Arc<SearchEngine>> {
|
||||
async fn engine(store: Option<Store>) -> Result<Arc<SearchEngine>> {
|
||||
ENGINE
|
||||
.get_or_try_init(|| async {
|
||||
tokio::task::spawn_blocking(|| SearchEngine::load_default(SembleConfig::default()))
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))?
|
||||
.map(Arc::new)
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))
|
||||
.get_or_try_init(|| async move {
|
||||
let builder = match store {
|
||||
Some(store) => crate::network::blocking_client_builder(&store).await?,
|
||||
None => reqwest::blocking::Client::builder().use_native_tls(),
|
||||
};
|
||||
tokio::task::spawn_blocking(move || {
|
||||
let client = builder.build()?;
|
||||
SearchEngine::load_default_with_client(SembleConfig::default(), &client)
|
||||
.map(Arc::new)
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))
|
||||
})
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))?
|
||||
})
|
||||
.await
|
||||
.cloned()
|
||||
|
||||
@@ -17,6 +17,7 @@ mod tests;
|
||||
|
||||
use crate::{
|
||||
model::{CanonicalMessage, MessageContent, Role, ToolCall},
|
||||
store::Store,
|
||||
web::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
@@ -32,6 +33,7 @@ pub struct ToolDispatcher {
|
||||
results: ToolResultSender,
|
||||
search: WebSearch,
|
||||
fetch: WebFetch,
|
||||
store: Option<Store>,
|
||||
edit_schedule: Arc<Mutex<EditSchedule>>,
|
||||
}
|
||||
|
||||
@@ -55,15 +57,27 @@ pub enum ClientToolEvent {
|
||||
impl ToolDispatcher {
|
||||
pub fn new(runtime: CursorToolRuntime) -> Self {
|
||||
let (results, _) = result::tool_result_channel();
|
||||
Self::with_results(runtime, results)
|
||||
}
|
||||
|
||||
pub fn with_results(runtime: CursorToolRuntime, results: ToolResultSender) -> Self {
|
||||
Self {
|
||||
runtime,
|
||||
results,
|
||||
search: WebSearch::built_in(),
|
||||
fetch: WebFetch::built_in(),
|
||||
store: None,
|
||||
edit_schedule: Arc::new(Mutex::new(EditSchedule::default())),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_results(
|
||||
runtime: CursorToolRuntime,
|
||||
results: ToolResultSender,
|
||||
store: Store,
|
||||
) -> Self {
|
||||
Self {
|
||||
runtime,
|
||||
results,
|
||||
search: WebSearch::managed(store.clone()),
|
||||
fetch: WebFetch::managed(store.clone()),
|
||||
store: Some(store),
|
||||
edit_schedule: Arc::new(Mutex::new(EditSchedule::default())),
|
||||
}
|
||||
}
|
||||
@@ -165,6 +179,7 @@ impl ToolDispatcher {
|
||||
message_index,
|
||||
dynamic_mcp,
|
||||
context,
|
||||
self.store.as_ref(),
|
||||
)
|
||||
.await?;
|
||||
messages.extend(started.messages);
|
||||
|
||||
+113
-1
@@ -4,7 +4,9 @@ use crate::{store::Store, Result};
|
||||
|
||||
pub async fn client_builder(store: &Store) -> Result<reqwest::ClientBuilder> {
|
||||
let settings = store.proxy_settings_secret().await?;
|
||||
let mut builder = reqwest::Client::builder();
|
||||
// Use the platform TLS stack for compatibility with provider gateways that
|
||||
// only offer legacy TLS 1.2 cipher suites unsupported by rustls.
|
||||
let mut builder = reqwest::Client::builder().use_native_tls();
|
||||
if settings.mode.is_custom() {
|
||||
let mut proxy = reqwest::Proxy::all(&settings.address)?;
|
||||
if settings.auth_enabled {
|
||||
@@ -18,3 +20,113 @@ pub async fn client_builder(store: &Store) -> Result<reqwest::ClientBuilder> {
|
||||
pub async fn client(store: &Store) -> Result<reqwest::Client> {
|
||||
Ok(client_builder(store).await?.build()?)
|
||||
}
|
||||
|
||||
pub async fn blocking_client_builder(store: &Store) -> Result<reqwest::blocking::ClientBuilder> {
|
||||
let settings = store.proxy_settings_secret().await?;
|
||||
let mut builder = reqwest::blocking::Client::builder().use_native_tls();
|
||||
if settings.mode.is_custom() {
|
||||
let mut proxy = reqwest::Proxy::all(&settings.address)?;
|
||||
if settings.auth_enabled {
|
||||
proxy = proxy.basic_auth(&settings.username, &settings.password);
|
||||
}
|
||||
builder = builder.no_proxy().proxy(proxy);
|
||||
}
|
||||
Ok(builder)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{
|
||||
io::{BufRead, BufReader, Write},
|
||||
net::TcpListener,
|
||||
sync::mpsc,
|
||||
thread,
|
||||
};
|
||||
|
||||
use crate::store::{ProxyMode, ProxySettingsInput};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn custom_proxy_applies_to_async_and_blocking_clients() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let database_url = format!("sqlite://{}", directory.path().join("test.db").display());
|
||||
let store = Store::connect(&database_url).await.unwrap();
|
||||
let (proxy_address, requests, proxy) = proxy_server(2);
|
||||
store
|
||||
.set_proxy_settings(ProxySettingsInput {
|
||||
mode: ProxyMode::Custom,
|
||||
address: proxy_address,
|
||||
auth_enabled: true,
|
||||
username: "proxy-user".into(),
|
||||
password: Some("proxy-password".into()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
client(&store)
|
||||
.await
|
||||
.unwrap()
|
||||
.get("http://provider.invalid/async")
|
||||
.send()
|
||||
.await
|
||||
.unwrap()
|
||||
.error_for_status()
|
||||
.unwrap();
|
||||
let blocking = blocking_client_builder(&store).await.unwrap();
|
||||
tokio::task::spawn_blocking(move || {
|
||||
blocking
|
||||
.build()
|
||||
.unwrap()
|
||||
.get("http://provider.invalid/blocking")
|
||||
.send()
|
||||
.unwrap()
|
||||
.error_for_status()
|
||||
.unwrap();
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let requests = [requests.recv().unwrap(), requests.recv().unwrap()];
|
||||
assert!(requests
|
||||
.iter()
|
||||
.any(|request| request.starts_with("GET http://provider.invalid/async ")));
|
||||
assert!(requests
|
||||
.iter()
|
||||
.any(|request| request.starts_with("GET http://provider.invalid/blocking ")));
|
||||
assert!(requests.iter().all(|request| request
|
||||
.to_ascii_lowercase()
|
||||
.contains("\r\nproxy-authorization: basic ")));
|
||||
proxy.join().unwrap();
|
||||
}
|
||||
|
||||
fn proxy_server(
|
||||
expected_requests: usize,
|
||||
) -> (String, mpsc::Receiver<String>, thread::JoinHandle<()>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
|
||||
let address = format!("http://{}", listener.local_addr().unwrap());
|
||||
let (sender, receiver) = mpsc::channel();
|
||||
let server = thread::spawn(move || {
|
||||
for stream in listener.incoming().take(expected_requests) {
|
||||
let mut stream = stream.unwrap();
|
||||
let mut request = String::new();
|
||||
let mut reader = BufReader::new(stream.try_clone().unwrap());
|
||||
loop {
|
||||
let mut line = String::new();
|
||||
reader.read_line(&mut line).unwrap();
|
||||
request.push_str(&line);
|
||||
if line == "\r\n" {
|
||||
break;
|
||||
}
|
||||
}
|
||||
sender.send(request).unwrap();
|
||||
stream
|
||||
.write_all(
|
||||
b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok",
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
});
|
||||
(address, receiver, server)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,8 @@ use std::{cmp::Ordering, collections::HashMap};
|
||||
|
||||
use futures_util::future::join_all;
|
||||
|
||||
use crate::store::Store;
|
||||
|
||||
use super::{catalog, SearchEngine, SearchHit};
|
||||
|
||||
const RRF_K: f64 = 60.0;
|
||||
@@ -9,10 +11,16 @@ const MAX_RESULTS: usize = 10;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct WebSearch {
|
||||
client: reqwest::Client,
|
||||
client: SearchClient,
|
||||
engines: Vec<SearchEngine>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
enum SearchClient {
|
||||
Managed(Store),
|
||||
Direct(reqwest::Client),
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
#[error("web search failed: {0}")]
|
||||
pub struct SearchError(String);
|
||||
@@ -22,13 +30,20 @@ impl WebSearch {
|
||||
Self::with_engines(catalog::engines())
|
||||
}
|
||||
|
||||
pub(crate) fn managed(store: Store) -> Self {
|
||||
Self {
|
||||
client: SearchClient::Managed(store),
|
||||
engines: catalog::engines(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_engines<I, E>(engines: I) -> Self
|
||||
where
|
||||
I: IntoIterator<Item = E>,
|
||||
E: Into<SearchEngine>,
|
||||
{
|
||||
Self {
|
||||
client: reqwest::Client::new(),
|
||||
client: SearchClient::Direct(reqwest::Client::new()),
|
||||
engines: engines.into_iter().map(Into::into).collect(),
|
||||
}
|
||||
}
|
||||
@@ -42,10 +57,16 @@ impl WebSearch {
|
||||
if query.is_empty() {
|
||||
return Err(SearchError("query is empty".into()));
|
||||
}
|
||||
let client = match &self.client {
|
||||
SearchClient::Managed(store) => crate::network::client(store)
|
||||
.await
|
||||
.map_err(|error| SearchError(format!("HTTP client failed: {error}")))?,
|
||||
SearchClient::Direct(client) => client.clone(),
|
||||
};
|
||||
let responses = join_all(
|
||||
self.engines
|
||||
.iter()
|
||||
.map(|engine| engine.search(&self.client, query)),
|
||||
.map(|engine| engine.search(&client, query)),
|
||||
)
|
||||
.await;
|
||||
let mut merged = HashMap::<String, SearchHit>::new();
|
||||
|
||||
+25
-1
@@ -14,6 +14,8 @@ use reqwest::{
|
||||
use tokio::{net::lookup_host, time::timeout};
|
||||
use url::{Host, Url};
|
||||
|
||||
use crate::store::Store;
|
||||
|
||||
const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
|
||||
const MAX_REDIRECTS: usize = 5;
|
||||
const FETCH_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
@@ -38,12 +40,27 @@ enum NetworkPolicy {
|
||||
#[derive(Clone)]
|
||||
pub struct WebFetch {
|
||||
network: NetworkPolicy,
|
||||
client: FetchClient,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
enum FetchClient {
|
||||
Managed(Store),
|
||||
Direct,
|
||||
}
|
||||
|
||||
impl WebFetch {
|
||||
pub fn built_in() -> Self {
|
||||
Self {
|
||||
network: NetworkPolicy::PublicOnly,
|
||||
client: FetchClient::Direct,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn managed(store: Store) -> Self {
|
||||
Self {
|
||||
network: NetworkPolicy::PublicOnly,
|
||||
client: FetchClient::Managed(store),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,6 +68,7 @@ impl WebFetch {
|
||||
pub(crate) fn for_test() -> Self {
|
||||
Self {
|
||||
network: NetworkPolicy::Any,
|
||||
client: FetchClient::Direct,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,7 +129,13 @@ impl WebFetch {
|
||||
return Err(failure("URL resolves to a non-public address"));
|
||||
}
|
||||
|
||||
let mut builder = reqwest::Client::builder()
|
||||
let builder = match &self.client {
|
||||
FetchClient::Managed(store) => crate::network::client_builder(store)
|
||||
.await
|
||||
.map_err(|error| failure(format!("HTTP client failed: {error}")))?,
|
||||
FetchClient::Direct => reqwest::Client::builder().use_native_tls(),
|
||||
};
|
||||
let mut builder = builder
|
||||
.redirect(Policy::none())
|
||||
.connect_timeout(Duration::from_secs(10));
|
||||
if domain {
|
||||
|
||||
Reference in New Issue
Block a user