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, LOCATION_KEY, LOG_EXTRA_KEY, LOG_LEVEL_KEY, LOG_MESSAGE_KEY, PROTOCOL_VERSION_KEY,
17 REQUEST_ID_KEY, REQUEST_VERSION, REQUEST_VERSION_KEY, RPC_METHOD_KEY, SERVER_ID_KEY,
18 SHM_OFFSET_KEY, SHM_SEGMENT_NAME_KEY,
19};
20#[cfg(feature = "shm")]
21use crate::shm::{is_shm_pointer_batch, maybe_write_to_shm, resolve_shm_batch, ShmSegment};
22
23#[cfg(not(feature = "shm"))]
25pub(crate) struct ShmSegment;
26
27#[cfg(feature = "shm")]
30fn maybe_attach_shm(req_md: &Metadata) -> Option<ShmSegment> {
31 let name = req_md.get(SHM_SEGMENT_NAME_KEY)?;
32 let size: usize = req_md.get(SHM_SEGMENT_SIZE_KEY)?.parse().ok()?;
33 match ShmSegment::attach(name, size, false) {
34 Ok(seg) => Some(seg),
35 Err(e) => {
36 tracing::warn!(target: "vgi_rpc.shm", "ignoring malformed SHM metadata ({e})");
37 None
38 }
39 }
40}
41
42#[cfg(not(feature = "shm"))]
43#[inline]
44fn maybe_attach_shm(_req_md: &Metadata) -> Option<ShmSegment> {
45 None
46}
47
48#[derive(Default)]
62pub(crate) struct ConnectionShm {
63 #[cfg(feature = "shm")]
64 name: Option<String>,
65 #[cfg(feature = "shm")]
66 segment: Option<ShmSegment>,
67}
68
69#[cfg(feature = "shm")]
70impl ConnectionShm {
71 fn refresh(&mut self, req_md: &Metadata) {
74 let Some(name) = req_md.get(SHM_SEGMENT_NAME_KEY) else {
75 return;
76 };
77 if self.name.as_deref() == Some(name.as_str()) {
78 return;
79 }
80 let Some(new) = maybe_attach_shm(req_md) else {
81 return;
82 };
83 self.segment = Some(new);
84 self.name = Some(name.clone());
85 }
86
87 fn segment(&self) -> Option<&ShmSegment> {
88 self.segment.as_ref()
89 }
90}
91
92#[cfg(not(feature = "shm"))]
93impl ConnectionShm {
94 #[inline]
99 fn segment(&self) -> Option<&ShmSegment> {
100 None
101 }
102}
103use crate::stream::{empty_schema, Emitted, OutputCollector, StreamResult, StreamStateKind};
104use crate::wire::{
105 empty_batch, md_get, Metadata, StreamReader, StreamWriter, INVALID_UTF8_METADATA_KEY,
106};
107
108pub(crate) fn serialize_request_batch(batch: &RecordBatch) -> std::io::Result<Vec<u8>> {
112 let mut buf = Vec::new();
113 {
114 let mut w = arrow_ipc::writer::StreamWriter::try_new(&mut buf, batch.schema_ref())
115 .map_err(|e| std::io::Error::other(e.to_string()))?;
116 w.write(batch)
117 .map_err(|e| std::io::Error::other(e.to_string()))?;
118 w.finish()
119 .map_err(|e| std::io::Error::other(e.to_string()))?;
120 }
121 Ok(buf)
122}
123
124fn lock_ok<T>(m: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
130 m.lock().unwrap_or_else(|e| e.into_inner())
131}
132
133pub(crate) fn call_guard<T>(f: impl FnOnce() -> T) -> Result<T> {
138 std::panic::catch_unwind(std::panic::AssertUnwindSafe(f))
139 .map_err(|_| RpcError::new("RuntimeError", "handler panicked"))
140}
141
142pub(crate) fn validate_protocol_version(
149 server_version: &str,
150 request_metadata: &Metadata,
151) -> Result<()> {
152 if server_version.is_empty() {
153 return Ok(());
154 }
155 let Some(client_version) = md_get(request_metadata, PROTOCOL_VERSION_KEY) else {
156 return Ok(());
157 };
158 let client_major = client_version.split('.').next().unwrap_or("");
159 let server_major = server_version.split('.').next().unwrap_or("");
160 if client_major != server_major {
161 return Err(RpcError::version_error(format!(
162 "protocol_version mismatch: client {:?} is incompatible with server {:?}",
163 client_version, server_version
164 )));
165 }
166 Ok(())
167}
168
169#[derive(Clone)]
171pub struct CallContext {
172 pub server_id: String,
173 pub method: String,
174 pub request_id: String,
175 pub transport_metadata: Arc<Metadata>,
176 pub auth: crate::auth::AuthContext,
179 pub cookies: std::collections::BTreeMap<String, String>,
181 pub kind: Option<crate::transport::TransportKind>,
185 pub(crate) log_sink: Arc<Mutex<Vec<LogMessage>>>,
186 pub(crate) tick_metadata: Arc<Mutex<Metadata>>,
189 pub(crate) sticky: Option<Arc<dyn StickySink>>,
194}
195
196pub trait StickySink: Send + Sync {
201 fn accept_opens(&self) -> bool;
203 fn current_state(&self) -> Option<Arc<dyn std::any::Any + Send + Sync>>;
205 fn current_session_id(&self) -> Option<String>;
207 fn open(
209 &self,
210 state: Arc<dyn std::any::Any + Send + Sync>,
211 ttl: Option<std::time::Duration>,
212 ) -> Result<()>;
213 fn close(&self) -> Result<bool>;
215}
216
217impl CallContext {
218 pub fn client_log(&self, level: LogLevel, message: impl Into<String>) {
219 lock_ok(&self.log_sink).push(LogMessage::new(level, message));
220 }
221
222 pub fn client_log_with(&self, msg: LogMessage) {
223 lock_ok(&self.log_sink).push(msg);
224 }
225
226 pub(crate) fn drain_logs(&self) -> Vec<LogMessage> {
227 std::mem::take(&mut *lock_ok(&self.log_sink))
228 }
229
230 pub fn tick_metadata(&self, key: &str) -> Option<String> {
233 lock_ok(&self.tick_metadata).get(key).cloned()
234 }
235
236 #[cfg(feature = "http")]
241 pub(crate) fn set_tick_metadata(&self, md: Metadata) {
242 *lock_ok(&self.tick_metadata) = md;
243 }
244
245 pub(crate) fn for_request(server: &RpcServer, req: &Request) -> Self {
250 Self {
251 server_id: server.server_id.clone(),
252 method: req.method.clone(),
253 request_id: req.request_id.clone(),
254 transport_metadata: req.metadata.clone(),
255 auth: crate::auth::AuthContext::anonymous(),
256 cookies: std::collections::BTreeMap::new(),
257 kind: server.transport_kind(),
258 log_sink: Arc::new(Mutex::new(Vec::new())),
259 tick_metadata: Arc::new(Mutex::new(Metadata::default())),
260 sticky: None,
261 }
262 }
263
264 #[cfg(feature = "http")]
269 pub(crate) fn with_auth_cookies(
270 server: &RpcServer,
271 req: &Request,
272 auth: crate::auth::AuthContext,
273 cookies: std::collections::BTreeMap<String, String>,
274 ) -> Self {
275 Self {
276 server_id: server.server_id.clone(),
277 method: req.method.clone(),
278 request_id: req.request_id.clone(),
279 transport_metadata: req.metadata.clone(),
280 auth,
281 cookies,
282 kind: server.transport_kind(),
283 log_sink: Arc::new(Mutex::new(Vec::new())),
284 tick_metadata: Arc::new(Mutex::new(Metadata::default())),
285 sticky: None,
286 }
287 }
288
289 #[cfg(feature = "http")]
294 pub(crate) fn set_sticky(&mut self, sink: Arc<dyn StickySink>) {
295 self.sticky = Some(sink);
296 }
297
298 pub fn session<T: std::any::Any + Send + Sync>(&self) -> Option<Arc<T>> {
306 let state = self.sticky.as_ref()?.current_state()?;
307 state.downcast::<T>().ok()
308 }
309
310 pub fn session_id(&self) -> Option<String> {
313 self.sticky.as_ref()?.current_session_id()
314 }
315
316 pub fn open_session(
327 &self,
328 state: Arc<dyn std::any::Any + Send + Sync>,
329 ttl: Option<std::time::Duration>,
330 ) -> Result<()> {
331 let sink = self.sticky.as_ref().ok_or_else(|| {
332 RpcError::runtime_error("sticky sessions not available on this transport")
333 })?;
334 if !sink.accept_opens() {
335 return Err(RpcError::runtime_error(
336 "client did not opt in to sticky sessions \
337 (missing VGI-Session-Accept: true header — open the call inside \
338 an HttpConnection.with_session_token() block)",
339 ));
340 }
341 if sink.current_state().is_some() {
342 return Err(RpcError::runtime_error(
343 "a sticky session is already active for this request",
344 ));
345 }
346 sink.open(state, ttl)
347 }
348
349 pub fn close_session(&self) -> Result<()> {
352 let sink = self.sticky.as_ref().ok_or_else(|| {
353 RpcError::runtime_error("sticky sessions not available on this transport")
354 })?;
355 sink.close()?;
356 Ok(())
357 }
358}
359
360pub struct Request {
362 pub method: String,
363 pub request_id: String,
364 pub batch: RecordBatch,
365 pub metadata: Arc<Metadata>,
369}
370
371impl Request {
372 pub fn column(&self, name: &str) -> Option<&dyn arrow_array::Array> {
373 let idx = self.batch.schema().index_of(name).ok()?;
374 Some(self.batch.column(idx).as_ref())
375 }
376
377 pub(crate) fn from_read_batch(
385 batch: RecordBatch,
386 metadata: Metadata,
387 require_method: bool,
388 ) -> Result<Self> {
389 if metadata.contains_key(INVALID_UTF8_METADATA_KEY) {
390 return Err(RpcError::protocol_error(
391 "Invalid UTF-8 in request batch custom_metadata",
392 ));
393 }
394 let method = if require_method {
395 md_get(&metadata, RPC_METHOD_KEY)
396 .ok_or_else(|| {
397 RpcError::protocol_error(
398 "Missing 'vgi_rpc.method' in request batch custom_metadata.",
399 )
400 })?
401 .to_string()
402 } else {
403 md_get(&metadata, RPC_METHOD_KEY).unwrap_or("").to_string()
404 };
405 let version = md_get(&metadata, REQUEST_VERSION_KEY).ok_or_else(|| {
406 RpcError::version_error(format!(
407 "Missing 'vgi_rpc.request_version' in request batch custom_metadata. Set it to {:?}.",
408 REQUEST_VERSION
409 ))
410 })?;
411 if version != REQUEST_VERSION {
412 return Err(RpcError::version_error(format!(
413 "Unsupported request version {:?}, expected {:?}.",
414 version, REQUEST_VERSION
415 )));
416 }
417 let external_pointer = md_get(&metadata, LOCATION_KEY).is_some();
418 if require_method
419 && !batch.schema().fields().is_empty()
420 && batch.num_rows() != 1
421 && !external_pointer
422 {
423 return Err(RpcError::protocol_error(format!(
424 "Expected 1 row in request batch, got {}",
425 batch.num_rows()
426 )));
427 }
428 let request_id = md_get(&metadata, REQUEST_ID_KEY).unwrap_or("").to_string();
429 Ok(Request {
430 method,
431 request_id,
432 batch,
433 metadata: Arc::new(metadata),
434 })
435 }
436}
437
438pub(crate) fn validate_parameter_batch(batch: &RecordBatch, expected: &Schema) -> Result<()> {
444 if batch.schema().fields() != expected.fields() {
445 return Err(RpcError::type_error(format!(
446 "parameter schema mismatch: expected {expected:?}, got {:?}",
447 batch.schema()
448 )));
449 }
450 if !expected.fields().is_empty() && batch.num_rows() != 1 {
451 return Err(RpcError::protocol_error(format!(
452 "Expected 1 row in request batch, got {}",
453 batch.num_rows()
454 )));
455 }
456 Ok(())
457}
458
459#[derive(Clone, Copy, Debug, PartialEq, Eq)]
461pub enum MethodType {
462 Unary,
463 Producer,
464 Exchange,
465 Dynamic,
467}
468
469pub type UnaryHandler =
471 Arc<dyn Fn(&Request, &CallContext) -> Result<Option<RecordBatch>> + Send + Sync>;
472
473pub type StreamHandler = Arc<dyn Fn(&Request, &CallContext) -> Result<StreamResult> + Send + Sync>;
475
476#[derive(Default)]
478pub struct RpcServerBuilder {
479 server_id: Option<String>,
480 server_version: Option<String>,
481 protocol_name: Option<String>,
482 protocol_version: Option<String>,
483 enable_describe: bool,
484 dispatch_hook: Option<Arc<dyn crate::hooks::DispatchHook>>,
485 on_serve_start: Option<crate::transport::ServeStartHook>,
486 #[cfg(feature = "http")]
487 external_config: Option<Arc<crate::external::ExternalLocationConfig>>,
488}
489
490impl RpcServerBuilder {
491 pub fn server_id(mut self, id: impl Into<String>) -> Self {
492 self.server_id = Some(id.into());
493 self
494 }
495
496 pub fn server_version(mut self, v: impl Into<String>) -> Self {
497 self.server_version = Some(v.into());
498 self
499 }
500
501 pub fn protocol_name(mut self, name: impl Into<String>) -> Self {
502 self.protocol_name = Some(name.into());
503 self
504 }
505
506 pub fn protocol_version(mut self, v: impl Into<String>) -> Self {
510 self.protocol_version = Some(v.into());
511 self
512 }
513
514 pub fn enable_describe(mut self, enabled: bool) -> Self {
515 self.enable_describe = enabled;
516 self
517 }
518
519 pub fn with_hook(mut self, hook: Arc<dyn crate::hooks::DispatchHook>) -> Self {
520 self.dispatch_hook = Some(hook);
521 self
522 }
523
524 pub fn on_serve_start(mut self, hook: crate::transport::ServeStartHook) -> Self {
534 self.on_serve_start = Some(hook);
535 self
536 }
537
538 #[cfg(feature = "http")]
542 pub fn with_external_location(mut self, cfg: crate::external::ExternalLocationConfig) -> Self {
543 self.external_config = Some(Arc::new(cfg));
544 self
545 }
546
547 pub fn build(self) -> RpcServer {
548 RpcServer {
549 methods: HashMap::new(),
550 server_id: self.server_id.unwrap_or_else(crate::util::short_random_id),
551 server_version: self.server_version.unwrap_or_default(),
552 protocol_name: self.protocol_name.unwrap_or_default(),
553 protocol_version: self.protocol_version.unwrap_or_default(),
554 protocol_hash: std::sync::OnceLock::new(),
555 describe_enabled: self.enable_describe,
556 dispatch_hook: self.dispatch_hook,
557 on_serve_start: self.on_serve_start,
558 transport_notify: Mutex::new(()),
559 transport_state: Mutex::new(None),
560 #[cfg(feature = "http")]
561 external_config: self.external_config,
562 }
563 }
564}
565
566pub struct MethodInfo {
573 pub name: String,
574 pub method_type: MethodType,
575 pub params_schema: SchemaRef,
577 pub result_schema: SchemaRef,
579 pub header_schema: Option<SchemaRef>,
581 pub doc: Option<String>,
583 pub param_types: Vec<(String, String)>,
586 pub param_defaults: Vec<(String, serde_json::Value)>,
588 pub param_docs: Vec<(String, String)>,
590 pub has_return: bool,
592 pub unary: Option<UnaryHandler>,
593 pub stream: Option<StreamHandler>,
594 pub state_decoder: Option<StateDecoder>,
600}
601
602pub type StateDecoder = Arc<dyn Fn(&[u8]) -> Result<crate::stream::StreamStateKind> + Send + Sync>;
605
606impl MethodInfo {
607 pub fn unary(
609 name: impl Into<String>,
610 params_schema: SchemaRef,
611 result_schema: SchemaRef,
612 handler: impl Fn(&Request, &CallContext) -> Result<Option<RecordBatch>> + Send + Sync + 'static,
613 ) -> Self {
614 let has_return = !result_schema.fields().is_empty();
615 Self {
616 name: name.into(),
617 method_type: MethodType::Unary,
618 params_schema,
619 result_schema,
620 header_schema: None,
621 doc: None,
622 param_types: Vec::new(),
623 param_defaults: Vec::new(),
624 param_docs: Vec::new(),
625 has_return,
626 unary: Some(Arc::new(handler)),
627 stream: None,
628 state_decoder: None,
629 }
630 }
631
632 pub fn stream(
642 name: impl Into<String>,
643 method_type: MethodType,
644 params_schema: SchemaRef,
645 handler: impl Fn(&Request, &CallContext) -> Result<StreamResult> + Send + Sync + 'static,
646 ) -> Self {
647 assert!(
648 matches!(
649 method_type,
650 MethodType::Producer | MethodType::Exchange | MethodType::Dynamic
651 ),
652 "stream methods must be Producer / Exchange / Dynamic"
653 );
654 Self {
655 name: name.into(),
656 method_type,
657 params_schema,
658 result_schema: empty_schema(),
659 header_schema: None,
660 doc: None,
661 param_types: Vec::new(),
662 param_defaults: Vec::new(),
663 param_docs: Vec::new(),
664 has_return: false,
665 unary: None,
666 stream: Some(Arc::new(handler)),
667 state_decoder: None,
668 }
669 }
670
671 pub fn with_state_decoder(mut self, decoder: StateDecoder) -> Self {
673 self.state_decoder = Some(decoder);
674 self
675 }
676
677 pub fn doc(mut self, s: impl Into<String>) -> Self {
678 self.doc = Some(s.into());
679 self
680 }
681
682 pub fn param_type(mut self, param: impl Into<String>, ty: impl Into<String>) -> Self {
683 self.param_types.push((param.into(), ty.into()));
684 self
685 }
686
687 pub fn param_default(mut self, param: impl Into<String>, value: serde_json::Value) -> Self {
688 self.param_defaults.push((param.into(), value));
689 self
690 }
691
692 pub fn param_doc(mut self, param: impl Into<String>, doc: impl Into<String>) -> Self {
693 self.param_docs.push((param.into(), doc.into()));
694 self
695 }
696
697 pub fn header_schema(mut self, schema: SchemaRef) -> Self {
698 self.header_schema = Some(schema);
699 self
700 }
701}
702
703pub struct RpcServer {
705 methods: HashMap<String, MethodInfo>,
706 pub server_id: String,
707 pub(crate) server_version: String,
708 pub(crate) protocol_name: String,
709 pub(crate) protocol_version: String,
710 pub(crate) protocol_hash: std::sync::OnceLock<String>,
711 pub(crate) describe_enabled: bool,
712 pub(crate) dispatch_hook: Option<Arc<dyn crate::hooks::DispatchHook>>,
713 on_serve_start: Option<crate::transport::ServeStartHook>,
716 transport_notify: Mutex<()>,
720 transport_state: Mutex<
723 Option<(
724 crate::transport::TransportKind,
725 crate::transport::TransportCapabilities,
726 )>,
727 >,
728 #[cfg(feature = "http")]
729 pub(crate) external_config: Option<Arc<crate::external::ExternalLocationConfig>>,
730}
731
732impl RpcServer {
733 pub fn new(server_id: impl Into<String>) -> Self {
735 Self::builder().server_id(server_id).build()
736 }
737
738 pub fn builder() -> RpcServerBuilder {
740 RpcServerBuilder::default()
741 }
742
743 pub fn protocol_name(&self) -> &str {
744 &self.protocol_name
745 }
746
747 pub fn describe_enabled(&self) -> bool {
748 self.describe_enabled
749 }
750
751 pub fn server_version(&self) -> &str {
752 &self.server_version
753 }
754
755 pub fn protocol_version(&self) -> &str {
756 &self.protocol_version
757 }
758
759 pub fn protocol_hash(&self) -> &str {
762 self.protocol_hash.get_or_init(|| {
763 match crate::introspect::build_describe(
764 &self.protocol_name,
765 &self.methods,
766 &self.server_id,
767 &self.protocol_version,
768 ) {
769 Ok((_, md)) => md
770 .get(crate::metadata::PROTOCOL_HASH_KEY)
771 .cloned()
772 .unwrap_or_default(),
773 Err(_) => String::new(),
774 }
775 })
776 }
777
778 #[cfg(feature = "http")]
779 pub fn external_config(&self) -> Option<&Arc<crate::external::ExternalLocationConfig>> {
780 self.external_config.as_ref()
781 }
782
783 pub fn transport_kind(&self) -> Option<crate::transport::TransportKind> {
787 lock_ok(&self.transport_state).as_ref().map(|(k, _)| *k)
788 }
789
790 pub fn transport_capabilities(&self) -> crate::transport::TransportCapabilities {
794 lock_ok(&self.transport_state)
795 .as_ref()
796 .map(|(_, c)| *c)
797 .unwrap_or_default()
798 }
799
800 pub fn notify_transport(
812 &self,
813 kind: crate::transport::TransportKind,
814 caps: crate::transport::TransportCapabilities,
815 ) {
816 let _notify = lock_ok(&self.transport_notify);
817 {
818 let guard = lock_ok(&self.transport_state);
819 if let Some((cur_kind, cur_caps)) = guard.as_ref() {
820 if *cur_kind == kind && *cur_caps == caps {
821 return;
822 }
823 }
824 }
825 if let Some(h) = self.on_serve_start.clone() {
826 h(kind, &caps);
827 }
828 *lock_ok(&self.transport_state) = Some((kind, caps));
831 }
832
833 pub fn register(&mut self, info: MethodInfo) {
835 self.methods.insert(info.name.clone(), info);
836 }
837
838 pub fn register_unary(
842 &mut self,
843 name: impl Into<String>,
844 result_schema: SchemaRef,
845 handler: impl Fn(&Request, &CallContext) -> Result<Option<RecordBatch>> + Send + Sync + 'static,
846 ) {
847 self.register(MethodInfo::unary(
848 name,
849 empty_schema(),
850 result_schema,
851 handler,
852 ));
853 }
854
855 pub fn register_stream(
859 &mut self,
860 name: impl Into<String>,
861 method_type: MethodType,
862 handler: impl Fn(&Request, &CallContext) -> Result<StreamResult> + Send + Sync + 'static,
863 ) {
864 self.register(MethodInfo::stream(
865 name,
866 method_type,
867 empty_schema(),
868 handler,
869 ));
870 }
871
872 pub fn method(&self, name: &str) -> Option<&MethodInfo> {
873 self.methods.get(name)
874 }
875
876 pub fn methods(&self) -> &HashMap<String, MethodInfo> {
877 &self.methods
878 }
879
880 pub fn method_names(&self) -> Vec<&str> {
881 self.sorted_method_names()
882 }
883
884 pub fn sorted_method_names(&self) -> Vec<&str> {
887 let mut names: Vec<_> = self.methods.keys().map(String::as_str).collect();
888 names.sort();
889 names
890 }
891
892 pub fn serve<R: Read, W: Write>(&self, mut r: R, mut w: W) {
903 let mut conn_shm = ConnectionShm::default();
907 loop {
908 match self.serve_one_conn(&mut r, &mut w, Some(&mut conn_shm)) {
909 Ok(keep_going) => {
910 if !keep_going {
911 return;
912 }
913 }
914 Err(e) => {
915 tracing::warn!(
921 target: "vgi_rpc.server",
922 error = %e,
923 "serve loop terminating connection on error"
924 );
925 return;
926 }
927 }
928 }
929 }
930
931 pub fn serve_with_shutdown<R, W, F>(&self, mut r: R, mut w: W, shutdown: F)
937 where
938 R: Read,
939 W: Write,
940 F: Fn() -> bool,
941 {
942 let mut conn_shm = ConnectionShm::default();
943 loop {
944 if shutdown() {
945 return;
946 }
947 match self.serve_one_conn(&mut r, &mut w, Some(&mut conn_shm)) {
948 Ok(true) => {}
949 _ => return,
950 }
951 }
952 }
953
954 pub fn serve_one<R: Read, W: Write>(&self, r: &mut R, w: &mut W) -> Result<bool> {
961 self.serve_one_conn(r, w, None)
962 }
963
964 fn serve_one_conn<R: Read, W: Write>(
965 &self,
966 r: &mut R,
967 w: &mut W,
968 shm_cache: Option<&mut ConnectionShm>,
969 ) -> Result<bool> {
970 let result = self._serve_one(r, w, shm_cache);
971 let _ = w.flush();
972 result
973 }
974
975 fn _serve_one<R: Read, W: Write>(
976 &self,
977 r: &mut R,
978 w: &mut W,
979 mut shm_cache: Option<&mut ConnectionShm>,
980 ) -> Result<bool> {
981 let (batch, metadata, request_used_shm) =
982 match self.read_request(r, shm_cache.as_deref_mut())? {
983 Some(frame) => frame,
984 None => return Ok(false),
985 };
986 let request_id = md_get(&metadata, REQUEST_ID_KEY).unwrap_or("").to_string();
992 let req = match Request::from_read_batch(batch, metadata, true) {
993 Ok(req) => req,
994 Err(err) => {
995 write_error_stream(w, &empty_schema(), &err, &self.server_id, &request_id)?;
996 return Ok(true);
997 }
998 };
999
1000 if req.method == crate::transport_options::TRANSPORT_OPTIONS_METHOD_NAME {
1007 let mut md = crate::transport_options::worker_transport_metadata();
1008 md.insert(REQUEST_VERSION_KEY.to_string(), REQUEST_VERSION.to_string());
1009 md.insert(SERVER_ID_KEY.to_string(), self.server_id.clone());
1010 let schema = empty_schema();
1011 let batch = empty_batch(&schema)?;
1012 let mut sw = StreamWriter::new(w, &schema)?;
1013 sw.write(&batch, Some(&md))?;
1014 sw.finish()?;
1015 return Ok(true);
1016 }
1017
1018 if let Err(err) = validate_protocol_version(&self.protocol_version, &req.metadata) {
1021 write_error_stream(w, &empty_schema(), &err, &self.server_id, &req.request_id)?;
1022 return Ok(true);
1023 }
1024
1025 let ctx = CallContext::for_request(self, &req);
1026
1027 let stats = Arc::new(Mutex::new(crate::hooks::CallStatistics::default()));
1028 {
1030 let mut s = lock_ok(&stats);
1031 s.input_batches = 1;
1032 s.input_rows = req.batch.num_rows() as u64;
1033 }
1034
1035 if self.describe_enabled && req.method == crate::introspect::DESCRIBE_METHOD_NAME {
1037 match crate::introspect::build_describe(
1038 &self.protocol_name,
1039 &self.methods,
1040 &self.server_id,
1041 &self.protocol_version,
1042 ) {
1043 Ok((batch, md)) => {
1044 crate::introspect::write_describe_response(w, &batch, &md)?;
1045 }
1046 Err(err) => {
1047 write_error_stream(w, &empty_schema(), &err, &self.server_id, &req.request_id)?;
1048 }
1049 }
1050 return Ok(true);
1051 }
1052
1053 let Some(info) = self.methods.get(&req.method) else {
1054 let names = self.sorted_method_names();
1055 let msg = format!(
1056 "Unknown method: '{}'. Available methods: {:?}",
1057 req.method, names
1058 );
1059 write_error_stream(
1060 w,
1061 &empty_schema(),
1062 &RpcError::attribute_error(msg),
1063 &self.server_id,
1064 &req.request_id,
1065 )?;
1066 return Ok(true);
1067 };
1068
1069 if let Err(err) = validate_parameter_batch(&req.batch, &info.params_schema) {
1070 write_error_stream(
1071 w,
1072 &info.result_schema,
1073 &err,
1074 &self.server_id,
1075 &req.request_id,
1076 )?;
1077 return Ok(true);
1078 }
1079
1080 let method_type = match info.method_type {
1081 MethodType::Unary => "unary",
1082 _ => "stream",
1083 };
1084 #[cfg_attr(not(feature = "http"), allow(unused_mut))]
1090 let mut dispatch_info = self.dispatch_hook.as_ref().map(|_| {
1091 let mut di =
1092 crate::hooks::DispatchInfo::from_request(self, &req, method_type, &ctx.auth);
1093 if let Ok(bytes) = serialize_request_batch(&req.batch) {
1097 di.request_data = bytes;
1098 }
1099 if method_type == "stream" {
1100 di.stream_id = crate::access_log::random_stream_id();
1101 }
1102 di
1103 });
1104 let hook_token = match (self.dispatch_hook.as_ref(), dispatch_info.as_ref()) {
1105 (Some(h), Some(di)) => Some(h.on_dispatch_start(di)),
1106 _ => None,
1107 };
1108
1109 let mut app_err: Option<RpcError> = None;
1110 let dynamic_shm: Option<ShmSegment>;
1122 let shm_ref: Option<&ShmSegment> = match shm_cache {
1123 Some(cache) => {
1124 if request_used_shm {
1125 cache.segment()
1126 } else {
1127 None
1128 }
1129 }
1130 None => {
1131 dynamic_shm = maybe_attach_shm(&req.metadata);
1132 dynamic_shm.as_ref()
1133 }
1134 };
1135 #[cfg(feature = "http")]
1140 let externalized = crate::external::ExternalizedScope::new();
1141 match info.method_type {
1142 MethodType::Unary => {
1143 self.serve_unary(w, &req, info, &ctx, &stats, &mut app_err, shm_ref)?
1144 }
1145 MethodType::Producer | MethodType::Exchange | MethodType::Dynamic => {
1146 self.serve_stream(r, w, &req, info, &ctx, &stats, &mut app_err, shm_ref)?
1147 }
1148 }
1149 #[cfg(feature = "http")]
1150 let externalized_bytes = externalized.finish();
1151 #[cfg(feature = "http")]
1152 if let Some(di) = dispatch_info.as_mut() {
1153 di.externalized_bytes = externalized_bytes;
1154 }
1155 if let (Some(hook), Some(di)) = (self.dispatch_hook.as_ref(), dispatch_info.as_ref()) {
1160 let token = hook_token.unwrap_or(0);
1161 let final_stats = lock_ok(&stats).clone();
1162 hook.on_dispatch_end(token, di, app_err.as_ref(), &final_stats);
1163 }
1164 Ok(true)
1165 }
1166
1167 fn read_request<R: Read>(
1172 &self,
1173 r: &mut R,
1174 shm_cache: Option<&mut ConnectionShm>,
1175 ) -> Result<Option<(RecordBatch, Metadata, bool)>> {
1176 let mut reader = match StreamReader::new(r) {
1177 Ok(r) => r,
1178 Err(e) => {
1179 let msg = e.message.to_lowercase();
1181 if msg.contains("empty ipc stream") || msg.contains("eof") {
1182 return Ok(None);
1183 }
1184 return Err(e);
1185 }
1186 };
1187 let (batch, metadata) = match reader.read_next()? {
1188 Some(b) => b,
1189 None => return Ok(None),
1190 };
1191 reader.drain()?;
1192 let request_used_shm =
1195 metadata.contains_key(SHM_OFFSET_KEY) || metadata.contains_key(SHM_SEGMENT_NAME_KEY);
1196 #[cfg(feature = "shm")]
1204 let (batch, metadata) = if is_shm_pointer_batch(&batch, &metadata) {
1205 let one_shot: Option<ShmSegment>;
1206 let seg: Option<&ShmSegment> = match shm_cache {
1207 Some(cache) => {
1208 cache.refresh(&metadata);
1209 cache.segment()
1210 }
1211 None => {
1212 one_shot = maybe_attach_shm(&metadata);
1213 one_shot.as_ref()
1214 }
1215 };
1216 let resolved = resolve_shm_batch(batch, metadata, seg)?;
1217 if let (Some(off), Some(seg)) = (resolved.release_offset, seg) {
1221 let _ = seg.free(off);
1222 }
1223 (resolved.batch, resolved.metadata)
1224 } else {
1225 if let Some(cache) = shm_cache {
1226 cache.refresh(&metadata);
1227 }
1228 (batch, metadata)
1229 };
1230 #[cfg(not(feature = "shm"))]
1231 let _ = shm_cache;
1232 Ok(Some((batch, metadata, request_used_shm)))
1233 }
1234
1235 #[allow(clippy::too_many_arguments)]
1236 fn serve_unary<W: Write>(
1237 &self,
1238 w: &mut W,
1239 req: &Request,
1240 info: &MethodInfo,
1241 ctx: &CallContext,
1242 stats: &Arc<Mutex<crate::hooks::CallStatistics>>,
1243 app_err: &mut Option<RpcError>,
1244 #[cfg_attr(not(feature = "shm"), allow(unused_variables))] shm: Option<&ShmSegment>,
1245 ) -> Result<()> {
1246 let result = call_guard(|| (info.unary.as_ref().unwrap())(req, ctx)).and_then(|r| r);
1250 let logs = ctx.drain_logs();
1251 let mut envelope = EnvelopeMeta::new(&self.server_id, &req.request_id);
1252 match result {
1253 Ok(maybe_batch) => {
1254 let mut sw = StreamWriter::new(w, &info.result_schema)?;
1255 for log in logs {
1256 let md = envelope.log(&log);
1257 sw.write(&empty_batch(&info.result_schema)?, Some(md))?;
1258 }
1259 let out_batch = match maybe_batch {
1260 Some(b) => b,
1261 None => empty_batch(&info.result_schema)?,
1262 };
1263 {
1264 let mut s = lock_ok(stats);
1265 s.output_batches = 1;
1266 s.output_rows = out_batch.num_rows() as u64;
1267 }
1268 #[cfg(feature = "shm")]
1269 if let Some(seg) = shm {
1270 let (written, written_md) =
1271 maybe_write_to_shm(out_batch.clone(), Metadata::new(), Some(seg))?;
1272 if written_md.contains_key(crate::metadata::SHM_OFFSET_KEY) {
1273 sw.write(&written, Some(&written_md))?;
1274 sw.finish()?;
1275 return Ok(());
1276 }
1277 }
1278 #[cfg(feature = "http")]
1279 if let Some(cfg) = self.external_config.as_ref() {
1280 if let Ok(Some((ptr, md))) = crate::external::maybe_externalize_batch(
1285 &out_batch,
1286 &info.result_schema,
1287 None,
1288 cfg,
1289 ) {
1290 sw.write(&ptr, Some(&md))?;
1291 sw.finish()?;
1292 return Ok(());
1293 }
1294 }
1295 #[cfg(not(feature = "shm"))]
1296 let _ = shm;
1297 sw.write(&out_batch, None)?;
1298 sw.finish()?;
1299 }
1300 Err(err) => {
1301 let mut sw = StreamWriter::new(w, &info.result_schema)?;
1302 for log in logs {
1303 let md = envelope.log(&log);
1304 sw.write(&empty_batch(&info.result_schema)?, Some(md))?;
1305 }
1306 let md = envelope.error(&err);
1307 sw.write(&empty_batch(&info.result_schema)?, Some(md))?;
1308 sw.finish()?;
1309 *app_err = Some(err);
1310 }
1311 }
1312 Ok(())
1313 }
1314
1315 #[allow(clippy::too_many_arguments)]
1316 #[allow(clippy::too_many_arguments)]
1317 fn serve_stream<R: Read, W: Write>(
1318 &self,
1319 r: &mut R,
1320 w: &mut W,
1321 req: &Request,
1322 info: &MethodInfo,
1323 ctx: &CallContext,
1324 stats: &Arc<Mutex<crate::hooks::CallStatistics>>,
1325 app_err: &mut Option<RpcError>,
1326 #[cfg_attr(not(feature = "shm"), allow(unused_variables))] shm: Option<&ShmSegment>,
1327 ) -> Result<()> {
1328 let init_result = call_guard(|| (info.stream.as_ref().unwrap())(req, ctx)).and_then(|r| r);
1329 let init_logs = ctx.drain_logs();
1330 let stream = match init_result {
1331 Ok(s) => s,
1332 Err(err) => {
1333 let output_schema = info.result_schema.clone();
1335 let mut sw = StreamWriter::new(w, &output_schema)?;
1336 let mut envelope = EnvelopeMeta::new(&self.server_id, &req.request_id);
1337 for log in init_logs {
1338 let md = envelope.log(&log);
1339 sw.write(&empty_batch(&output_schema)?, Some(md))?;
1340 }
1341 let md = envelope.error(&err);
1342 sw.write(&empty_batch(&output_schema)?, Some(md))?;
1343 sw.finish()?;
1344 let _ = drain_input(r);
1347 *app_err = Some(err);
1348 return Ok(());
1349 }
1350 };
1351
1352 let StreamResult {
1353 output_schema,
1354 input_schema,
1355 state,
1356 header,
1357 header_metadata,
1358 } = stream;
1359
1360 let mut envelope = EnvelopeMeta::new(&self.server_id, &req.request_id);
1364
1365 let wrote_header = header.is_some();
1367 if let Some(header_batch) = header {
1368 let mut hw = StreamWriter::new(&mut *w, header_batch.schema().as_ref())?;
1369 for log in &init_logs {
1370 let md = envelope.log(log);
1371 hw.write(&empty_batch(header_batch.schema().as_ref())?, Some(md))?;
1372 }
1373 hw.write(&header_batch, header_metadata.as_ref())?;
1374 hw.finish()?;
1375 }
1376 let _ = w.flush();
1377
1378 let mut out_writer = StreamWriter::new(&mut *w, output_schema.as_ref())?;
1382 out_writer.flush()?;
1383
1384 let mut input_reader = StreamReader::new(&mut *r)?;
1386
1387 let empty_out = empty_batch(output_schema.as_ref())?;
1391
1392 if !wrote_header {
1394 for log in &init_logs {
1395 let md = envelope.log(log);
1396 out_writer.write(&empty_out, Some(md))?;
1397 }
1398 }
1399 let _ = header_metadata;
1400
1401 let mut state = state;
1402 let mut cancelled = false;
1403
1404 'lockstep: loop {
1405 let read = match input_reader.read_next() {
1406 Ok(x) => x,
1407 Err(_) => break,
1408 };
1409 let Some((input_batch, input_md)) = read else {
1410 break;
1411 };
1412
1413 #[cfg(feature = "shm")]
1418 let (input_batch, input_md) = {
1419 let resolved = resolve_shm_batch(input_batch, input_md, shm)?;
1420 if let (Some(off), Some(seg)) = (resolved.release_offset, shm) {
1421 let _ = seg.free(off);
1422 }
1423 (resolved.batch, resolved.metadata)
1424 };
1425
1426 {
1427 let mut s = lock_ok(stats);
1428 s.input_batches += 1;
1429 s.input_rows += input_batch.num_rows() as u64;
1430 }
1431
1432 let is_cancel = md_get(&input_md, CANCEL_KEY).is_some();
1434
1435 *lock_ok(&ctx.tick_metadata) = input_md;
1440
1441 if is_cancel {
1443 cancelled = true;
1444 let cancel_result = call_guard(|| match &mut state {
1445 StreamStateKind::Producer(p) => p.on_cancel(ctx),
1446 StreamStateKind::Exchange(e) => e.on_cancel(ctx),
1447 });
1448 for log in ctx.drain_logs() {
1449 let md = envelope.log(&log);
1450 out_writer.write(&empty_out, Some(md))?;
1451 }
1452 if let Err(err) = cancel_result {
1453 let md = envelope.error(&err);
1454 out_writer.write(&empty_out, Some(md))?;
1455 *app_err = Some(err);
1456 }
1457 break;
1458 }
1459
1460 let casted = match &input_schema {
1462 Some(expected) if input_batch.schema() != *expected => {
1463 match cast_batch(&input_batch, expected) {
1464 Ok(b) => b,
1465 Err(e) => {
1466 let md = envelope.error(&e);
1467 out_writer.write(&empty_out, Some(md))?;
1468 break 'lockstep;
1469 }
1470 }
1471 }
1472 _ => input_batch,
1473 };
1474
1475 let mut out = OutputCollector::new(output_schema.clone(), input_schema.is_none());
1476
1477 let iter_result = call_guard(|| match &mut state {
1478 StreamStateKind::Producer(p) => p.produce(&mut out, ctx),
1479 StreamStateKind::Exchange(e) => e.exchange(&casted, &mut out, ctx),
1480 })
1481 .and_then(|r| r);
1482
1483 let iter_logs = ctx.drain_logs();
1485 for log in iter_logs {
1486 let md = envelope.log(&log);
1487 out_writer.write(&empty_out, Some(md))?;
1488 }
1489
1490 if let Err(err) = iter_result {
1491 let md = envelope.error(&err);
1492 out_writer.write(&empty_out, Some(md))?;
1493 *app_err = Some(err);
1494 break;
1495 }
1496
1497 let finished = out.finished();
1498
1499 for item in out.items.drain(..) {
1501 match item {
1502 Emitted::Log(log) => {
1503 let md = envelope.log(&log);
1504 out_writer.write(&empty_out, Some(md))?;
1505 }
1506 Emitted::Batch { batch, metadata } => {
1507 {
1508 let mut s = lock_ok(stats);
1509 s.output_batches += 1;
1510 s.output_rows += batch.num_rows() as u64;
1511 }
1512 #[cfg(feature = "shm")]
1513 if let Some(seg) = shm {
1514 let md_in = metadata.clone().unwrap_or_default();
1515 let (written, written_md) =
1516 maybe_write_to_shm(batch.clone(), md_in, Some(seg))?;
1517 if written_md.contains_key(crate::metadata::SHM_OFFSET_KEY) {
1518 out_writer.write(&written, Some(&written_md))?;
1519 continue;
1520 }
1521 }
1522 #[cfg(feature = "http")]
1523 if let Some(cfg) = self.external_config.as_ref() {
1524 match crate::external::maybe_externalize_batch(
1525 &batch,
1526 output_schema.as_ref(),
1527 metadata.as_ref(),
1528 cfg,
1529 ) {
1530 Ok(Some((ptr, md))) => {
1531 out_writer.write(&ptr, Some(&md))?;
1532 continue;
1533 }
1534 Ok(None) => {}
1535 Err(e) => {
1536 *app_err = Some(e);
1539 }
1540 }
1541 }
1542 out_writer.write(&batch, metadata.as_ref())?;
1543 }
1544 }
1545 }
1546 out_writer.flush()?;
1549
1550 if finished {
1551 break;
1552 }
1553 }
1554 let _ = cancelled;
1555 out_writer.finish()?;
1556
1557 let _ = input_reader.drain();
1559 Ok(())
1560 }
1561}
1562
1563fn drain_input<R: Read>(r: &mut R) -> Result<()> {
1564 let mut rdr = StreamReader::new(r)?;
1565 rdr.drain()?;
1566 Ok(())
1567}
1568
1569pub(crate) fn cast_batch(batch: &RecordBatch, target: &SchemaRef) -> Result<RecordBatch> {
1570 if batch.num_columns() != target.fields().len() {
1571 return Err(RpcError::type_error(format!(
1572 "Input schema mismatch: expected {} fields, got {}",
1573 target.fields().len(),
1574 batch.num_columns()
1575 )));
1576 }
1577 let src_schema = batch.schema();
1578 for (i, field) in target.fields().iter().enumerate() {
1579 let src_name = src_schema.field(i).name();
1580 if src_name != field.name() {
1581 return Err(RpcError::type_error(format!(
1582 "Input schema mismatch: expected field {:?}, got {:?}",
1583 field.name(),
1584 src_name
1585 )));
1586 }
1587 }
1588 let opts = arrow_cast::CastOptions::default();
1589 let mut cols = Vec::with_capacity(batch.num_columns());
1590 for (i, field) in target.fields().iter().enumerate() {
1591 let src = batch.column(i);
1592 if src.data_type() == field.data_type() {
1593 cols.push(src.clone());
1594 continue;
1595 }
1596 let c = cast_with_options(src.as_ref(), field.data_type(), &opts)
1597 .map_err(|e| RpcError::type_error(format!("cast field {}: {}", field.name(), e)))?;
1598 cols.push(c);
1599 }
1600 RecordBatch::try_new(target.clone(), cols).map_err(RpcError::from)
1603}
1604
1605pub(crate) struct EnvelopeMeta<'a> {
1616 server_id: &'a str,
1617 request_id: &'a str,
1618 md: Option<Metadata>,
1621}
1622
1623impl<'a> EnvelopeMeta<'a> {
1624 pub(crate) fn new(server_id: &'a str, request_id: &'a str) -> Self {
1625 Self {
1626 server_id,
1627 request_id,
1628 md: None,
1629 }
1630 }
1631
1632 fn map(&mut self) -> &mut Metadata {
1635 if self.md.is_none() {
1636 let mut md = Metadata::with_capacity(5);
1637 if !self.server_id.is_empty() {
1638 md.insert(SERVER_ID_KEY.to_string(), self.server_id.to_string());
1639 }
1640 if !self.request_id.is_empty() {
1641 md.insert(REQUEST_ID_KEY.to_string(), self.request_id.to_string());
1642 }
1643 self.md = Some(md);
1644 }
1645 self.md.as_mut().unwrap()
1646 }
1647
1648 fn set(&mut self, key: &'static str, val: String) {
1651 let md = self.map();
1652 if let Some(slot) = md.get_mut(key) {
1653 *slot = val;
1654 } else {
1655 md.insert(key.to_string(), val);
1656 }
1657 }
1658
1659 pub(crate) fn log(&mut self, msg: &LogMessage) -> &Metadata {
1661 self.set(LOG_LEVEL_KEY, msg.level.as_str().to_string());
1662 self.set(LOG_MESSAGE_KEY, msg.message.clone());
1663 if !msg.extras.is_empty() {
1664 self.set(LOG_EXTRA_KEY, msg.extras_json());
1665 } else {
1666 self.map().remove(LOG_EXTRA_KEY);
1668 }
1669 self.md.as_ref().unwrap()
1670 }
1671
1672 pub(crate) fn error(&mut self, err: &RpcError) -> &Metadata {
1674 let extra = serde_json::json!({
1675 "exception_type": err.error_type,
1676 "exception_message": err.message,
1677 "traceback": err.traceback,
1678 })
1679 .to_string();
1680 self.set(LOG_LEVEL_KEY, "EXCEPTION".to_string());
1681 self.set(LOG_MESSAGE_KEY, err.message.clone());
1682 self.set(LOG_EXTRA_KEY, extra);
1683 self.md.as_ref().unwrap()
1684 }
1685}
1686
1687#[cfg(feature = "http")]
1691pub(crate) fn build_log_metadata(msg: &LogMessage, server_id: &str, request_id: &str) -> Metadata {
1692 let mut e = EnvelopeMeta::new(server_id, request_id);
1693 e.log(msg);
1694 e.md.unwrap()
1695}
1696
1697pub(crate) fn build_error_metadata(err: &RpcError, server_id: &str, request_id: &str) -> Metadata {
1698 let mut e = EnvelopeMeta::new(server_id, request_id);
1699 e.error(err);
1700 e.md.unwrap()
1701}
1702
1703pub(crate) fn write_error_stream<W: Write>(
1705 w: &mut W,
1706 schema: &Schema,
1707 err: &RpcError,
1708 server_id: &str,
1709 request_id: &str,
1710) -> Result<()> {
1711 let mut sw = StreamWriter::new(w, schema)?;
1712 let md = build_error_metadata(err, server_id, request_id);
1713 sw.write(&empty_batch(schema)?, Some(&md))?;
1714 sw.finish()?;
1715 Ok(())
1716}
1717
1718#[cfg(test)]
1719mod tests {
1720 use super::*;
1721 use std::io::Cursor;
1722 use std::sync::atomic::{AtomicBool, Ordering};
1723
1724 fn request_bytes(method: &str) -> Vec<u8> {
1727 let schema = empty_schema();
1728 let batch = empty_batch(&schema).unwrap();
1729 let mut buf = Vec::new();
1730 {
1731 let mut w = StreamWriter::new(&mut buf, &schema).unwrap();
1732 let mut md = Metadata::new();
1733 md.insert(RPC_METHOD_KEY.into(), method.into());
1734 md.insert(REQUEST_VERSION_KEY.into(), REQUEST_VERSION.into());
1735 md.insert(REQUEST_ID_KEY.into(), format!("req-{method}"));
1736 w.write(&batch, Some(&md)).unwrap();
1737 w.finish().unwrap();
1738 }
1739 buf
1740 }
1741
1742 #[test]
1743 fn panicking_handler_yields_error_envelope_and_loop_survives() {
1744 let mut server = RpcServer::new("test-srv");
1745 server.register(MethodInfo::unary(
1746 "boom",
1747 empty_schema(),
1748 empty_schema(),
1749 |_req, _ctx| panic!("handler exploded"),
1750 ));
1751 let ran_second = Arc::new(AtomicBool::new(false));
1752 let flag = ran_second.clone();
1753 server.register(MethodInfo::unary(
1754 "ok",
1755 empty_schema(),
1756 empty_schema(),
1757 move |_req, _ctx| {
1758 flag.store(true, Ordering::SeqCst);
1759 Ok(None)
1760 },
1761 ));
1762
1763 let mut input = request_bytes("boom");
1766 input.extend(request_bytes("ok"));
1767 let mut output: Vec<u8> = Vec::new();
1768 server.serve(Cursor::new(input), &mut output);
1769
1770 assert!(
1771 ran_second.load(Ordering::SeqCst),
1772 "serve loop aborted after a handler panic"
1773 );
1774
1775 let mut r = StreamReader::new(output.as_slice()).unwrap();
1778 let (_b, md) = r.read_next().unwrap().expect("error batch");
1779 assert_eq!(md_get(&md, LOG_LEVEL_KEY), Some("EXCEPTION"));
1780 }
1781
1782 #[test]
1783 fn transport_options_reports_shm_capability_unregistered() {
1784 use crate::metadata::TRANSPORT_SHM_KEY;
1785 use crate::transport_options::{shm_available, TRANSPORT_OPTIONS_METHOD_NAME};
1786
1787 let mut server = RpcServer::new("test-srv");
1788 server.register(MethodInfo::unary(
1789 "noop",
1790 empty_schema(),
1791 empty_schema(),
1792 |_req, _ctx| Ok(None),
1793 ));
1794 assert!(!server.methods.contains_key(TRANSPORT_OPTIONS_METHOD_NAME));
1796
1797 let input = request_bytes(TRANSPORT_OPTIONS_METHOD_NAME);
1798 let mut output: Vec<u8> = Vec::new();
1799 server.serve(Cursor::new(input), &mut output);
1800
1801 let mut r = StreamReader::new(output.as_slice()).unwrap();
1802 let (_b, md) = r.read_next().unwrap().expect("transport options batch");
1803 let expected = if shm_available() { "true" } else { "false" };
1804 assert_eq!(md_get(&md, TRANSPORT_SHM_KEY), Some(expected));
1805 assert_eq!(md_get(&md, REQUEST_VERSION_KEY), Some(REQUEST_VERSION));
1806 assert_eq!(md_get(&md, SERVER_ID_KEY), Some("test-srv"));
1807 }
1808
1809 #[cfg(feature = "shm")]
1819 mod shm_requests {
1820 use super::*;
1821 use crate::metadata::{SHM_OFFSET_KEY, SHM_SEGMENT_NAME_KEY, SHM_SEGMENT_SIZE_KEY};
1822 use crate::shm::{
1823 is_shm_pointer_batch, make_shm_pointer_batch, maybe_write_to_shm, ShmSegment,
1824 };
1825 use arrow_array::{BinaryArray, Int64Array};
1826 use arrow_schema::{DataType, Field};
1827
1828 fn params_schema() -> SchemaRef {
1829 Arc::new(Schema::new(vec![Field::new(
1830 "request",
1831 DataType::Binary,
1832 false,
1833 )]))
1834 }
1835
1836 fn result_schema() -> SchemaRef {
1837 Arc::new(Schema::new(vec![Field::new("n", DataType::Int64, false)]))
1838 }
1839
1840 fn request_batch(payload: &[u8]) -> RecordBatch {
1841 RecordBatch::try_new(
1842 params_schema(),
1843 vec![Arc::new(BinaryArray::from(vec![Some(payload)]))],
1844 )
1845 .unwrap()
1846 }
1847
1848 fn dispatch_md(seg: Option<&ShmSegment>) -> Metadata {
1851 let mut md = Metadata::new();
1852 md.insert(RPC_METHOD_KEY.into(), "do_thing".into());
1853 md.insert(REQUEST_VERSION_KEY.into(), REQUEST_VERSION.into());
1854 if let Some(seg) = seg {
1855 md.insert(SHM_SEGMENT_NAME_KEY.into(), seg.name().to_string());
1856 md.insert(SHM_SEGMENT_SIZE_KEY.into(), seg.size().to_string());
1857 }
1858 md
1859 }
1860
1861 fn pointer_request(seg: &ShmSegment, payload: &[u8], advertise: bool) -> Vec<u8> {
1866 let md = dispatch_md(advertise.then_some(seg));
1867 let (ptr, ptr_md) = maybe_write_to_shm(request_batch(payload), md, Some(seg)).unwrap();
1868 assert!(
1869 is_shm_pointer_batch(&ptr, &ptr_md),
1870 "request batch should have routed through shm"
1871 );
1872 let mut buf = Vec::new();
1873 {
1874 let mut w = StreamWriter::new(&mut buf, ptr.schema().as_ref()).unwrap();
1875 w.write(&ptr, Some(&ptr_md)).unwrap();
1876 w.finish().unwrap();
1877 }
1878 buf
1879 }
1880
1881 fn inline_request(payload: &[u8], seg: Option<&ShmSegment>) -> Vec<u8> {
1883 let batch = request_batch(payload);
1884 let md = dispatch_md(seg);
1885 let mut buf = Vec::new();
1886 {
1887 let mut w = StreamWriter::new(&mut buf, batch.schema().as_ref()).unwrap();
1888 w.write(&batch, Some(&md)).unwrap();
1889 w.finish().unwrap();
1890 }
1891 buf
1892 }
1893
1894 fn payload_server(seen: Arc<Mutex<Vec<Vec<u8>>>>) -> RpcServer {
1898 let mut server = RpcServer::new("shm-srv");
1899 let rs = result_schema();
1900 server.register(MethodInfo::unary(
1901 "do_thing",
1902 params_schema(),
1903 result_schema(),
1904 move |req, _ctx| {
1905 let col = req
1906 .column("request")
1907 .expect("request column")
1908 .as_any()
1909 .downcast_ref::<BinaryArray>()
1910 .unwrap();
1911 lock_ok(&seen).push(col.value(0).to_vec());
1912 Ok(Some(RecordBatch::try_new(
1913 rs.clone(),
1914 vec![Arc::new(Int64Array::from(vec![1i64]))],
1915 )?))
1916 },
1917 ));
1918 server
1919 }
1920
1921 fn response_metadata(output: &[u8]) -> Vec<Metadata> {
1924 let mut out = Vec::new();
1925 let mut cursor = Cursor::new(output);
1926 while (cursor.position() as usize) < output.len() {
1927 let mut reader = StreamReader::new(&mut cursor).unwrap();
1928 while let Some((_b, md)) = reader.read_next().unwrap() {
1929 out.push(md);
1930 }
1931 }
1932 out
1933 }
1934
1935 #[test]
1939 fn pointer_request_batch_resolves_via_segment_named_in_metadata() {
1940 let seg = ShmSegment::create(1024 * 1024).unwrap();
1941 let payload = b"serialized-request-blob";
1942
1943 let seen = Arc::new(Mutex::new(Vec::new()));
1944 let server = payload_server(seen.clone());
1945 let mut output: Vec<u8> = Vec::new();
1946 server.serve(
1947 Cursor::new(pointer_request(&seg, payload, true)),
1948 &mut output,
1949 );
1950 assert_eq!(lock_ok(&seen).as_slice(), &[payload.to_vec()]);
1951
1952 let seen2 = Arc::new(Mutex::new(Vec::new()));
1954 let server2 = payload_server(seen2.clone());
1955 let mut input = Cursor::new(pointer_request(&seg, payload, true));
1956 let mut output2: Vec<u8> = Vec::new();
1957 assert!(server2.serve_one(&mut input, &mut output2).unwrap());
1958 assert_eq!(lock_ok(&seen2).as_slice(), &[payload.to_vec()]);
1959 }
1960
1961 #[test]
1964 fn serve_caches_client_segment_for_offset_only_requests() {
1965 let seg = ShmSegment::create(1024 * 1024).unwrap();
1966
1967 let mut input = inline_request(b"first", Some(&seg));
1969
1970 let (off, len) = seg
1972 .allocate_and_write(&request_batch(b"second"))
1973 .unwrap()
1974 .expect("payload fits");
1975 let (ptr, mut ptr_md) =
1976 make_shm_pointer_batch(params_schema().as_ref(), off, len).unwrap();
1977 ptr_md.insert(RPC_METHOD_KEY.into(), "do_thing".into());
1978 ptr_md.insert(REQUEST_VERSION_KEY.into(), REQUEST_VERSION.into());
1979 {
1980 let mut w = StreamWriter::new(&mut input, ptr.schema().as_ref()).unwrap();
1981 w.write(&ptr, Some(&ptr_md)).unwrap();
1982 w.finish().unwrap();
1983 }
1984
1985 let seen = Arc::new(Mutex::new(Vec::new()));
1986 let server = payload_server(seen.clone());
1987 let mut output: Vec<u8> = Vec::new();
1988 server.serve(Cursor::new(input), &mut output);
1989 assert_eq!(
1990 lock_ok(&seen).as_slice(),
1991 &[b"first".to_vec(), b"second".to_vec()]
1992 );
1993 }
1994
1995 #[test]
1998 fn pointer_request_without_segment_trips_single_row_guard() {
1999 let seg = ShmSegment::create(1024 * 1024).unwrap();
2000 let seen = Arc::new(Mutex::new(Vec::new()));
2001 let server = payload_server(seen.clone());
2002 let mut output: Vec<u8> = Vec::new();
2003 server.serve(
2004 Cursor::new(pointer_request(&seg, b"orphan", false)),
2005 &mut output,
2006 );
2007 assert!(lock_ok(&seen).is_empty(), "guard should reject dispatch");
2008 }
2009
2010 #[test]
2016 fn response_routed_through_shm_only_when_request_signalled_shm() {
2017 let seg = ShmSegment::create(1024 * 1024).unwrap();
2018
2019 let mut input = inline_request(b"a", Some(&seg));
2022 input.extend(inline_request(b"b", None));
2023
2024 let seen = Arc::new(Mutex::new(Vec::new()));
2025 let server = payload_server(seen.clone());
2026 let mut output: Vec<u8> = Vec::new();
2027 server.serve(Cursor::new(input), &mut output);
2028 assert_eq!(lock_ok(&seen).len(), 2);
2029
2030 let mds = response_metadata(&output);
2031 assert_eq!(mds.len(), 2, "one data batch per response");
2032 assert!(
2033 mds[0].contains_key(SHM_OFFSET_KEY),
2034 "response A (segment advertised) should route through shm"
2035 );
2036 assert!(
2037 !mds[1].contains_key(SHM_OFFSET_KEY),
2038 "response B (no shm signal) must stay inline despite the cached segment"
2039 );
2040 }
2041 }
2042}