Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 69 additions & 2 deletions clashctl-core/src/api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,15 +4,17 @@ use std::{
time::Duration,
};

use log::{debug, trace};
use log::{debug, trace, warn};
use serde::de::DeserializeOwned;
use serde_json::{from_str, json};
use ureq::{Agent, Request};
use url::Url;

use crate::{
model::{Config, Connections, Delay, Log, Proxies, Proxy, Rules, Traffic, Version},
Error, Result,
model::{
Config, Connections, Delay, Log, Proxies, Proxy, ProxyProviders, Rules, Traffic, Version,
},
};

trait Convert<T: DeserializeOwned> {
Expand Down Expand Up @@ -233,6 +235,30 @@ impl Clash {
self.get("proxies")
}

/// Get Mihomo proxy-provider information.
pub fn get_proxy_providers(&self) -> Result<ProxyProviders> {
self.get("providers/proxies")
}

/// Get proxies, including nodes that are only exposed through
/// `/providers/proxies`.
///
/// Older controllers may not implement the provider endpoint. Its failure
/// is non-fatal: the complete top-level `/proxies` response is returned.
pub fn get_proxies_with_providers(&self) -> Result<Proxies> {
let proxies = self.get_proxies()?;
match self.get_proxy_providers() {
Ok(proxy_providers) => Ok(proxies.merge_proxy_providers(proxy_providers)),
Err(error) => {
warn!(
"Could not retrieve proxy providers; using /proxies response only: {}",
error
);
Ok(proxies)
}
}
}

/// Get rules information
pub fn get_rules(&self) -> Result<Rules> {
self.get("rules")
Expand Down Expand Up @@ -300,6 +326,47 @@ impl Clash {
}
}

#[cfg(test)]
mod tests {
use std::{
io::{Read as _, Write as _},
net::TcpListener,
thread,
};

use super::Clash;

#[test]
fn provider_endpoint_failure_falls_back_to_top_level_proxies() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
let server = thread::spawn(move || {
for (index, stream) in listener.incoming().take(2).enumerate() {
let mut stream = stream.unwrap();
let mut request = [0; 1024];
stream.read(&mut request).unwrap();
let response = if index == 0 {
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nConnection: \
close\r\n\r\n{\"proxies\":{\"Group\":{\"type\":\"Selector\",\"history\":[],\"\
all\":[\"provider-node\"],\"now\":\"provider-node\"}}}"
} else {
"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
};
stream.write_all(response.as_bytes()).unwrap();
}
});
let clash = Clash::builder(format!("http://{}", address))
.unwrap()
.build();

let proxies = clash.get_proxies_with_providers().unwrap();

assert!(proxies.contains_key("Group"));
assert!(!proxies.contains_key("provider-node"));
server.join().unwrap();
}
}

pub struct LongHaul<T: DeserializeOwned> {
reader: BufReader<Box<dyn Read + Send>>,
ty: PhantomData<T>,
Expand Down
132 changes: 130 additions & 2 deletions clashctl-core/src/model/proxy.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
use std::collections::HashMap;
use std::ops::Deref;
use std::{
collections::{BTreeMap, HashMap},
ops::Deref,
};

use serde::{Deserialize, Serialize};

Expand All @@ -11,6 +13,17 @@ pub struct Proxies {
}

impl Proxies {
/// Add proxies returned by Mihomo proxy providers without replacing the
/// canonical entries returned by `/proxies`.
pub fn merge_proxy_providers(mut self, proxy_providers: ProxyProviders) -> Self {
for provider in proxy_providers.providers.into_values() {
for proxy in provider.proxies {
self.proxies.entry(proxy.name).or_insert(proxy.proxy);
}
}
self
}

pub fn normal(&self) -> impl Iterator<Item = (&String, &Proxy)> {
self.iter().filter(|(_, x)| x.proxy_type.is_normal())
}
Expand All @@ -28,8 +41,29 @@ impl Proxies {
}
}

/// Response returned by Mihomo's `/providers/proxies` endpoint.
#[derive(Serialize, Deserialize, Clone, Debug, Default, PartialEq, Eq)]
pub struct ProxyProviders {
/// A `BTreeMap` makes duplicate names across providers resolve in a
/// deterministic provider order when merged.
pub providers: BTreeMap<String, ProxyProvider>,
}

#[derive(Serialize, Deserialize, Clone, Debug, Default, PartialEq, Eq)]
pub struct ProxyProvider {
pub proxies: Vec<ProviderProxy>,
}

#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, Eq)]
pub struct ProviderProxy {
pub name: String,
#[serde(flatten)]
pub proxy: Proxy,
}

impl Deref for Proxies {
type Target = HashMap<String, Proxy>;

fn deref(&self) -> &Self::Target {
&self.proxies
}
Expand Down Expand Up @@ -181,3 +215,97 @@ fn test_proxies() {
vec!["test_c"]
);
}

#[cfg(test)]
mod provider_tests {
use serde_json::from_str;

use super::*;

fn proxy(proxy_type: ProxyType) -> Proxy {
Proxy {
proxy_type,
history: vec![],
udp: None,
all: None,
now: None,
}
}

#[test]
fn deserializes_and_merges_provider_only_proxies() {
let providers: ProxyProviders = from_str(
r#"{"providers":{"subscription":{"proxies":[{"name":"provider-node","type":"Shadowsocks","history":[]}]}}}"#,
)
.unwrap();
let merged = Proxies::default().merge_proxy_providers(providers);

assert_eq!(merged["provider-node"].proxy_type, ProxyType::Shadowsocks);
}

#[test]
fn top_level_proxy_takes_precedence_over_provider_proxy() {
let providers = ProxyProviders {
providers: BTreeMap::from([(
"subscription".to_owned(),
ProxyProvider {
proxies: vec![ProviderProxy {
name: "duplicate".to_owned(),
proxy: proxy(ProxyType::Shadowsocks),
}],
},
)]),
};
let top_level = Proxies {
proxies: HashMap::from([("duplicate".to_owned(), proxy(ProxyType::Direct))]),
};

let merged = top_level.merge_proxy_providers(providers);

assert_eq!(merged["duplicate"].proxy_type, ProxyType::Direct);
}

#[test]
fn first_provider_in_deterministic_order_wins_duplicate_names() {
let providers = ProxyProviders {
providers: BTreeMap::from([
(
"a-provider".to_owned(),
ProxyProvider {
proxies: vec![ProviderProxy {
name: "duplicate".to_owned(),
proxy: proxy(ProxyType::Shadowsocks),
}],
},
),
(
"z-provider".to_owned(),
ProxyProvider {
proxies: vec![ProviderProxy {
name: "duplicate".to_owned(),
proxy: proxy(ProxyType::Trojan),
}],
},
),
]),
};

let merged = Proxies::default().merge_proxy_providers(providers);

assert_eq!(merged["duplicate"].proxy_type, ProxyType::Shadowsocks);
}

#[test]
fn empty_provider_response_preserves_proxy_map() {
let top_level = Proxies {
proxies: HashMap::from([("node".to_owned(), proxy(ProxyType::Trojan))]),
};

assert_eq!(
top_level
.clone()
.merge_proxy_providers(ProxyProviders::default()),
top_level
);
}
}
11 changes: 6 additions & 5 deletions clashctl/src/command/proxy.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,11 @@ use clap::{Parser, Subcommand};
use clashctl_core::{model::ProxyType, strum::VariantNames};
use log::{error, info, warn};
use owo_colors::OwoColorize;
use requestty::{prompt_one, Answer, ListItem, Question};
use requestty::{Answer, ListItem, Question, prompt_one};

use crate::{
interactive::{Flags, ProxySortBy, SortOrder},
RenderList, Result,
interactive::{Flags, ProxySortBy, SortOrder},
};
// use crate::{Result};

Expand Down Expand Up @@ -101,11 +101,11 @@ impl ProxySubcommand {

match self {
ProxySubcommand::List(opt) => {
let proxies = clash.get_proxies()?;
let proxies = clash.get_proxies_with_providers()?;
proxies.render_list(opt);
}
ProxySubcommand::Use => {
let proxies = clash.get_proxies()?;
let proxies = clash.get_proxies_with_providers()?;
let mut groups = proxies
.iter()
.filter(|(_, p)| p.proxy_type.is_selector())
Expand All @@ -127,7 +127,8 @@ impl ProxySubcommand {
};
let proxy = clash.get_proxy(&group_selected)?;

// all / now only occurs when proxy_type is [`ProxyType::Selector`]
// all / now only occurs when proxy_type is
// [`ProxyType::Selector`]
let members = proxy.all.unwrap();
let now = proxy.now.unwrap();
let cur_index = members.iter().position(|x| x == &now).unwrap();
Expand Down
Loading