use pyo3::prelude::*;
use tokio::runtime::Runtime;
use crate::builder::DnsResolverBuilder as RustDnsResolverBuilder;
use crate::builder::strategy::QueryStrategy as RustQueryStrategy;
use crate::upstream_handler::{UpstreamSpec, UpstreamManager};
use super::resolver::PyDnsResolver;
use super::types::PyQueryStrategy;
#[pyclass(name = "DnsResolverBuilder")]
pub struct PyDnsResolverBuilder {
inner: RustDnsResolverBuilder,
}
#[pymethods]
impl PyDnsResolverBuilder {
#[new]
pub fn new() -> Self {
let default_strategy = RustQueryStrategy::Smart;
let default_edns = true;
let default_region = "CN".to_string();
Self {
inner: RustDnsResolverBuilder::new(
default_strategy,
default_edns,
default_region,
),
}
}
pub fn query_strategy(&mut self, strategy: &PyQueryStrategy) -> PyResult<()> {
self.inner = self.inner.clone().query_strategy(strategy.to_rust());
Ok(())
}
pub fn add_udp_upstream(&mut self, name: String, server: String) -> PyResult<()> {
self.inner = self.inner.clone().add_udp_upstream(name, server);
Ok(())
}
pub fn add_tcp_upstream(&mut self, name: String, server: String) -> PyResult<()> {
self.inner = self.inner.clone().add_tcp_upstream(name, server);
Ok(())
}
pub fn add_doh_upstream(&mut self, name: String, url: String) -> PyResult<()> {
if !url.starts_with("https://") {
return Err(PyErr::new::<pyo3::exceptions::PyValueError, _>(
"DoH URL must use HTTPS".to_string()
));
}
self.inner = self.inner.clone().add_doh_upstream(name, url);
Ok(())
}
pub fn add_dot_upstream(&mut self, name: String, server: String) -> PyResult<()> {
self.inner = self.inner.clone().add_dot_upstream(name, server);
Ok(())
}
pub fn timeout(&mut self, timeout_secs: f64) -> PyResult<()> {
let duration = std::time::Duration::from_secs_f64(timeout_secs);
self.inner = self.inner.clone().with_timeout(duration);
Ok(())
}
pub fn region(&mut self, region: String) -> PyResult<()> {
self.inner = self.inner.clone().region(region);
Ok(())
}
pub fn with_public_dns(&mut self) -> PyResult<()> {
self.inner = self.inner.clone().with_public_dns().map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!("Failed to add public DNS: {}", e))
})?;
Ok(())
}
pub fn enable_edns(&mut self, enable: bool) -> PyResult<()> {
self.inner = self.inner.clone().enable_edns(enable);
Ok(())
}
pub fn enable_upstream_monitoring(&mut self, enable: bool) -> PyResult<()> {
self.inner = self.inner.clone().with_upstream_monitoring(enable);
Ok(())
}
pub fn enable_health_checker(&mut self, enable: bool) -> PyResult<()> {
self.enable_upstream_monitoring(enable)
}
pub fn round_robin_timeout(&mut self, timeout_secs: f64) -> PyResult<()> {
let duration = std::time::Duration::from_secs_f64(timeout_secs);
self.inner = self.inner.clone().with_round_robin_timeout(duration);
Ok(())
}
pub fn optimize_for_round_robin(&mut self) -> PyResult<()> {
self.inner = self.inner.clone().optimize_for_round_robin();
Ok(())
}
pub fn with_debug_logger_init(&mut self) -> PyResult<()> {
self.inner = self.inner.clone().with_debug_logger_init();
Ok(())
}
pub fn with_silent_logger_init(&mut self) -> PyResult<()> {
self.inner = self.inner.clone().with_silent_logger_init();
Ok(())
}
pub fn with_auto_logger_init(&mut self) -> PyResult<()> {
self.inner = self.inner.clone().with_auto_logger_init();
Ok(())
}
pub fn build(&self, py: Python) -> pyo3::PyResult<PyDnsResolver> {
py.allow_threads(|| {
let rt = Runtime::new().map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(
format!("Failed to create runtime: {}", e)
)
})?;
rt.block_on(async {
let resolver = self.inner.clone().build().await.map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(
format!("Failed to build resolver: {}", e)
)
})?;
Ok(PyDnsResolver::new(resolver))
})
})
}
fn __str__(&self) -> String {
"DnsResolverBuilder".to_string()
}
fn __repr__(&self) -> String {
"DnsResolverBuilder()".to_string()
}
}
impl PyDnsResolverBuilder {
pub fn retries(&mut self, count: usize) -> PyResult<()> {
self.inner = self.inner.clone().with_retry_count(count);
Ok(())
}
pub fn cache(&mut self, enable: bool) -> PyResult<()> {
self.inner = self.inner.clone().with_cache(enable);
Ok(())
}
pub fn port(&mut self, port: u16) -> PyResult<()> {
self.inner = self.inner.clone().with_port(port);
Ok(())
}
pub fn concurrent_queries(&mut self, count: usize) -> PyResult<()> {
self.inner = self.inner.clone().with_concurrent_queries(count);
Ok(())
}
pub fn recursion_desired(&mut self, enable: bool) -> PyResult<()> {
self.inner = self.inner.clone().with_recursion(enable);
Ok(())
}
pub fn buffer_size(&mut self, size: usize) -> PyResult<()> {
self.inner = self.inner.clone().with_buffer_size(size);
Ok(())
}
pub fn inner(&self) -> &RustDnsResolverBuilder {
&self.inner
}
}