Skip to main content

spvirit_client/
client.rs

1use tokio::io::AsyncWriteExt;
2use tokio::net::TcpStream;
3use tokio::time::timeout;
4
5use crate::auth::{resolved_authnz_host, resolved_authnz_user};
6use crate::search::resolve_pv_server;
7use crate::transport::{read_packet, read_until};
8use crate::types::{PvGetError, PvGetOptions, PvGetResult};
9use spvirit_codec::SegmentReassembler;
10use spvirit_codec::epics_decode::{
11    DecodeMode, PvaPacket, PvaPacketCommand,
12    decode_op_response_status as codec_decode_op_response_status,
13};
14use spvirit_codec::spvd_encode::encode_pv_request;
15use spvirit_codec::spvirit_encode::encode_client_connection_validation;
16pub use spvirit_codec::spvirit_encode::{
17    encode_create_channel_request, encode_get_field_request, encode_get_request,
18    encode_monitor_request, encode_put_request,
19};
20
21pub fn build_client_validation(
22    opts: &crate::types::PvGetOptions,
23    version: u8,
24    is_be: bool,
25) -> Vec<u8> {
26    let user = resolved_authnz_user(opts);
27    let host = resolved_authnz_host(opts);
28    encode_client_connection_validation(87_040, 32_767, 0, "ca", &user, &host, version, is_be)
29}
30
31pub fn op_response_status(
32    raw: &[u8],
33    is_be: bool,
34) -> Result<Option<spvirit_codec::epics_decode::PvaStatus>, PvGetError> {
35    codec_decode_op_response_status(raw, is_be).map_err(PvGetError::Protocol)
36}
37
38pub fn ensure_status_ok(raw: &[u8], is_be: bool, step: &str) -> Result<(), PvGetError> {
39    match op_response_status(raw, is_be)? {
40        None => Ok(()),
41        Some(st) if st.code == 0 => Ok(()),
42        Some(st) => Err(PvGetError::Protocol(format!(
43            "{} failed: {}",
44            step,
45            st.message.unwrap_or_else(|| format!("code={}", st.code))
46        ))),
47    }
48}
49
50pub struct ChannelConn {
51    pub stream: TcpStream,
52    pub sid: u32,
53    pub version: u8,
54    pub is_be: bool,
55    pub server_addr: std::net::SocketAddr,
56    /// Segment reassembly state for this connection.
57    ///
58    /// It lives on the connection rather than on each read call because a
59    /// message's segments may be separated by control frames — and by the
60    /// boundary between [`establish_channel`] and whatever the caller does
61    /// next on the same socket.
62    pub reassembler: SegmentReassembler,
63}
64
65pub async fn establish_channel(
66    target: std::net::SocketAddr,
67    opts: &PvGetOptions,
68) -> Result<ChannelConn, PvGetError> {
69    let mut stream = timeout(opts.timeout, TcpStream::connect(target))
70        .await
71        .map_err(|_| PvGetError::Timeout("connect"))??;
72
73    // One reassembler for the lifetime of this connection; it is handed to
74    // the caller in `ChannelConn` so pending segments survive the handshake.
75    let mut reassembler = SegmentReassembler::new();
76
77    let mut version = 2u8;
78    let mut is_be = false;
79
80    for _ in 0..2 {
81        if let Ok(bytes) = read_packet(&mut stream, opts.timeout, &mut reassembler).await {
82            let mut pkt = PvaPacket::new(&bytes);
83            if let Some(cmd) = pkt.decode_payload() {
84                match cmd {
85                    PvaPacketCommand::Control(payload) => {
86                        if payload.command == 2 {
87                            is_be = pkt.header.flags.is_msb;
88                        }
89                    }
90                    PvaPacketCommand::ConnectionValidation(_) => {
91                        version = pkt.header.version;
92                        is_be = pkt.header.flags.is_msb;
93                    }
94                    _ => {}
95                }
96            }
97        }
98    }
99
100    let validation = build_client_validation(opts, version, is_be);
101    stream.write_all(&validation).await?;
102
103    let _ = read_until(&mut stream, opts.timeout, &mut reassembler, |cmd| {
104        matches!(cmd, PvaPacketCommand::ConnectionValidated(_))
105    })
106    .await?;
107
108    let cid = 1u32;
109    let create = encode_create_channel_request(cid, &opts.pv_name, version, is_be);
110    stream.write_all(&create).await?;
111
112    let create_resp = read_until(&mut stream, opts.timeout, &mut reassembler, |cmd| {
113        matches!(cmd, PvaPacketCommand::CreateChannel(_))
114    })
115    .await?;
116    let mut pkt = PvaPacket::new(&create_resp);
117    let cmd = pkt.decode_payload().ok_or(PvGetError::Protocol(
118        "create_channel decode failed".to_string(),
119    ))?;
120    let sid = match cmd {
121        PvaPacketCommand::CreateChannel(payload) => {
122            if payload.status.as_ref().is_some_and(|s| s.is_error()) {
123                let detail = payload
124                    .status
125                    .as_ref()
126                    .map(ToString::to_string)
127                    .unwrap_or_default();
128                return Err(PvGetError::Protocol(format!(
129                    "create_channel error: {}",
130                    detail
131                )));
132            }
133            payload.sid
134        }
135        _ => {
136            return Err(PvGetError::Protocol(
137                "unexpected create_channel response".to_string(),
138            ));
139        }
140    };
141
142    Ok(ChannelConn {
143        stream,
144        sid,
145        version,
146        is_be,
147        server_addr: target,
148        reassembler,
149    })
150}
151
152/// Convenience wrapper: GET with no field filtering.
153pub async fn pvget(opts: &PvGetOptions) -> Result<PvGetResult, PvGetError> {
154    pvget_fields(opts, &[]).await
155}
156
157/// GET with optional field filtering.
158///
159/// If `fields` is empty, requests all fields (equivalent to `-r ""`).
160/// Otherwise, encodes a pvRequest like `field(value,alarm,timeStamp)`.
161pub async fn pvget_fields(opts: &PvGetOptions, fields: &[&str]) -> Result<PvGetResult, PvGetError> {
162    let target = resolve_pv_server(opts).await?;
163
164    let conn = establish_channel(target, opts).await?;
165    let ChannelConn {
166        mut stream,
167        sid,
168        version,
169        is_be,
170        mut reassembler,
171        ..
172    } = conn;
173
174    let ioid = 1u32;
175    let pv_request = if fields.is_empty() {
176        // Empty pvRequest — request all fields
177        vec![0xfd, 0x02, 0x00, 0x80, 0x00, 0x00]
178    } else {
179        encode_pv_request(fields, is_be)
180    };
181    let get_init_req = encode_get_request(sid, ioid, 0x08, &pv_request, version, is_be);
182    stream.write_all(&get_init_req).await?;
183
184    let init_resp = read_until(
185        &mut stream,
186        opts.timeout,
187        &mut reassembler,
188        |cmd| matches!(cmd, PvaPacketCommand::Op(op) if (op.subcmd & 0x08) != 0),
189    )
190    .await?;
191    let mut pkt = PvaPacket::new(&init_resp);
192    let cmd = pkt.decode_payload().ok_or(PvGetError::Protocol(
193        "get init response decode failed".to_string(),
194    ))?;
195
196    let desc = match cmd {
197        PvaPacketCommand::Op(op) => op
198            .introspection
199            .ok_or_else(|| PvGetError::Decode("missing introspection".to_string()))?,
200        _ => {
201            return Err(PvGetError::Protocol(
202                "unexpected get init response".to_string(),
203            ));
204        }
205    };
206
207    let get_data_req = encode_get_request(sid, ioid, 0x00, &[], version, is_be);
208    stream.write_all(&get_data_req).await?;
209
210    let data_resp = read_until(
211        &mut stream,
212        opts.timeout,
213        &mut reassembler,
214        |cmd| matches!(cmd, PvaPacketCommand::Op(op) if op.subcmd == 0x00),
215    )
216    .await?;
217    let mut pkt = PvaPacket::new(&data_resp);
218    let cmd = pkt.decode_payload().ok_or(PvGetError::Protocol(
219        "get data response decode failed".to_string(),
220    ))?;
221
222    match cmd {
223        PvaPacketCommand::Op(mut op) => {
224            // A decode failure leaves decoded_value None, handled below.
225            let _ = op.decode_with_field_desc(&desc, is_be, DecodeMode::Strict);
226            if let Some(value) = op.decoded_value {
227                return Ok(PvGetResult {
228                    pv_name: opts.pv_name.clone(),
229                    value,
230                    raw_pva: data_resp,
231                    raw_pvd: op.body,
232                    introspection: desc,
233                });
234            }
235            Err(PvGetError::Decode("no decoded value".to_string()))
236        }
237        _ => Err(PvGetError::Protocol(
238            "unexpected get data response".to_string(),
239        )),
240    }
241}
242
243#[cfg(test)]
244mod tests {
245    use super::*;
246    use spvirit_codec::epics_decode::{PvaPacket, PvaPacketCommand, PvaStatus};
247
248    #[test]
249    fn encode_decode_monitor_request_roundtrip() {
250        let msg =
251            encode_monitor_request(1, 2, 0x08, &[0xfd, 0x02, 0x00, 0x80, 0x00, 0x00], 2, false);
252        let mut pkt = PvaPacket::new(&msg);
253        let cmd = pkt.decode_payload().expect("decoded");
254        match cmd {
255            PvaPacketCommand::Op(op) => {
256                assert_eq!(op.command, 13);
257                assert_eq!(op.subcmd, 0x08);
258                assert_eq!(op.sid_or_cid, 1);
259                assert_eq!(op.ioid, 2);
260            }
261            other => panic!("unexpected decode: {:?}", other),
262        }
263    }
264
265    #[test]
266    fn encode_decode_put_init_roundtrip() {
267        let msg = encode_put_request(5, 6, 0x08, &[0xfd, 0x02, 0x00, 0x80, 0x00, 0x00], 2, false);
268        let mut pkt = PvaPacket::new(&msg);
269        let cmd = pkt.decode_payload().expect("decoded");
270        match cmd {
271            PvaPacketCommand::Op(op) => {
272                assert_eq!(op.command, 11);
273                assert_eq!(op.subcmd, 0x08);
274                assert_eq!(op.sid_or_cid, 5);
275                assert_eq!(op.ioid, 6);
276            }
277            other => panic!("unexpected decode: {:?}", other),
278        }
279    }
280
281    #[test]
282    fn encode_decode_get_field_request_roundtrip() {
283        let msg = encode_get_field_request(7, 1, Some("*"), 2, false);
284        let mut pkt = PvaPacket::new(&msg);
285        let cmd = pkt.decode_payload().expect("decoded");
286        match cmd {
287            PvaPacketCommand::GetField(payload) => {
288                assert!(!payload.is_server);
289                assert_eq!(payload.sid, Some(7));
290                assert_eq!(payload.ioid, Some(1));
291                assert_eq!(payload.field_name.as_deref(), Some("*"));
292            }
293            other => panic!("unexpected decode: {:?}", other),
294        }
295    }
296
297    #[test]
298    fn encode_decode_get_field_request_empty_field_roundtrip() {
299        let msg = encode_get_field_request(7, 1, None, 2, false);
300        let mut pkt = PvaPacket::new(&msg);
301        let cmd = pkt.decode_payload().expect("decoded");
302        match cmd {
303            PvaPacketCommand::GetField(payload) => {
304                assert!(!payload.is_server);
305                assert_eq!(payload.sid, Some(7));
306                assert_eq!(payload.ioid, Some(1));
307                assert_eq!(payload.field_name.as_deref(), Some(""));
308            }
309            other => panic!("unexpected decode: {:?}", other),
310        }
311    }
312
313    #[test]
314    fn pva_status_code_zero_is_not_an_error() {
315        let ok = PvaStatus {
316            code: 0,
317            message: None,
318            stack: None,
319        };
320        let err = PvaStatus {
321            code: 1,
322            message: Some("bad".to_string()),
323            stack: None,
324        };
325        assert!(!None::<&PvaStatus>.is_some_and(|s| s.is_error()));
326        assert!(!Some(&ok).is_some_and(|s| s.is_error()));
327        assert!(Some(&err).is_some_and(|s| s.is_error()));
328    }
329}