From 995a12df446d7b581a2c40ddc4af6b042c7a1ed7 Mon Sep 17 00:00:00 2001 From: leookun Date: Thu, 27 Aug 2026 00:34:35 +0800 Subject: [PATCH] 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. --- Cargo.lock | 102 ++++++++++- apps/desktop/package-lock.json | 30 ---- apps/desktop/package.json | 2 - apps/desktop/src/api.ts | 4 +- apps/desktop/src/components/charts/EChart.tsx | 4 +- .../metrics/CacheHitRateChart.module.scss | 16 +- .../components/metrics/CacheHitRateChart.tsx | 113 ++++++------ .../src/components/metrics/HomeMetrics.tsx | 4 +- apps/desktop/src/pages/HomePage.tsx | 2 +- crates/semble-core/src/embedding/assets.rs | 25 ++- crates/semble-core/src/search/engine.rs | 12 +- server/Cargo.toml | 2 +- server/src/control/service.rs | 163 ++++++++++++++++-- server/src/cursor/actor.rs | 6 +- server/src/cursor/tools/dispatch/mod.rs | 4 +- server/src/cursor/tools/dispatch/semble.rs | 35 ++-- server/src/cursor/tools/mod.rs | 23 ++- server/src/network.rs | 114 +++++++++++- server/src/web/federation.rs | 27 ++- server/src/web/fetch.rs | 26 ++- 20 files changed, 560 insertions(+), 154 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 9dac26b..50f296b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -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" diff --git a/apps/desktop/package-lock.json b/apps/desktop/package-lock.json index fe3f47f..af7962c 100644 --- a/apps/desktop/package-lock.json +++ b/apps/desktop/package-lock.json @@ -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", diff --git a/apps/desktop/package.json b/apps/desktop/package.json index 259cbed..2cc6a7d 100644 --- a/apps/desktop/package.json +++ b/apps/desktop/package.json @@ -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", diff --git a/apps/desktop/src/api.ts b/apps/desktop/src/api.ts index 8226b99..9eabcad 100644 --- a/apps/desktop/src/api.ts +++ b/apps/desktop/src/api.ts @@ -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${query}`); + const query = params.toString(); + return request(`/overview${query ? `?${query}` : ""}`); }, cursorHarness: () => request("/harness/cursor/status"), initializeCursorCa: () => request("/harness/cursor/ca/initialize", { method: "POST" }), diff --git a/apps/desktop/src/components/charts/EChart.tsx b/apps/desktop/src/components/charts/EChart.tsx index 7c5776e..940e737 100644 --- a/apps/desktop/src/components/charts/EChart.tsx +++ b/apps/desktop/src/components/charts/EChart.tsx @@ -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; diff --git a/apps/desktop/src/components/metrics/CacheHitRateChart.module.scss b/apps/desktop/src/components/metrics/CacheHitRateChart.module.scss index 4f213e9..ecf6b7a 100644 --- a/apps/desktop/src/components/metrics/CacheHitRateChart.module.scss +++ b/apps/desktop/src/components/metrics/CacheHitRateChart.module.scss @@ -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 { diff --git a/apps/desktop/src/components/metrics/CacheHitRateChart.tsx b/apps/desktop/src/components/metrics/CacheHitRateChart.tsx index 4eb34bb..05c2c68 100644 --- a/apps/desktop/src/components/metrics/CacheHitRateChart.tsx +++ b/apps/desktop/src/components/metrics/CacheHitRateChart.tsx @@ -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(() => ({ + 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
- +
{label}
; } diff --git a/apps/desktop/src/components/metrics/HomeMetrics.tsx b/apps/desktop/src/components/metrics/HomeMetrics.tsx index acd6994..e20a418 100644 --- a/apps/desktop/src/components/metrics/HomeMetrics.tsx +++ b/apps/desktop/src/components/metrics/HomeMetrics.tsx @@ -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 }) {
{t("缓存命中率")}
- +
{t("LLM 调用")}
diff --git a/apps/desktop/src/pages/HomePage.tsx b/apps/desktop/src/pages/HomePage.tsx index 7d7bff8..354eaa0 100644 --- a/apps/desktop/src/pages/HomePage.tsx +++ b/apps/desktop/src/pages/HomePage.tsx @@ -118,7 +118,7 @@ export function HomePage() { { key: "metrics", estimatedHeight: 130, - content: , + content: , }, { diff --git a/crates/semble-core/src/embedding/assets.rs b/crates/semble-core/src/embedding/assets.rs index 935e1d4..cdc982b 100644 --- a/crates/semble-core/src/embedding/assets.rs +++ b/crates/semble-core/src/embedding/assets.rs @@ -29,24 +29,37 @@ impl ModelAssets { } pub fn ensure(cache_root: &Path) -> Result { + 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 { 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) diff --git a/crates/semble-core/src/search/engine.rs b/crates/semble-core/src/search/engine.rs index e638ff9..1800cbd 100644 --- a/crates/semble-core/src/search/engine.rs +++ b/crates/semble-core/src/search/engine.rs @@ -35,12 +35,22 @@ static EMBEDDERS: LazyLock>>> = impl SearchEngine { pub fn load_default(config: SembleConfig) -> Result { + 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 { 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() diff --git a/server/Cargo.toml b/server/Cargo.toml index 9b39ed3..1745655 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -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" diff --git a/server/src/control/service.rs b/server/src/control/service.rs index 1f6b81f..e861185 100644 --- a/server/src/control/service.rs +++ b/server/src/control/service.rs @@ -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 { Ok(url) } +fn model_discovery_urls(base_url: &str) -> Result> { + 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> { - 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> { + 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> { + 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> { let mut after_id = None::; 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::>(); + 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>>, } diff --git a/server/src/cursor/actor.rs b/server/src/cursor/actor.rs index de789fb..85c158a 100644 --- a/server/src/cursor/actor.rs +++ b/server/src/cursor/actor.rs @@ -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 { diff --git a/server/src/cursor/tools/dispatch/mod.rs b/server/src/cursor/tools/dispatch/mod.rs index 16dfe6b..eb4ea88 100644 --- a/server/src/cursor/tools/dispatch/mod.rs +++ b/server/src/cursor/tools/dispatch/mod.rs @@ -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, context: &ExecContext, + store: Option<&Store>, ) -> Result { 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))), } } diff --git a/server/src/cursor/tools/dispatch/semble.rs b/server/src/cursor/tools/dispatch/semble.rs index 89a41f4..1852e80 100644 --- a/server/src/cursor/tools/dispatch/semble.rs +++ b/server/src/cursor/tools/dispatch/semble.rs @@ -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 { +pub(super) fn start( + results: &ToolResultSender, + call: &ToolCall, + store: Option, +) -> Result { 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 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 { - let engine = engine().await.map_err(|error| error.to_string())?; +async fn execute(operation: Operation, store: Option) -> std::result::Result { + 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 { .map_err(|error| error.to_string()) } -async fn engine() -> Result> { +async fn engine(store: Option) -> Result> { 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() diff --git a/server/src/cursor/tools/mod.rs b/server/src/cursor/tools/mod.rs index f46f107..2003766 100644 --- a/server/src/cursor/tools/mod.rs +++ b/server/src/cursor/tools/mod.rs @@ -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, edit_schedule: Arc>, } @@ -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); diff --git a/server/src/network.rs b/server/src/network.rs index 9f86a17..72b41b5 100644 --- a/server/src/network.rs +++ b/server/src/network.rs @@ -4,7 +4,9 @@ use crate::{store::Store, Result}; pub async fn client_builder(store: &Store) -> Result { 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 { pub async fn client(store: &Store) -> Result { Ok(client_builder(store).await?.build()?) } + +pub async fn blocking_client_builder(store: &Store) -> Result { + 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, 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) + } +} diff --git a/server/src/web/federation.rs b/server/src/web/federation.rs index 43ff444..da6633a 100644 --- a/server/src/web/federation.rs +++ b/server/src/web/federation.rs @@ -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, } +#[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(engines: I) -> Self where I: IntoIterator, E: Into, { 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::::new(); diff --git a/server/src/web/fetch.rs b/server/src/web/fetch.rs index 0327c32..76c4dc6 100644 --- a/server/src/web/fetch.rs +++ b/server/src/web/fetch.rs @@ -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 {