1use std::collections::HashMap;
4use std::io::{Read, Write};
5use std::sync::{Arc, Mutex};
6
7use arrow_array::RecordBatch;
8use arrow_cast::cast_with_options;
9use arrow_schema::{Schema, SchemaRef};
10
11use crate::errors::{Result, RpcError};
12use crate::log::{LogLevel, LogMessage};
13#[cfg(feature = "shm")]
14use crate::metadata::SHM_SEGMENT_SIZE_KEY;
15use crate::metadata::{
16 CANCEL_KEY, LOG_EXTRA_KEY, LOG_LEVEL_KEY, LOG_MESSAGE_KEY, REQUEST_ID_KEY, REQUEST_VERSION,
17 REQUEST_VERSION_KEY, RPC_METHOD_KEY, SERVER_ID_KEY, SHM_OFFSET_KEY, SHM_SEGMENT_NAME_KEY,
18};
19#[cfg(feature = "shm")]
20use crate::shm::{is_shm_pointer_batch, maybe_write_to_shm, resolve_shm_batch, ShmSegment};
21
22#[cfg(not(feature = "shm"))]
24pub(crate) struct ShmSegment;
25
26#[cfg(feature = "shm")]
29fn maybe_attach_shm(req_md: &Metadata) -> Option<ShmSegment> {
30 let name = req_md.get(SHM_SEGMENT_NAME_KEY)?;
31 let size: usize = req_md.get(SHM_SEGMENT_SIZE_KEY)?.parse().ok()?;
32 match ShmSegment::attach(name, size, false) {
33 Ok(seg) => Some(seg),
34 Err(e) => {
35 tracing::warn!(target: "vgi_rpc.shm", "ignoring malformed SHM metadata ({e})");
36 None
37 }
38 }
39}
40
41#[cfg(not(feature = "shm"))]
42#[inline]
43fn maybe_attach_shm(_req_md: &Metadata) -> Option<ShmSegment> {
44 None
45}
46
47#[derive(Default)]
61pub(crate) struct ConnectionShm {
62 #[cfg(feature = "shm")]
63 name: Option<String>,
64 #[cfg(feature = "shm")]
65 segment: Option<ShmSegment>,
66}
67
68#[cfg(feature = "shm")]
69impl ConnectionShm {
70 fn refresh(&mut self, req_md: &Metadata) {
73 let Some(name) = req_md.get(SHM_SEGMENT_NAME_KEY) else {
74 return;
75 };
76 if self.name.as_deref() == Some(name.as_str()) {
77 return;
78 }
79 let Some(new) = maybe_attach_shm(req_md) else {
80 return;
81 };
82 self.segment = Some(new);
83 self.name = Some(name.clone());
84 }
85
86 fn segment(&self) -> Option<&ShmSegment> {
87 self.segment.as_ref()
88 }
89}
90
91#[cfg(not(feature = "shm"))]
92impl ConnectionShm {
93 #[inline]
98 fn segment(&self) -> Option<&ShmSegment> {
99 None
100 }
101}
102use crate::stream::{empty_schema, Emitted, OutputCollector, StreamResult, StreamStateKind};
103use crate::wire::{empty_batch, md_get, Metadata, StreamReader, StreamWriter};
104
105pub(crate) fn serialize_request_batch(batch: &RecordBatch) -> std::io::Result<Vec<u8>> {
109 let mut buf = Vec::new();
110 {
111 let mut w = arrow_ipc::writer::StreamWriter::try_new(&mut buf, batch.schema_ref())
112 .map_err(|e| std::io::Error::other(e.to_string()))?;
113 w.write(batch)
114 .map_err(|e| std::io::Error::other(e.to_string()))?;
115 w.finish()
116 .map_err(|e| std::io::Error::other(e.to_string()))?;
117 }
118 Ok(buf)
119}
120
121fn lock_ok<T>(m: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
127 m.lock().unwrap_or_else(|e| e.into_inner())
128}
129
130pub(crate) fn call_guard<T>(f: impl FnOnce() -> T) -> Result<T> {
135 std::panic::catch_unwind(std::panic::AssertUnwindSafe(f))
136 .map_err(|_| RpcError::new("RuntimeError", "handler panicked"))
137}
138
139#[derive(Clone)]
141pub struct CallContext {
142 pub server_id: String,
143 pub method: String,
144 pub request_id: String,
145 pub transport_metadata: Arc<Metadata>,
146 pub auth: crate::auth::AuthContext,
149 pub cookies: std::collections::BTreeMap<String, String>,
151 pub kind: Option<crate::transport::TransportKind>,
155 pub(crate) log_sink: Arc<Mutex<Vec<LogMessage>>>,
156 pub(crate) tick_metadata: Arc<Mutex<Metadata>>,
159 pub(crate) sticky: Option<Arc<dyn StickySink>>,
164}
165
166pub trait StickySink: Send + Sync {
171 fn accept_opens(&self) -> bool;
173 fn current_state(&self) -> Option<Arc<dyn std::any::Any + Send + Sync>>;
175 fn current_session_id(&self) -> Option<String>;
177 fn open(
179 &self,
180 state: Arc<dyn std::any::Any + Send + Sync>,
181 ttl: Option<std::time::Duration>,
182 ) -> Result<()>;
183 fn close(&self) -> Result<bool>;
185}
186
187impl CallContext {
188 pub fn client_log(&self, level: LogLevel, message: impl Into<String>) {
189 lock_ok(&self.log_sink).push(LogMessage::new(level, message));
190 }
191
192 pub fn client_log_with(&self, msg: LogMessage) {
193 lock_ok(&self.log_sink).push(msg);
194 }
195
196 pub(crate) fn drain_logs(&self) -> Vec<LogMessage> {
197 std::mem::take(&mut *lock_ok(&self.log_sink))
198 }
199
200 pub fn tick_metadata(&self, key: &str) -> Option<String> {
203 lock_ok(&self.tick_metadata).get(key).cloned()
204 }
205
206 #[cfg(feature = "http")]
211 pub(crate) fn set_tick_metadata(&self, md: Metadata) {
212 *lock_ok(&self.tick_metadata) = md;
213 }
214
215 pub(crate) fn for_request(server: &RpcServer, req: &Request) -> Self {
220 Self {
221 server_id: server.server_id.clone(),
222 method: req.method.clone(),
223 request_id: req.request_id.clone(),
224 transport_metadata: req.metadata.clone(),
225 auth: crate::auth::AuthContext::anonymous(),
226 cookies: std::collections::BTreeMap::new(),
227 kind: server.transport_kind(),
228 log_sink: Arc::new(Mutex::new(Vec::new())),
229 tick_metadata: Arc::new(Mutex::new(Metadata::default())),
230 sticky: None,
231 }
232 }
233
234 #[cfg(feature = "http")]
239 pub(crate) fn with_auth_cookies(
240 server: &RpcServer,
241 req: &Request,
242 auth: crate::auth::AuthContext,
243 cookies: std::collections::BTreeMap<String, String>,
244 ) -> Self {
245 Self {
246 server_id: server.server_id.clone(),
247 method: req.method.clone(),
248 request_id: req.request_id.clone(),
249 transport_metadata: req.metadata.clone(),
250 auth,
251 cookies,
252 kind: server.transport_kind(),
253 log_sink: Arc::new(Mutex::new(Vec::new())),
254 tick_metadata: Arc::new(Mutex::new(Metadata::default())),
255 sticky: None,
256 }
257 }
258
259 #[cfg(feature = "http")]
264 pub(crate) fn set_sticky(&mut self, sink: Arc<dyn StickySink>) {
265 self.sticky = Some(sink);
266 }
267
268 pub fn session<T: std::any::Any + Send + Sync>(&self) -> Option<Arc<T>> {
276 let state = self.sticky.as_ref()?.current_state()?;
277 state.downcast::<T>().ok()
278 }
279
280 pub fn session_id(&self) -> Option<String> {
283 self.sticky.as_ref()?.current_session_id()
284 }
285
286 pub fn open_session(
297 &self,
298 state: Arc<dyn std::any::Any + Send + Sync>,
299 ttl: Option<std::time::Duration>,
300 ) -> Result<()> {
301 let sink = self.sticky.as_ref().ok_or_else(|| {
302 RpcError::runtime_error("sticky sessions not available on this transport")
303 })?;
304 if !sink.accept_opens() {
305 return Err(RpcError::runtime_error(
306 "client did not opt in to sticky sessions \
307 (missing VGI-Session-Accept: true header — open the call inside \
308 an HttpConnection.with_session_token() block)",
309 ));
310 }
311 if sink.current_state().is_some() {
312 return Err(RpcError::runtime_error(
313 "a sticky session is already active for this request",
314 ));
315 }
316 sink.open(state, ttl)
317 }
318
319 pub fn close_session(&self) -> Result<()> {
322 let sink = self.sticky.as_ref().ok_or_else(|| {
323 RpcError::runtime_error("sticky sessions not available on this transport")
324 })?;
325 sink.close()?;
326 Ok(())
327 }
328}
329
330pub struct Request {
332 pub method: String,
333 pub request_id: String,
334 pub batch: RecordBatch,
335 pub metadata: Arc<Metadata>,
339}
340
341impl Request {
342 pub fn column(&self, name: &str) -> Option<&dyn arrow_array::Array> {
343 let idx = self.batch.schema().index_of(name).ok()?;
344 Some(self.batch.column(idx).as_ref())
345 }
346
347 pub(crate) fn from_read_batch(
355 batch: RecordBatch,
356 metadata: Metadata,
357 require_method: bool,
358 ) -> Result<Self> {
359 let method = if require_method {
360 md_get(&metadata, RPC_METHOD_KEY)
361 .ok_or_else(|| {
362 RpcError::protocol_error(
363 "Missing 'vgi_rpc.method' in request batch custom_metadata.",
364 )
365 })?
366 .to_string()
367 } else {
368 md_get(&metadata, RPC_METHOD_KEY).unwrap_or("").to_string()
369 };
370 let version = md_get(&metadata, REQUEST_VERSION_KEY).ok_or_else(|| {
371 RpcError::version_error(format!(
372 "Missing 'vgi_rpc.request_version' in request batch custom_metadata. Set it to {:?}.",
373 REQUEST_VERSION
374 ))
375 })?;
376 if version != REQUEST_VERSION {
377 return Err(RpcError::version_error(format!(
378 "Unsupported request version {:?}, expected {:?}.",
379 version, REQUEST_VERSION
380 )));
381 }
382 if require_method && !batch.schema().fields().is_empty() && batch.num_rows() != 1 {
383 return Err(RpcError::protocol_error(format!(
384 "Expected 1 row in request batch, got {}",
385 batch.num_rows()
386 )));
387 }
388 let request_id = md_get(&metadata, REQUEST_ID_KEY).unwrap_or("").to_string();
389 Ok(Request {
390 method,
391 request_id,
392 batch,
393 metadata: Arc::new(metadata),
394 })
395 }
396}
397
398#[derive(Clone, Copy, Debug, PartialEq, Eq)]
400pub enum MethodType {
401 Unary,
402 Producer,
403 Exchange,
404 Dynamic,
406}
407
408pub type UnaryHandler =
410 Arc<dyn Fn(&Request, &CallContext) -> Result<Option<RecordBatch>> + Send + Sync>;
411
412pub type StreamHandler = Arc<dyn Fn(&Request, &CallContext) -> Result<StreamResult> + Send + Sync>;
414
415#[derive(Default)]
417pub struct RpcServerBuilder {
418 server_id: Option<String>,
419 server_version: Option<String>,
420 protocol_name: Option<String>,
421 protocol_version: Option<String>,
422 enable_describe: bool,
423 dispatch_hook: Option<Arc<dyn crate::hooks::DispatchHook>>,
424 on_serve_start: Option<crate::transport::ServeStartHook>,
425 #[cfg(feature = "http")]
426 external_config: Option<Arc<crate::external::ExternalLocationConfig>>,
427}
428
429impl RpcServerBuilder {
430 pub fn server_id(mut self, id: impl Into<String>) -> Self {
431 self.server_id = Some(id.into());
432 self
433 }
434
435 pub fn server_version(mut self, v: impl Into<String>) -> Self {
436 self.server_version = Some(v.into());
437 self
438 }
439
440 pub fn protocol_name(mut self, name: impl Into<String>) -> Self {
441 self.protocol_name = Some(name.into());
442 self
443 }
444
445 pub fn protocol_version(mut self, v: impl Into<String>) -> Self {
449 self.protocol_version = Some(v.into());
450 self
451 }
452
453 pub fn enable_describe(mut self, enabled: bool) -> Self {
454 self.enable_describe = enabled;
455 self
456 }
457
458 pub fn with_hook(mut self, hook: Arc<dyn crate::hooks::DispatchHook>) -> Self {
459 self.dispatch_hook = Some(hook);
460 self
461 }
462
463 pub fn on_serve_start(mut self, hook: crate::transport::ServeStartHook) -> Self {
473 self.on_serve_start = Some(hook);
474 self
475 }
476
477 #[cfg(feature = "http")]
481 pub fn with_external_location(mut self, cfg: crate::external::ExternalLocationConfig) -> Self {
482 self.external_config = Some(Arc::new(cfg));
483 self
484 }
485
486 pub fn build(self) -> RpcServer {
487 RpcServer {
488 methods: HashMap::new(),
489 server_id: self.server_id.unwrap_or_else(crate::util::short_random_id),
490 server_version: self.server_version.unwrap_or_default(),
491 protocol_name: self.protocol_name.unwrap_or_default(),
492 protocol_version: self.protocol_version.unwrap_or_default(),
493 protocol_hash: std::sync::OnceLock::new(),
494 describe_enabled: self.enable_describe,
495 dispatch_hook: self.dispatch_hook,
496 on_serve_start: self.on_serve_start,
497 transport_state: Mutex::new(None),
498 #[cfg(feature = "http")]
499 external_config: self.external_config,
500 }
501 }
502}
503
504pub struct MethodInfo {
511 pub name: String,
512 pub method_type: MethodType,
513 pub params_schema: SchemaRef,
515 pub result_schema: SchemaRef,
517 pub header_schema: Option<SchemaRef>,
519 pub doc: Option<String>,
521 pub param_types: Vec<(String, String)>,
524 pub param_defaults: Vec<(String, serde_json::Value)>,
526 pub param_docs: Vec<(String, String)>,
528 pub has_return: bool,
530 pub unary: Option<UnaryHandler>,
531 pub stream: Option<StreamHandler>,
532 pub state_decoder: Option<StateDecoder>,
538}
539
540pub type StateDecoder = Arc<dyn Fn(&[u8]) -> Result<crate::stream::StreamStateKind> + Send + Sync>;
543
544impl MethodInfo {
545 pub fn unary(
547 name: impl Into<String>,
548 params_schema: SchemaRef,
549 result_schema: SchemaRef,
550 handler: impl Fn(&Request, &CallContext) -> Result<Option<RecordBatch>> + Send + Sync + 'static,
551 ) -> Self {
552 let has_return = !result_schema.fields().is_empty();
553 Self {
554 name: name.into(),
555 method_type: MethodType::Unary,
556 params_schema,
557 result_schema,
558 header_schema: None,
559 doc: None,
560 param_types: Vec::new(),
561 param_defaults: Vec::new(),
562 param_docs: Vec::new(),
563 has_return,
564 unary: Some(Arc::new(handler)),
565 stream: None,
566 state_decoder: None,
567 }
568 }
569
570 pub fn stream(
580 name: impl Into<String>,
581 method_type: MethodType,
582 params_schema: SchemaRef,
583 handler: impl Fn(&Request, &CallContext) -> Result<StreamResult> + Send + Sync + 'static,
584 ) -> Self {
585 assert!(
586 matches!(
587 method_type,
588 MethodType::Producer | MethodType::Exchange | MethodType::Dynamic
589 ),
590 "stream methods must be Producer / Exchange / Dynamic"
591 );
592 Self {
593 name: name.into(),
594 method_type,
595 params_schema,
596 result_schema: empty_schema(),
597 header_schema: None,
598 doc: None,
599 param_types: Vec::new(),
600 param_defaults: Vec::new(),
601 param_docs: Vec::new(),
602 has_return: false,
603 unary: None,
604 stream: Some(Arc::new(handler)),
605 state_decoder: None,
606 }
607 }
608
609 pub fn with_state_decoder(mut self, decoder: StateDecoder) -> Self {
611 self.state_decoder = Some(decoder);
612 self
613 }
614
615 pub fn doc(mut self, s: impl Into<String>) -> Self {
616 self.doc = Some(s.into());
617 self
618 }
619
620 pub fn param_type(mut self, param: impl Into<String>, ty: impl Into<String>) -> Self {
621 self.param_types.push((param.into(), ty.into()));
622 self
623 }
624
625 pub fn param_default(mut self, param: impl Into<String>, value: serde_json::Value) -> Self {
626 self.param_defaults.push((param.into(), value));
627 self
628 }
629
630 pub fn param_doc(mut self, param: impl Into<String>, doc: impl Into<String>) -> Self {
631 self.param_docs.push((param.into(), doc.into()));
632 self
633 }
634
635 pub fn header_schema(mut self, schema: SchemaRef) -> Self {
636 self.header_schema = Some(schema);
637 self
638 }
639}
640
641pub struct RpcServer {
643 methods: HashMap<String, MethodInfo>,
644 pub server_id: String,
645 pub(crate) server_version: String,
646 pub(crate) protocol_name: String,
647 pub(crate) protocol_version: String,
648 pub(crate) protocol_hash: std::sync::OnceLock<String>,
649 pub(crate) describe_enabled: bool,
650 pub(crate) dispatch_hook: Option<Arc<dyn crate::hooks::DispatchHook>>,
651 on_serve_start: Option<crate::transport::ServeStartHook>,
654 transport_state: Mutex<
657 Option<(
658 crate::transport::TransportKind,
659 crate::transport::TransportCapabilities,
660 )>,
661 >,
662 #[cfg(feature = "http")]
663 pub(crate) external_config: Option<Arc<crate::external::ExternalLocationConfig>>,
664}
665
666impl RpcServer {
667 pub fn new(server_id: impl Into<String>) -> Self {
669 Self::builder().server_id(server_id).build()
670 }
671
672 pub fn builder() -> RpcServerBuilder {
674 RpcServerBuilder::default()
675 }
676
677 pub fn protocol_name(&self) -> &str {
678 &self.protocol_name
679 }
680
681 pub fn describe_enabled(&self) -> bool {
682 self.describe_enabled
683 }
684
685 pub fn server_version(&self) -> &str {
686 &self.server_version
687 }
688
689 pub fn protocol_version(&self) -> &str {
690 &self.protocol_version
691 }
692
693 pub fn protocol_hash(&self) -> &str {
696 self.protocol_hash.get_or_init(|| {
697 match crate::introspect::build_describe(
698 &self.protocol_name,
699 &self.methods,
700 &self.server_id,
701 &self.protocol_version,
702 ) {
703 Ok((_, md)) => md
704 .get(crate::metadata::PROTOCOL_HASH_KEY)
705 .cloned()
706 .unwrap_or_default(),
707 Err(_) => String::new(),
708 }
709 })
710 }
711
712 #[cfg(feature = "http")]
713 pub fn external_config(&self) -> Option<&Arc<crate::external::ExternalLocationConfig>> {
714 self.external_config.as_ref()
715 }
716
717 pub fn transport_kind(&self) -> Option<crate::transport::TransportKind> {
721 lock_ok(&self.transport_state).as_ref().map(|(k, _)| *k)
722 }
723
724 pub fn transport_capabilities(&self) -> crate::transport::TransportCapabilities {
728 lock_ok(&self.transport_state)
729 .as_ref()
730 .map(|(_, c)| *c)
731 .unwrap_or_default()
732 }
733
734 pub fn notify_transport(
746 &self,
747 kind: crate::transport::TransportKind,
748 caps: crate::transport::TransportCapabilities,
749 ) {
750 let hook = {
751 let mut guard = lock_ok(&self.transport_state);
752 if let Some((cur_kind, cur_caps)) = guard.as_ref() {
753 if *cur_kind == kind && *cur_caps == caps {
754 return;
755 }
756 }
757 *guard = Some((kind, caps));
758 self.on_serve_start.clone()
759 };
760 if let Some(h) = hook {
761 h(kind, &caps);
762 }
763 }
764
765 pub fn register(&mut self, info: MethodInfo) {
767 self.methods.insert(info.name.clone(), info);
768 }
769
770 pub fn register_unary(
774 &mut self,
775 name: impl Into<String>,
776 result_schema: SchemaRef,
777 handler: impl Fn(&Request, &CallContext) -> Result<Option<RecordBatch>> + Send + Sync + 'static,
778 ) {
779 self.register(MethodInfo::unary(
780 name,
781 empty_schema(),
782 result_schema,
783 handler,
784 ));
785 }
786
787 pub fn register_stream(
791 &mut self,
792 name: impl Into<String>,
793 method_type: MethodType,
794 handler: impl Fn(&Request, &CallContext) -> Result<StreamResult> + Send + Sync + 'static,
795 ) {
796 self.register(MethodInfo::stream(
797 name,
798 method_type,
799 empty_schema(),
800 handler,
801 ));
802 }
803
804 pub fn method(&self, name: &str) -> Option<&MethodInfo> {
805 self.methods.get(name)
806 }
807
808 pub fn methods(&self) -> &HashMap<String, MethodInfo> {
809 &self.methods
810 }
811
812 pub fn method_names(&self) -> Vec<&str> {
813 self.sorted_method_names()
814 }
815
816 pub fn sorted_method_names(&self) -> Vec<&str> {
819 let mut names: Vec<_> = self.methods.keys().map(String::as_str).collect();
820 names.sort();
821 names
822 }
823
824 pub fn serve<R: Read, W: Write>(&self, mut r: R, mut w: W) {
835 let mut conn_shm = ConnectionShm::default();
839 loop {
840 match self.serve_one_conn(&mut r, &mut w, Some(&mut conn_shm)) {
841 Ok(keep_going) => {
842 if !keep_going {
843 return;
844 }
845 }
846 Err(e) => {
847 tracing::warn!(
853 target: "vgi_rpc.server",
854 error = %e,
855 "serve loop terminating connection on error"
856 );
857 return;
858 }
859 }
860 }
861 }
862
863 pub fn serve_with_shutdown<R, W, F>(&self, mut r: R, mut w: W, shutdown: F)
869 where
870 R: Read,
871 W: Write,
872 F: Fn() -> bool,
873 {
874 let mut conn_shm = ConnectionShm::default();
875 loop {
876 if shutdown() {
877 return;
878 }
879 match self.serve_one_conn(&mut r, &mut w, Some(&mut conn_shm)) {
880 Ok(true) => {}
881 _ => return,
882 }
883 }
884 }
885
886 pub fn serve_one<R: Read, W: Write>(&self, r: &mut R, w: &mut W) -> Result<bool> {
893 self.serve_one_conn(r, w, None)
894 }
895
896 fn serve_one_conn<R: Read, W: Write>(
897 &self,
898 r: &mut R,
899 w: &mut W,
900 shm_cache: Option<&mut ConnectionShm>,
901 ) -> Result<bool> {
902 let result = self._serve_one(r, w, shm_cache);
903 let _ = w.flush();
904 result
905 }
906
907 fn _serve_one<R: Read, W: Write>(
908 &self,
909 r: &mut R,
910 w: &mut W,
911 mut shm_cache: Option<&mut ConnectionShm>,
912 ) -> Result<bool> {
913 let (req, request_used_shm) = match self.read_request(r, shm_cache.as_deref_mut())? {
914 Some(rq) => rq,
915 None => return Ok(false),
916 };
917
918 if req.method == crate::transport_options::TRANSPORT_OPTIONS_METHOD_NAME {
925 let mut md = crate::transport_options::worker_transport_metadata();
926 md.insert(REQUEST_VERSION_KEY.to_string(), REQUEST_VERSION.to_string());
927 md.insert(SERVER_ID_KEY.to_string(), self.server_id.clone());
928 let schema = empty_schema();
929 let batch = empty_batch(&schema)?;
930 let mut sw = StreamWriter::new(w, &schema)?;
931 sw.write(&batch, Some(&md))?;
932 sw.finish()?;
933 return Ok(true);
934 }
935
936 if !self.protocol_version.is_empty() {
940 if let Some(client_v) = md_get(&req.metadata, crate::metadata::PROTOCOL_VERSION_KEY) {
941 let major = |v: &str| v.split('.').next().unwrap_or("").to_string();
942 if major(client_v) != major(&self.protocol_version) {
943 let err = RpcError::version_error(format!(
944 "protocol_version mismatch: client {:?} is incompatible with server {:?}",
945 client_v, self.protocol_version
946 ));
947 write_error_stream(w, &empty_schema(), &err, &self.server_id, &req.request_id)?;
948 return Ok(true);
949 }
950 }
951 }
952
953 let ctx = CallContext::for_request(self, &req);
954
955 let stats = Arc::new(Mutex::new(crate::hooks::CallStatistics::default()));
956 {
958 let mut s = lock_ok(&stats);
959 s.input_batches = 1;
960 s.input_rows = req.batch.num_rows() as u64;
961 }
962
963 if self.describe_enabled && req.method == crate::introspect::DESCRIBE_METHOD_NAME {
965 match crate::introspect::build_describe(
966 &self.protocol_name,
967 &self.methods,
968 &self.server_id,
969 &self.protocol_version,
970 ) {
971 Ok((batch, md)) => {
972 crate::introspect::write_describe_response(w, &batch, &md)?;
973 }
974 Err(err) => {
975 write_error_stream(w, &empty_schema(), &err, &self.server_id, &req.request_id)?;
976 }
977 }
978 return Ok(true);
979 }
980
981 let Some(info) = self.methods.get(&req.method) else {
982 let names = self.sorted_method_names();
983 let msg = format!(
984 "Unknown method: '{}'. Available methods: {:?}",
985 req.method, names
986 );
987 write_error_stream(
988 w,
989 &empty_schema(),
990 &RpcError::attribute_error(msg),
991 &self.server_id,
992 &req.request_id,
993 )?;
994 return Ok(true);
995 };
996
997 let method_type = match info.method_type {
998 MethodType::Unary => "unary",
999 _ => "stream",
1000 };
1001 #[cfg_attr(not(feature = "http"), allow(unused_mut))]
1007 let mut dispatch_info = self.dispatch_hook.as_ref().map(|_| {
1008 let mut di =
1009 crate::hooks::DispatchInfo::from_request(self, &req, method_type, &ctx.auth);
1010 if let Ok(bytes) = serialize_request_batch(&req.batch) {
1014 di.request_data = bytes;
1015 }
1016 if method_type == "stream" {
1017 di.stream_id = crate::access_log::random_stream_id();
1018 }
1019 di
1020 });
1021 let hook_token = match (self.dispatch_hook.as_ref(), dispatch_info.as_ref()) {
1022 (Some(h), Some(di)) => Some(h.on_dispatch_start(di)),
1023 _ => None,
1024 };
1025
1026 let mut app_err: Option<RpcError> = None;
1027 let dynamic_shm: Option<ShmSegment>;
1039 let shm_ref: Option<&ShmSegment> = match shm_cache {
1040 Some(cache) => {
1041 if request_used_shm {
1042 cache.segment()
1043 } else {
1044 None
1045 }
1046 }
1047 None => {
1048 dynamic_shm = maybe_attach_shm(&req.metadata);
1049 dynamic_shm.as_ref()
1050 }
1051 };
1052 #[cfg(feature = "http")]
1057 let externalized = crate::external::ExternalizedScope::new();
1058 match info.method_type {
1059 MethodType::Unary => {
1060 self.serve_unary(w, &req, info, &ctx, &stats, &mut app_err, shm_ref)?
1061 }
1062 MethodType::Producer | MethodType::Exchange | MethodType::Dynamic => {
1063 self.serve_stream(r, w, &req, info, &ctx, &stats, &mut app_err, shm_ref)?
1064 }
1065 }
1066 #[cfg(feature = "http")]
1067 let externalized_bytes = externalized.finish();
1068 #[cfg(feature = "http")]
1069 if let Some(di) = dispatch_info.as_mut() {
1070 di.externalized_bytes = externalized_bytes;
1071 }
1072 if let (Some(hook), Some(di)) = (self.dispatch_hook.as_ref(), dispatch_info.as_ref()) {
1077 let token = hook_token.unwrap_or(0);
1078 let final_stats = lock_ok(&stats).clone();
1079 hook.on_dispatch_end(token, di, app_err.as_ref(), &final_stats);
1080 }
1081 Ok(true)
1082 }
1083
1084 fn read_request<R: Read>(
1089 &self,
1090 r: &mut R,
1091 shm_cache: Option<&mut ConnectionShm>,
1092 ) -> Result<Option<(Request, bool)>> {
1093 let mut reader = match StreamReader::new(r) {
1094 Ok(r) => r,
1095 Err(e) => {
1096 let msg = e.message.to_lowercase();
1098 if msg.contains("empty ipc stream") || msg.contains("eof") {
1099 return Ok(None);
1100 }
1101 return Err(e);
1102 }
1103 };
1104 let (batch, metadata) = match reader.read_next()? {
1105 Some(b) => b,
1106 None => return Ok(None),
1107 };
1108 reader.drain()?;
1109 let request_used_shm =
1112 metadata.contains_key(SHM_OFFSET_KEY) || metadata.contains_key(SHM_SEGMENT_NAME_KEY);
1113 #[cfg(feature = "shm")]
1121 let (batch, metadata) = if is_shm_pointer_batch(&batch, &metadata) {
1122 let one_shot: Option<ShmSegment>;
1123 let seg: Option<&ShmSegment> = match shm_cache {
1124 Some(cache) => {
1125 cache.refresh(&metadata);
1126 cache.segment()
1127 }
1128 None => {
1129 one_shot = maybe_attach_shm(&metadata);
1130 one_shot.as_ref()
1131 }
1132 };
1133 let resolved = resolve_shm_batch(batch, metadata, seg)?;
1134 if let (Some(off), Some(seg)) = (resolved.release_offset, seg) {
1137 let _ = seg.free(off);
1138 }
1139 (resolved.batch, resolved.metadata)
1140 } else {
1141 if let Some(cache) = shm_cache {
1142 cache.refresh(&metadata);
1143 }
1144 (batch, metadata)
1145 };
1146 #[cfg(not(feature = "shm"))]
1147 let _ = shm_cache;
1148 Ok(Some((
1149 Request::from_read_batch(batch, metadata, true)?,
1150 request_used_shm,
1151 )))
1152 }
1153
1154 #[allow(clippy::too_many_arguments)]
1155 fn serve_unary<W: Write>(
1156 &self,
1157 w: &mut W,
1158 req: &Request,
1159 info: &MethodInfo,
1160 ctx: &CallContext,
1161 stats: &Arc<Mutex<crate::hooks::CallStatistics>>,
1162 app_err: &mut Option<RpcError>,
1163 #[cfg_attr(not(feature = "shm"), allow(unused_variables))] shm: Option<&ShmSegment>,
1164 ) -> Result<()> {
1165 let result = call_guard(|| (info.unary.as_ref().unwrap())(req, ctx)).and_then(|r| r);
1169 let logs = ctx.drain_logs();
1170 let mut envelope = EnvelopeMeta::new(&self.server_id, &req.request_id);
1171 match result {
1172 Ok(maybe_batch) => {
1173 let mut sw = StreamWriter::new(w, &info.result_schema)?;
1174 for log in logs {
1175 let md = envelope.log(&log);
1176 sw.write(&empty_batch(&info.result_schema)?, Some(md))?;
1177 }
1178 let out_batch = match maybe_batch {
1179 Some(b) => b,
1180 None => empty_batch(&info.result_schema)?,
1181 };
1182 {
1183 let mut s = lock_ok(stats);
1184 s.output_batches = 1;
1185 s.output_rows = out_batch.num_rows() as u64;
1186 }
1187 #[cfg(feature = "shm")]
1188 if let Some(seg) = shm {
1189 let (written, written_md) =
1190 maybe_write_to_shm(out_batch.clone(), Metadata::new(), Some(seg))?;
1191 if written_md.contains_key(crate::metadata::SHM_OFFSET_KEY) {
1192 sw.write(&written, Some(&written_md))?;
1193 sw.finish()?;
1194 return Ok(());
1195 }
1196 }
1197 #[cfg(feature = "http")]
1198 if let Some(cfg) = self.external_config.as_ref() {
1199 if let Ok(Some((ptr, md))) = crate::external::maybe_externalize_batch(
1204 &out_batch,
1205 &info.result_schema,
1206 None,
1207 cfg,
1208 ) {
1209 sw.write(&ptr, Some(&md))?;
1210 sw.finish()?;
1211 return Ok(());
1212 }
1213 }
1214 #[cfg(not(feature = "shm"))]
1215 let _ = shm;
1216 sw.write(&out_batch, None)?;
1217 sw.finish()?;
1218 }
1219 Err(err) => {
1220 let mut sw = StreamWriter::new(w, &info.result_schema)?;
1221 for log in logs {
1222 let md = envelope.log(&log);
1223 sw.write(&empty_batch(&info.result_schema)?, Some(md))?;
1224 }
1225 let md = envelope.error(&err);
1226 sw.write(&empty_batch(&info.result_schema)?, Some(md))?;
1227 sw.finish()?;
1228 *app_err = Some(err);
1229 }
1230 }
1231 Ok(())
1232 }
1233
1234 #[allow(clippy::too_many_arguments)]
1235 #[allow(clippy::too_many_arguments)]
1236 fn serve_stream<R: Read, W: Write>(
1237 &self,
1238 r: &mut R,
1239 w: &mut W,
1240 req: &Request,
1241 info: &MethodInfo,
1242 ctx: &CallContext,
1243 stats: &Arc<Mutex<crate::hooks::CallStatistics>>,
1244 app_err: &mut Option<RpcError>,
1245 #[cfg_attr(not(feature = "shm"), allow(unused_variables))] shm: Option<&ShmSegment>,
1246 ) -> Result<()> {
1247 let init_result = call_guard(|| (info.stream.as_ref().unwrap())(req, ctx)).and_then(|r| r);
1248 let init_logs = ctx.drain_logs();
1249 let stream = match init_result {
1250 Ok(s) => s,
1251 Err(err) => {
1252 let output_schema = info.result_schema.clone();
1254 let mut sw = StreamWriter::new(w, &output_schema)?;
1255 let mut envelope = EnvelopeMeta::new(&self.server_id, &req.request_id);
1256 for log in init_logs {
1257 let md = envelope.log(&log);
1258 sw.write(&empty_batch(&output_schema)?, Some(md))?;
1259 }
1260 let md = envelope.error(&err);
1261 sw.write(&empty_batch(&output_schema)?, Some(md))?;
1262 sw.finish()?;
1263 let _ = drain_input(r);
1266 *app_err = Some(err);
1267 return Ok(());
1268 }
1269 };
1270
1271 let StreamResult {
1272 output_schema,
1273 input_schema,
1274 state,
1275 header,
1276 header_metadata,
1277 } = stream;
1278
1279 let mut envelope = EnvelopeMeta::new(&self.server_id, &req.request_id);
1283
1284 let wrote_header = header.is_some();
1286 if let Some(header_batch) = header {
1287 let mut hw = StreamWriter::new(&mut *w, header_batch.schema().as_ref())?;
1288 for log in &init_logs {
1289 let md = envelope.log(log);
1290 hw.write(&empty_batch(header_batch.schema().as_ref())?, Some(md))?;
1291 }
1292 hw.write(&header_batch, header_metadata.as_ref())?;
1293 hw.finish()?;
1294 }
1295 let _ = w.flush();
1296
1297 let mut out_writer = StreamWriter::new(&mut *w, output_schema.as_ref())?;
1301 out_writer.flush()?;
1302
1303 let mut input_reader = StreamReader::new(&mut *r)?;
1305
1306 let empty_out = empty_batch(output_schema.as_ref())?;
1310
1311 if !wrote_header {
1313 for log in &init_logs {
1314 let md = envelope.log(log);
1315 out_writer.write(&empty_out, Some(md))?;
1316 }
1317 }
1318 let _ = header_metadata;
1319
1320 let mut state = state;
1321 let mut cancelled = false;
1322
1323 'lockstep: loop {
1324 let read = match input_reader.read_next() {
1325 Ok(x) => x,
1326 Err(_) => break,
1327 };
1328 let Some((input_batch, input_md)) = read else {
1329 break;
1330 };
1331
1332 #[cfg(feature = "shm")]
1337 let (input_batch, input_md) = {
1338 let resolved = resolve_shm_batch(input_batch, input_md, shm)?;
1339 if let (Some(off), Some(seg)) = (resolved.release_offset, shm) {
1340 let _ = seg.free(off);
1341 }
1342 (resolved.batch, resolved.metadata)
1343 };
1344
1345 {
1346 let mut s = lock_ok(stats);
1347 s.input_batches += 1;
1348 s.input_rows += input_batch.num_rows() as u64;
1349 }
1350
1351 let is_cancel = md_get(&input_md, CANCEL_KEY).is_some();
1353
1354 *lock_ok(&ctx.tick_metadata) = input_md;
1359
1360 if is_cancel {
1362 cancelled = true;
1363 match &mut state {
1364 StreamStateKind::Producer(p) => p.on_cancel(ctx),
1365 StreamStateKind::Exchange(e) => e.on_cancel(ctx),
1366 }
1367 break;
1368 }
1369
1370 let casted = match &input_schema {
1372 Some(expected) if input_batch.schema() != *expected => {
1373 match cast_batch(&input_batch, expected) {
1374 Ok(b) => b,
1375 Err(e) => {
1376 let md = envelope.error(&e);
1377 out_writer.write(&empty_out, Some(md))?;
1378 break 'lockstep;
1379 }
1380 }
1381 }
1382 _ => input_batch,
1383 };
1384
1385 let mut out = OutputCollector::new(output_schema.clone(), input_schema.is_none());
1386
1387 let iter_result = call_guard(|| match &mut state {
1388 StreamStateKind::Producer(p) => p.produce(&mut out, ctx),
1389 StreamStateKind::Exchange(e) => e.exchange(&casted, &mut out, ctx),
1390 })
1391 .and_then(|r| r);
1392
1393 let iter_logs = ctx.drain_logs();
1395 for log in iter_logs {
1396 let md = envelope.log(&log);
1397 out_writer.write(&empty_out, Some(md))?;
1398 }
1399
1400 if let Err(err) = iter_result {
1401 let md = envelope.error(&err);
1402 out_writer.write(&empty_out, Some(md))?;
1403 *app_err = Some(err);
1404 break;
1405 }
1406
1407 let finished = out.finished();
1408
1409 for item in out.items.drain(..) {
1411 match item {
1412 Emitted::Log(log) => {
1413 let md = envelope.log(&log);
1414 out_writer.write(&empty_out, Some(md))?;
1415 }
1416 Emitted::Batch { batch, metadata } => {
1417 {
1418 let mut s = lock_ok(stats);
1419 s.output_batches += 1;
1420 s.output_rows += batch.num_rows() as u64;
1421 }
1422 #[cfg(feature = "shm")]
1423 if let Some(seg) = shm {
1424 let md_in = metadata.clone().unwrap_or_default();
1425 let (written, written_md) =
1426 maybe_write_to_shm(batch.clone(), md_in, Some(seg))?;
1427 if written_md.contains_key(crate::metadata::SHM_OFFSET_KEY) {
1428 out_writer.write(&written, Some(&written_md))?;
1429 continue;
1430 }
1431 }
1432 #[cfg(feature = "http")]
1433 if let Some(cfg) = self.external_config.as_ref() {
1434 match crate::external::maybe_externalize_batch(
1435 &batch,
1436 output_schema.as_ref(),
1437 metadata.as_ref(),
1438 cfg,
1439 ) {
1440 Ok(Some((ptr, md))) => {
1441 out_writer.write(&ptr, Some(&md))?;
1442 continue;
1443 }
1444 Ok(None) => {}
1445 Err(e) => {
1446 *app_err = Some(e);
1449 }
1450 }
1451 }
1452 out_writer.write(&batch, metadata.as_ref())?;
1453 }
1454 }
1455 }
1456 out_writer.flush()?;
1459
1460 if finished {
1461 break;
1462 }
1463 }
1464 let _ = cancelled;
1465 out_writer.finish()?;
1466
1467 let _ = input_reader.drain();
1469 Ok(())
1470 }
1471}
1472
1473fn drain_input<R: Read>(r: &mut R) -> Result<()> {
1474 let mut rdr = StreamReader::new(r)?;
1475 rdr.drain()?;
1476 Ok(())
1477}
1478
1479pub(crate) fn cast_batch(batch: &RecordBatch, target: &SchemaRef) -> Result<RecordBatch> {
1480 if batch.num_columns() != target.fields().len() {
1481 return Err(RpcError::type_error(format!(
1482 "Input schema mismatch: expected {} fields, got {}",
1483 target.fields().len(),
1484 batch.num_columns()
1485 )));
1486 }
1487 let src_schema = batch.schema();
1488 for (i, field) in target.fields().iter().enumerate() {
1489 let src_name = src_schema.field(i).name();
1490 if src_name != field.name() {
1491 return Err(RpcError::type_error(format!(
1492 "Input schema mismatch: expected field {:?}, got {:?}",
1493 field.name(),
1494 src_name
1495 )));
1496 }
1497 }
1498 let opts = arrow_cast::CastOptions::default();
1499 let mut cols = Vec::with_capacity(batch.num_columns());
1500 for (i, field) in target.fields().iter().enumerate() {
1501 let src = batch.column(i);
1502 if src.data_type() == field.data_type() {
1503 cols.push(src.clone());
1504 continue;
1505 }
1506 let c = cast_with_options(src.as_ref(), field.data_type(), &opts)
1507 .map_err(|e| RpcError::type_error(format!("cast field {}: {}", field.name(), e)))?;
1508 cols.push(c);
1509 }
1510 RecordBatch::try_new(target.clone(), cols).map_err(RpcError::from)
1513}
1514
1515pub(crate) struct EnvelopeMeta<'a> {
1526 server_id: &'a str,
1527 request_id: &'a str,
1528 md: Option<Metadata>,
1531}
1532
1533impl<'a> EnvelopeMeta<'a> {
1534 pub(crate) fn new(server_id: &'a str, request_id: &'a str) -> Self {
1535 Self {
1536 server_id,
1537 request_id,
1538 md: None,
1539 }
1540 }
1541
1542 fn map(&mut self) -> &mut Metadata {
1545 if self.md.is_none() {
1546 let mut md = Metadata::with_capacity(5);
1547 if !self.server_id.is_empty() {
1548 md.insert(SERVER_ID_KEY.to_string(), self.server_id.to_string());
1549 }
1550 if !self.request_id.is_empty() {
1551 md.insert(REQUEST_ID_KEY.to_string(), self.request_id.to_string());
1552 }
1553 self.md = Some(md);
1554 }
1555 self.md.as_mut().unwrap()
1556 }
1557
1558 fn set(&mut self, key: &'static str, val: String) {
1561 let md = self.map();
1562 if let Some(slot) = md.get_mut(key) {
1563 *slot = val;
1564 } else {
1565 md.insert(key.to_string(), val);
1566 }
1567 }
1568
1569 pub(crate) fn log(&mut self, msg: &LogMessage) -> &Metadata {
1571 self.set(LOG_LEVEL_KEY, msg.level.as_str().to_string());
1572 self.set(LOG_MESSAGE_KEY, msg.message.clone());
1573 if !msg.extras.is_empty() {
1574 self.set(LOG_EXTRA_KEY, msg.extras_json());
1575 } else {
1576 self.map().remove(LOG_EXTRA_KEY);
1578 }
1579 self.md.as_ref().unwrap()
1580 }
1581
1582 pub(crate) fn error(&mut self, err: &RpcError) -> &Metadata {
1584 let extra = serde_json::json!({
1585 "exception_type": err.error_type,
1586 "exception_message": err.message,
1587 "traceback": err.traceback,
1588 })
1589 .to_string();
1590 self.set(LOG_LEVEL_KEY, "EXCEPTION".to_string());
1591 self.set(LOG_MESSAGE_KEY, err.message.clone());
1592 self.set(LOG_EXTRA_KEY, extra);
1593 self.md.as_ref().unwrap()
1594 }
1595}
1596
1597#[cfg(feature = "http")]
1601pub(crate) fn build_log_metadata(msg: &LogMessage, server_id: &str, request_id: &str) -> Metadata {
1602 let mut e = EnvelopeMeta::new(server_id, request_id);
1603 e.log(msg);
1604 e.md.unwrap()
1605}
1606
1607pub(crate) fn build_error_metadata(err: &RpcError, server_id: &str, request_id: &str) -> Metadata {
1608 let mut e = EnvelopeMeta::new(server_id, request_id);
1609 e.error(err);
1610 e.md.unwrap()
1611}
1612
1613pub(crate) fn write_error_stream<W: Write>(
1615 w: &mut W,
1616 schema: &Schema,
1617 err: &RpcError,
1618 server_id: &str,
1619 request_id: &str,
1620) -> Result<()> {
1621 let mut sw = StreamWriter::new(w, schema)?;
1622 let md = build_error_metadata(err, server_id, request_id);
1623 sw.write(&empty_batch(schema)?, Some(&md))?;
1624 sw.finish()?;
1625 Ok(())
1626}
1627
1628#[cfg(test)]
1629mod tests {
1630 use super::*;
1631 use std::io::Cursor;
1632 use std::sync::atomic::{AtomicBool, Ordering};
1633
1634 fn request_bytes(method: &str) -> Vec<u8> {
1637 let schema = empty_schema();
1638 let batch = empty_batch(&schema).unwrap();
1639 let mut buf = Vec::new();
1640 {
1641 let mut w = StreamWriter::new(&mut buf, &schema).unwrap();
1642 let mut md = Metadata::new();
1643 md.insert(RPC_METHOD_KEY.into(), method.into());
1644 md.insert(REQUEST_VERSION_KEY.into(), REQUEST_VERSION.into());
1645 md.insert(REQUEST_ID_KEY.into(), format!("req-{method}"));
1646 w.write(&batch, Some(&md)).unwrap();
1647 w.finish().unwrap();
1648 }
1649 buf
1650 }
1651
1652 #[test]
1653 fn panicking_handler_yields_error_envelope_and_loop_survives() {
1654 let mut server = RpcServer::new("test-srv");
1655 server.register(MethodInfo::unary(
1656 "boom",
1657 empty_schema(),
1658 empty_schema(),
1659 |_req, _ctx| panic!("handler exploded"),
1660 ));
1661 let ran_second = Arc::new(AtomicBool::new(false));
1662 let flag = ran_second.clone();
1663 server.register(MethodInfo::unary(
1664 "ok",
1665 empty_schema(),
1666 empty_schema(),
1667 move |_req, _ctx| {
1668 flag.store(true, Ordering::SeqCst);
1669 Ok(None)
1670 },
1671 ));
1672
1673 let mut input = request_bytes("boom");
1676 input.extend(request_bytes("ok"));
1677 let mut output: Vec<u8> = Vec::new();
1678 server.serve(Cursor::new(input), &mut output);
1679
1680 assert!(
1681 ran_second.load(Ordering::SeqCst),
1682 "serve loop aborted after a handler panic"
1683 );
1684
1685 let mut r = StreamReader::new(output.as_slice()).unwrap();
1688 let (_b, md) = r.read_next().unwrap().expect("error batch");
1689 assert_eq!(md_get(&md, LOG_LEVEL_KEY), Some("EXCEPTION"));
1690 }
1691
1692 #[test]
1693 fn transport_options_reports_shm_capability_unregistered() {
1694 use crate::metadata::TRANSPORT_SHM_KEY;
1695 use crate::transport_options::{shm_available, TRANSPORT_OPTIONS_METHOD_NAME};
1696
1697 let mut server = RpcServer::new("test-srv");
1698 server.register(MethodInfo::unary(
1699 "noop",
1700 empty_schema(),
1701 empty_schema(),
1702 |_req, _ctx| Ok(None),
1703 ));
1704 assert!(!server.methods.contains_key(TRANSPORT_OPTIONS_METHOD_NAME));
1706
1707 let input = request_bytes(TRANSPORT_OPTIONS_METHOD_NAME);
1708 let mut output: Vec<u8> = Vec::new();
1709 server.serve(Cursor::new(input), &mut output);
1710
1711 let mut r = StreamReader::new(output.as_slice()).unwrap();
1712 let (_b, md) = r.read_next().unwrap().expect("transport options batch");
1713 let expected = if shm_available() { "true" } else { "false" };
1714 assert_eq!(md_get(&md, TRANSPORT_SHM_KEY), Some(expected));
1715 assert_eq!(md_get(&md, REQUEST_VERSION_KEY), Some(REQUEST_VERSION));
1716 assert_eq!(md_get(&md, SERVER_ID_KEY), Some("test-srv"));
1717 }
1718
1719 #[cfg(feature = "shm")]
1729 mod shm_requests {
1730 use super::*;
1731 use crate::metadata::{SHM_OFFSET_KEY, SHM_SEGMENT_NAME_KEY, SHM_SEGMENT_SIZE_KEY};
1732 use crate::shm::{
1733 is_shm_pointer_batch, make_shm_pointer_batch, maybe_write_to_shm, ShmSegment,
1734 };
1735 use arrow_array::{BinaryArray, Int64Array};
1736 use arrow_schema::{DataType, Field};
1737
1738 fn params_schema() -> SchemaRef {
1739 Arc::new(Schema::new(vec![Field::new(
1740 "request",
1741 DataType::Binary,
1742 false,
1743 )]))
1744 }
1745
1746 fn result_schema() -> SchemaRef {
1747 Arc::new(Schema::new(vec![Field::new("n", DataType::Int64, false)]))
1748 }
1749
1750 fn request_batch(payload: &[u8]) -> RecordBatch {
1751 RecordBatch::try_new(
1752 params_schema(),
1753 vec![Arc::new(BinaryArray::from(vec![Some(payload)]))],
1754 )
1755 .unwrap()
1756 }
1757
1758 fn dispatch_md(seg: Option<&ShmSegment>) -> Metadata {
1761 let mut md = Metadata::new();
1762 md.insert(RPC_METHOD_KEY.into(), "do_thing".into());
1763 md.insert(REQUEST_VERSION_KEY.into(), REQUEST_VERSION.into());
1764 if let Some(seg) = seg {
1765 md.insert(SHM_SEGMENT_NAME_KEY.into(), seg.name().to_string());
1766 md.insert(SHM_SEGMENT_SIZE_KEY.into(), seg.size().to_string());
1767 }
1768 md
1769 }
1770
1771 fn pointer_request(seg: &ShmSegment, payload: &[u8], advertise: bool) -> Vec<u8> {
1776 let md = dispatch_md(advertise.then_some(seg));
1777 let (ptr, ptr_md) = maybe_write_to_shm(request_batch(payload), md, Some(seg)).unwrap();
1778 assert!(
1779 is_shm_pointer_batch(&ptr, &ptr_md),
1780 "request batch should have routed through shm"
1781 );
1782 let mut buf = Vec::new();
1783 {
1784 let mut w = StreamWriter::new(&mut buf, ptr.schema().as_ref()).unwrap();
1785 w.write(&ptr, Some(&ptr_md)).unwrap();
1786 w.finish().unwrap();
1787 }
1788 buf
1789 }
1790
1791 fn inline_request(payload: &[u8], seg: Option<&ShmSegment>) -> Vec<u8> {
1793 let batch = request_batch(payload);
1794 let md = dispatch_md(seg);
1795 let mut buf = Vec::new();
1796 {
1797 let mut w = StreamWriter::new(&mut buf, batch.schema().as_ref()).unwrap();
1798 w.write(&batch, Some(&md)).unwrap();
1799 w.finish().unwrap();
1800 }
1801 buf
1802 }
1803
1804 fn payload_server(seen: Arc<Mutex<Vec<Vec<u8>>>>) -> RpcServer {
1808 let mut server = RpcServer::new("shm-srv");
1809 let rs = result_schema();
1810 server.register(MethodInfo::unary(
1811 "do_thing",
1812 params_schema(),
1813 result_schema(),
1814 move |req, _ctx| {
1815 let col = req
1816 .column("request")
1817 .expect("request column")
1818 .as_any()
1819 .downcast_ref::<BinaryArray>()
1820 .unwrap();
1821 lock_ok(&seen).push(col.value(0).to_vec());
1822 Ok(Some(RecordBatch::try_new(
1823 rs.clone(),
1824 vec![Arc::new(Int64Array::from(vec![1i64]))],
1825 )?))
1826 },
1827 ));
1828 server
1829 }
1830
1831 fn response_metadata(output: &[u8]) -> Vec<Metadata> {
1834 let mut out = Vec::new();
1835 let mut cursor = Cursor::new(output);
1836 while (cursor.position() as usize) < output.len() {
1837 let mut reader = StreamReader::new(&mut cursor).unwrap();
1838 while let Some((_b, md)) = reader.read_next().unwrap() {
1839 out.push(md);
1840 }
1841 }
1842 out
1843 }
1844
1845 #[test]
1849 fn pointer_request_batch_resolves_via_segment_named_in_metadata() {
1850 let seg = ShmSegment::create(1024 * 1024).unwrap();
1851 let payload = b"serialized-request-blob";
1852
1853 let seen = Arc::new(Mutex::new(Vec::new()));
1854 let server = payload_server(seen.clone());
1855 let mut output: Vec<u8> = Vec::new();
1856 server.serve(
1857 Cursor::new(pointer_request(&seg, payload, true)),
1858 &mut output,
1859 );
1860 assert_eq!(lock_ok(&seen).as_slice(), &[payload.to_vec()]);
1861
1862 let seen2 = Arc::new(Mutex::new(Vec::new()));
1864 let server2 = payload_server(seen2.clone());
1865 let mut input = Cursor::new(pointer_request(&seg, payload, true));
1866 let mut output2: Vec<u8> = Vec::new();
1867 assert!(server2.serve_one(&mut input, &mut output2).unwrap());
1868 assert_eq!(lock_ok(&seen2).as_slice(), &[payload.to_vec()]);
1869 }
1870
1871 #[test]
1874 fn serve_caches_client_segment_for_offset_only_requests() {
1875 let seg = ShmSegment::create(1024 * 1024).unwrap();
1876
1877 let mut input = inline_request(b"first", Some(&seg));
1879
1880 let (off, len) = seg
1882 .allocate_and_write(&request_batch(b"second"))
1883 .unwrap()
1884 .expect("payload fits");
1885 let (ptr, mut ptr_md) =
1886 make_shm_pointer_batch(params_schema().as_ref(), off, len).unwrap();
1887 ptr_md.insert(RPC_METHOD_KEY.into(), "do_thing".into());
1888 ptr_md.insert(REQUEST_VERSION_KEY.into(), REQUEST_VERSION.into());
1889 {
1890 let mut w = StreamWriter::new(&mut input, ptr.schema().as_ref()).unwrap();
1891 w.write(&ptr, Some(&ptr_md)).unwrap();
1892 w.finish().unwrap();
1893 }
1894
1895 let seen = Arc::new(Mutex::new(Vec::new()));
1896 let server = payload_server(seen.clone());
1897 let mut output: Vec<u8> = Vec::new();
1898 server.serve(Cursor::new(input), &mut output);
1899 assert_eq!(
1900 lock_ok(&seen).as_slice(),
1901 &[b"first".to_vec(), b"second".to_vec()]
1902 );
1903 }
1904
1905 #[test]
1908 fn pointer_request_without_segment_trips_single_row_guard() {
1909 let seg = ShmSegment::create(1024 * 1024).unwrap();
1910 let seen = Arc::new(Mutex::new(Vec::new()));
1911 let server = payload_server(seen.clone());
1912 let mut output: Vec<u8> = Vec::new();
1913 server.serve(
1914 Cursor::new(pointer_request(&seg, b"orphan", false)),
1915 &mut output,
1916 );
1917 assert!(lock_ok(&seen).is_empty(), "guard should reject dispatch");
1918 }
1919
1920 #[test]
1926 fn response_routed_through_shm_only_when_request_signalled_shm() {
1927 let seg = ShmSegment::create(1024 * 1024).unwrap();
1928
1929 let mut input = inline_request(b"a", Some(&seg));
1932 input.extend(inline_request(b"b", None));
1933
1934 let seen = Arc::new(Mutex::new(Vec::new()));
1935 let server = payload_server(seen.clone());
1936 let mut output: Vec<u8> = Vec::new();
1937 server.serve(Cursor::new(input), &mut output);
1938 assert_eq!(lock_ok(&seen).len(), 2);
1939
1940 let mds = response_metadata(&output);
1941 assert_eq!(mds.len(), 2, "one data batch per response");
1942 assert!(
1943 mds[0].contains_key(SHM_OFFSET_KEY),
1944 "response A (segment advertised) should route through shm"
1945 );
1946 assert!(
1947 !mds[1].contains_key(SHM_OFFSET_KEY),
1948 "response B (no shm signal) must stay inline despite the cached segment"
1949 );
1950 }
1951 }
1952}