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 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 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
152pub async fn pvget(opts: &PvGetOptions) -> Result<PvGetResult, PvGetError> {
154 pvget_fields(opts, &[]).await
155}
156
157pub 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 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 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}