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:
leookun
2026-08-27 00:34:35 +08:00
parent 7058d0193d
commit 995a12df44
20 changed files with 560 additions and 154 deletions
Generated
+100 -2
View File
@@ -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"
-30
View File
@@ -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",
-2
View File
@@ -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",
+2 -2
View File
@@ -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>
+1 -1
View File
@@ -118,7 +118,7 @@ export function HomePage() {
{
key: "metrics",
estimatedHeight: 130,
content: <HomeMetrics data={metrics} />,
content: <HomeMetrics data={metrics} refreshVersion={refreshVersion} />,
},
{
+19 -6
View File
@@ -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)
+11 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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>>>,
}
+5 -1
View File
@@ -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 {
+3 -1
View File
@@ -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))),
}
}
+23 -12
View File
@@ -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()
+19 -4
View File
@@ -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
View File
@@ -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)
}
}
+24 -3
View File
@@ -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
View File
@@ -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 {