mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 20:44:07 +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:
+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