use std::net::SocketAddr;
use std::ops::ControlFlow;
use std::time::Duration;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyList};
use spvirit_client::pva_client::PvaClient;
use spvirit_client::search::{build_auto_broadcast_targets, discover_servers};
use crate::convert::{decoded_to_py, py_to_json};
use crate::errors::to_py_err;
use crate::runtime::RUNTIME;
#[pyclass(name = "GetResult")]
pub struct PyGetResult {
#[pyo3(get)]
pub pv_name: String,
value: PyObject,
#[pyo3(get)]
pub raw_pva: Vec<u8>,
#[pyo3(get)]
pub raw_pvd: Vec<u8>,
}
impl PyGetResult {
pub(crate) fn new(pv_name: String, value: PyObject, raw_pva: Vec<u8>, raw_pvd: Vec<u8>) -> Self {
Self {
pv_name,
value,
raw_pva,
raw_pvd,
}
}
}
#[pymethods]
impl PyGetResult {
#[getter]
fn value(&self, py: Python<'_>) -> PyObject {
self.value.clone_ref(py)
}
fn __repr__(&self) -> String {
format!("GetResult(pv_name={:?})", self.pv_name)
}
}
#[pyclass(name = "MonitorEvent")]
pub struct PyMonitorEvent {
#[pyo3(get)]
pub pv_name: String,
value: PyObject,
}
#[pymethods]
impl PyMonitorEvent {
#[getter]
fn value(&self, py: Python<'_>) -> PyObject {
self.value.clone_ref(py)
}
fn __repr__(&self) -> String {
format!("MonitorEvent(pv_name={:?})", self.pv_name)
}
}
#[pyclass(name = "DiscoveredServer")]
#[derive(Clone)]
pub struct PyDiscoveredServer {
#[pyo3(get)]
pub guid: Vec<u8>,
#[pyo3(get)]
pub tcp_addr: String,
}
#[pymethods]
impl PyDiscoveredServer {
fn __repr__(&self) -> String {
format!("DiscoveredServer(tcp_addr={:?})", self.tcp_addr)
}
}
#[pyclass(name = "ClientBuilder")]
pub struct PyClientBuilder {
udp_port: u16,
tcp_port: u16,
timeout_secs: f64,
no_broadcast: bool,
name_servers: Vec<String>,
authnz_user: Option<String>,
authnz_host: Option<String>,
server_addr: Option<String>,
search_addr: Option<String>,
bind_addr: Option<String>,
debug: bool,
}
#[pymethods]
impl PyClientBuilder {
#[new]
fn new() -> Self {
Self {
udp_port: 5076,
tcp_port: 5075,
timeout_secs: 5.0,
no_broadcast: false,
name_servers: Vec::new(),
authnz_user: None,
authnz_host: None,
server_addr: None,
search_addr: None,
bind_addr: None,
debug: false,
}
}
fn port(mut slf: PyRefMut<'_, Self>, port: u16) -> PyRefMut<'_, Self> {
slf.tcp_port = port;
slf
}
fn udp_port(mut slf: PyRefMut<'_, Self>, port: u16) -> PyRefMut<'_, Self> {
slf.udp_port = port;
slf
}
fn timeout(mut slf: PyRefMut<'_, Self>, secs: f64) -> PyRefMut<'_, Self> {
slf.timeout_secs = secs;
slf
}
fn no_broadcast(mut slf: PyRefMut<'_, Self>, enabled: bool) -> PyRefMut<'_, Self> {
slf.no_broadcast = enabled;
slf
}
fn name_server(mut slf: PyRefMut<'_, Self>, addr: String) -> PyResult<PyRefMut<'_, Self>> {
let _: SocketAddr = addr
.parse()
.map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("invalid address: {e}")))?;
slf.name_servers.push(addr);
Ok(slf)
}
fn authnz_user(mut slf: PyRefMut<'_, Self>, user: String) -> PyRefMut<'_, Self> {
slf.authnz_user = Some(user);
slf
}
fn authnz_host(mut slf: PyRefMut<'_, Self>, host: String) -> PyRefMut<'_, Self> {
slf.authnz_host = Some(host);
slf
}
fn server_addr(mut slf: PyRefMut<'_, Self>, addr: String) -> PyResult<PyRefMut<'_, Self>> {
let _: SocketAddr = addr
.parse()
.map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("invalid address: {e}")))?;
slf.server_addr = Some(addr);
Ok(slf)
}
fn search_addr(mut slf: PyRefMut<'_, Self>, addr: String) -> PyRefMut<'_, Self> {
slf.search_addr = Some(addr);
slf
}
fn bind_addr(mut slf: PyRefMut<'_, Self>, addr: String) -> PyRefMut<'_, Self> {
slf.bind_addr = Some(addr);
slf
}
fn debug(mut slf: PyRefMut<'_, Self>, enabled: bool) -> PyRefMut<'_, Self> {
slf.debug = enabled;
slf
}
fn build(&self) -> PyResult<PyClient> {
let mut b = PvaClient::builder()
.port(self.tcp_port)
.udp_port(self.udp_port)
.timeout(Duration::from_secs_f64(self.timeout_secs));
if self.no_broadcast {
b = b.no_broadcast();
}
for ns in &self.name_servers {
let addr: SocketAddr = ns.parse().map_err(|e| {
pyo3::exceptions::PyValueError::new_err(format!("invalid address: {e}"))
})?;
b = b.name_server(addr);
}
if let Some(ref user) = self.authnz_user {
b = b.authnz_user(user);
}
if let Some(ref host) = self.authnz_host {
b = b.authnz_host(host);
}
if let Some(ref addr) = self.server_addr {
let sa: SocketAddr = addr.parse().map_err(|e| {
pyo3::exceptions::PyValueError::new_err(format!("invalid address: {e}"))
})?;
b = b.server_addr(sa);
}
if let Some(ref addr) = self.search_addr {
let ip: std::net::IpAddr = addr.parse().map_err(|e| {
pyo3::exceptions::PyValueError::new_err(format!("invalid IP: {e}"))
})?;
b = b.search_addr(ip);
}
if let Some(ref addr) = self.bind_addr {
let ip: std::net::IpAddr = addr.parse().map_err(|e| {
pyo3::exceptions::PyValueError::new_err(format!("invalid IP: {e}"))
})?;
b = b.bind_addr(ip);
}
if self.debug {
b = b.debug();
}
Ok(PyClient {
inner: b.build(),
})
}
}
#[pyclass(name = "Client")]
pub struct PyClient {
inner: PvaClient,
}
#[pymethods]
impl PyClient {
#[new]
fn new() -> Self {
Self {
inner: PvaClient::builder().build(),
}
}
#[staticmethod]
fn builder() -> PyClientBuilder {
PyClientBuilder::new()
}
#[pyo3(signature = (pv_name, fields=None))]
fn get(
&self,
py: Python<'_>,
pv_name: String,
fields: Option<Vec<String>>,
) -> PyResult<PyGetResult> {
let client = self.inner.clone();
let result = py
.allow_threads(|| {
RUNTIME.block_on(async {
match fields {
None => client.pvget(&pv_name).await,
Some(ref f) => {
let refs: Vec<&str> = f.iter().map(String::as_str).collect();
client.pvget_fields(&pv_name, &refs).await
}
}
})
})
.map_err(to_py_err)?;
let value = decoded_to_py(py, &result.value);
Ok(PyGetResult {
pv_name: result.pv_name,
value,
raw_pva: result.raw_pva,
raw_pvd: result.raw_pvd,
})
}
#[pyo3(signature = (pv_name, value, fields=None))]
fn put(
&self,
py: Python<'_>,
pv_name: String,
value: PyObject,
fields: Option<Vec<String>>,
) -> PyResult<()> {
let json_val = py_to_json(value.bind(py))?;
let client = self.inner.clone();
py.allow_threads(|| {
RUNTIME.block_on(async {
match fields {
None => client.pvput(&pv_name, json_val).await,
Some(ref f) => {
let refs: Vec<&str> = f.iter().map(String::as_str).collect();
client.pvput_fields(&pv_name, json_val, &refs).await
}
}
})
})
.map_err(to_py_err)
}
#[pyo3(signature = (pv_name, callback, fields=None))]
fn monitor(
&self,
_py: Python<'_>,
pv_name: String,
callback: PyObject,
fields: Option<Vec<String>>,
) -> PyResult<()> {
let client = self.inner.clone();
let fields = fields.unwrap_or_default();
let result = RUNTIME.block_on(async {
let refs: Vec<&str> = fields.iter().map(String::as_str).collect();
client
.pvmonitor_fields(&pv_name, &refs, |decoded| {
let keep_going = Python::with_gil(|py| {
let py_val = decoded_to_py(py, decoded);
match callback.call1(py, (py_val,)) {
Ok(ret) => {
ret.extract::<bool>(py).unwrap_or(true)
}
Err(_) => false,
}
});
if keep_going {
ControlFlow::Continue(())
} else {
ControlFlow::Break(())
}
})
.await
});
result.map_err(to_py_err)
}
fn info(&self, py: Python<'_>, pv_name: String) -> PyResult<PyObject> {
let client = self.inner.clone();
let desc = py
.allow_threads(|| RUNTIME.block_on(client.pvinfo(&pv_name)))
.map_err(to_py_err)?;
let dict = PyDict::new(py);
dict.set_item("struct_id", &desc.struct_id)?;
let fields: Vec<PyObject> = desc
.fields
.iter()
.map(|f| {
let fd = PyDict::new(py);
fd.set_item("name", &f.name).expect("set");
fd.set_item("field_type", format!("{:?}", f.field_type))
.expect("set");
fd.into_any().unbind()
})
.collect();
dict.set_item("fields", PyList::new(py, &fields)?)?;
Ok(dict.into_any().unbind())
}
fn pvlist(&self, py: Python<'_>, server_addr: String) -> PyResult<Vec<String>> {
let addr: SocketAddr = server_addr
.parse()
.map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("invalid address: {e}")))?;
let client = self.inner.clone();
py.allow_threads(|| RUNTIME.block_on(client.pvlist(addr)))
.map_err(to_py_err)
}
}
#[pyfunction]
#[pyo3(signature = (udp_port=5076, timeout=2.0, debug=false))]
pub fn py_discover_servers(
py: Python<'_>,
udp_port: u16,
timeout: f64,
debug: bool,
) -> PyResult<Vec<PyDiscoveredServer>> {
let targets = build_auto_broadcast_targets();
let dur = Duration::from_secs_f64(timeout);
let servers = py
.allow_threads(|| RUNTIME.block_on(discover_servers(udp_port, dur, &targets, debug)))
.map_err(to_py_err)?;
Ok(servers
.into_iter()
.map(|s| PyDiscoveredServer {
guid: s.guid.to_vec(),
tcp_addr: s.tcp_addr.to_string(),
})
.collect())
}